feat: 优化模特上身图提示词

This commit is contained in:
hh
2026-07-14 15:18:54 +08:00
parent da58f99293
commit 0a58b89919
8 changed files with 1973 additions and 15 deletions
@@ -0,0 +1,271 @@
"""模特上身图 V2 同输入 A/B:默认只预估,显式 --execute 才创建计费任务。"""
from __future__ import annotations
import json
import uuid
from decimal import Decimal, InvalidOperation
from pathlib import Path
from time import monotonic
import requests
from django.core.management.base import BaseCommand, CommandError
class Command(BaseCommand):
help = "对同一商品、模特、提示词和比例提交 Seedream/GPT 模特上身图 A/B;默认 dry-run"
def add_arguments(self, parser):
parser.add_argument("--product", default="", help="商品 UUID")
parser.add_argument("--model-asset", default="", help="模特肖像 Asset UUID")
parser.add_argument("--model-entity", default="", help="可选:模特库 Model UUID,仅作溯源")
parser.add_argument("--prompt", default="", help="两组共同使用的用户提示词")
parser.add_argument("--ratio", default="3:4", help="两组共同使用的比例,默认 3:4")
parser.add_argument("--count", type=int, default=1, choices=(1, 2, 4), help="每个模型生成张数")
parser.add_argument(
"--variants",
nargs="+",
default=("volcano", "gpt-image"),
help="模型路由键,默认 volcano gpt-image",
)
parser.add_argument("--max-points", default="", help="硬预算上限;--execute 时必填")
parser.add_argument("--experiment", default="", help="可选实验 UUID;默认自动生成")
parser.add_argument("--execute", action="store_true", help="真正预扣积分并提交任务")
parser.add_argument("--report", default="", help="只读查看指定实验 UUID 的任务与结果")
parser.add_argument("--export-dir", default="", help="配合 --report 匿名导出 A/B 图片与独立答案文件")
def handle(self, *args, **opts):
if opts["report"]:
self._report(opts["report"], opts["export_dir"])
return
from apps.ai.services import resolve_image_model
from apps.assets.models import Asset
from apps.billing.models import CreditAccount
from apps.billing.pricing import quote_flat
from apps.products.models import Product
if not opts["product"] or not opts["model_asset"] or not opts["prompt"].strip():
raise CommandError("dry-run/execute 必须提供 --product、--model-asset 和 --prompt")
product = Product.objects.select_related("team", "created_by", "team__owner").filter(
id=opts["product"], purged_at__isnull=True
).first()
if product is None:
raise CommandError("商品不存在或已清理")
model_asset = Asset.objects.filter(
id=opts["model_asset"],
team=product.team,
is_deleted=False,
purged_at__isnull=True,
).first()
if model_asset is None:
raise CommandError("模特肖像不存在、已删除或不属于商品团队")
from apps.ai.services import _asset_preview_url, _product_reference_urls
product_reference_count = len(_product_reference_urls(product, limit=3))
if product_reference_count < 1:
raise CommandError("商品没有可用的真实参考图,不能执行同输入上身图 A/B")
if not _asset_preview_url(model_asset):
raise CommandError("模特肖像没有可用图片文件")
user = product.created_by or product.team.owner
if user is None:
raise CommandError("商品团队没有可用于计费的提交用户")
variants: list[tuple[str, object, Decimal]] = []
seen: set[str] = set()
for key in opts["variants"]:
key = str(key).strip()
if not key or key in seen:
continue
seen.add(key)
model = resolve_image_model(key)
if model is None:
raise CommandError(f"找不到图像模型路由:{key}")
quote = quote_flat(model, team=product.team)
variants.append((key, model, quote.points))
if len(variants) < 2:
raise CommandError("A/B 至少需要两个不同模型路由")
total_points = sum((points * opts["count"] for _, _, points in variants), Decimal("0"))
account, _ = CreditAccount.objects.get_or_create(team=product.team)
available = account.balance - account.reserved_balance
experiment_id = self._experiment_id(opts["experiment"])
self.stdout.write(f"experiment={experiment_id}")
self.stdout.write(
f"product={product.id} {product.title} · model_asset={model_asset.id} · "
f"product_refs={product_reference_count} · ratio={opts['ratio']} · count/model={opts['count']}"
)
for key, model, points in variants:
self.stdout.write(
f" {key}: {model.provider.name}:{model.name} · {points} 积分/张 · "
f"小计 {points * opts['count']}"
)
self.stdout.write(f"总预算={total_points} 积分 · 当前可用={available} 积分")
if not opts["execute"]:
self.stdout.write(self.style.WARNING("DRY-RUN:未创建任务、未预扣积分;加 --execute --max-points N 才会提交"))
return
max_points = self._max_points(opts["max_points"])
if total_points > max_points:
raise CommandError(f"预计 {total_points} 积分,超过硬预算 {max_points},已拒绝提交")
if available < total_points:
raise CommandError(f"当前可用 {available} 积分,低于预计 {total_points},已拒绝提交")
from apps.ai.services import enqueue_standalone_images
submitted = []
for key, model, points in variants:
tasks = enqueue_standalone_images(
team=product.team,
user=user,
prompt=opts["prompt"].strip(),
mode="model",
count=opts["count"],
product_id=str(product.id),
model_id=str(model_asset.id),
model_entity_id=str(opts["model_entity"] or "") or None,
ratio=str(opts["ratio"] or "").strip() or None,
image_model=key,
tryon_prompt_v2_override=True,
tryon_ab={"experiment_id": experiment_id, "variant": key},
dispatch=False,
)
submitted.append((key, model, points, tasks))
for key, model, points, tasks in submitted:
batch_id = (tasks[0].request_payload or {}).get("batch_id") if tasks else ""
ids = ",".join(str(task.id) for task in tasks)
self.stdout.write(
self.style.SUCCESS(
f"RESERVED {key} · model={model.provider.name}:{model.name} · "
f"batch={batch_id} · tasks={ids} · reserved={points * len(tasks)}"
)
)
# 不投递在线 Worker:它可能仍是未部署 V2 的旧版本。由当前命令进程串行执行,确保代码版本一致。
from apps.ai.services import run_standalone_image_task
for key, model, _points, tasks in submitted:
for task in tasks:
started = monotonic()
run_standalone_image_task(task_id=str(task.id))
elapsed = round(monotonic() - started, 3)
task.refresh_from_db()
payload = dict(task.request_payload or {})
ab = dict(payload.get("tryon_ab") or {})
ab["elapsed_seconds"] = elapsed
payload["tryon_ab"] = ab
task.request_payload = payload
task.save(update_fields=["request_payload", "updated_at"])
assets = ",".join(str(asset.id) for asset in task.generated_assets.all()) or "-"
self.stdout.write(
f"FINISHED {key} index={payload.get('index')} status={task.status} "
f"elapsed={elapsed}s actual={task.actual_cost} asset={assets} "
f"error={task.error_code or '-'}"
)
def _max_points(self, raw: str) -> Decimal:
if not str(raw or "").strip():
raise CommandError("--execute 必须同时提供正数 --max-points")
try:
value = Decimal(str(raw))
except InvalidOperation as exc:
raise CommandError("--max-points 必须是数字") from exc
if value <= 0:
raise CommandError("--max-points 必须大于 0")
return value
def _experiment_id(self, raw: str) -> str:
if not str(raw or "").strip():
return str(uuid.uuid4())
try:
return str(uuid.UUID(str(raw)))
except ValueError as exc:
raise CommandError("--experiment 必须是合法 UUID") from exc
def _report(self, experiment_id: str, export_dir: str = "") -> None:
from apps.ai.models import AITask
experiment_id = self._experiment_id(experiment_id)
tasks = list(
AITask.objects.filter(request_payload__tryon_ab__experiment_id=experiment_id)
.select_related("model_config", "model_config__provider")
.prefetch_related("generated_assets", "generated_assets__files")
.order_by("created_at")
)
if not tasks:
raise CommandError("找不到该 A/B 实验")
self.stdout.write(f"experiment={experiment_id} · tasks={len(tasks)}")
for task in tasks:
payload = task.request_payload or {}
ab = payload.get("tryon_ab") or {}
trace = payload.get("tryon_prompt") or {}
assets = ",".join(str(asset.id) for asset in task.generated_assets.all()) or "-"
self.stdout.write(
f" {ab.get('variant')} index={payload.get('index')} status={task.status} "
f"model={task.model_config.provider.name}:{task.model_config.name} "
f"estimated={task.estimated_cost} actual={task.actual_cost} "
f"prompt_v2={trace.get('applied')} elapsed={ab.get('elapsed_seconds')}s "
f"asset={assets} error={task.error_code or '-'}"
)
if export_dir:
self._blind_export(experiment_id, tasks, export_dir)
def _blind_export(self, experiment_id: str, tasks: list, export_dir: str) -> None:
from apps.ai.services import _asset_preview_url
output = Path(export_dir).expanduser().resolve()
output.mkdir(parents=True, exist_ok=True)
variants = sorted({str((task.request_payload.get("tryon_ab") or {}).get("variant") or "") for task in tasks})
variants = [variant for variant in variants if variant]
if len(variants) < 2:
raise CommandError("匿名导出至少需要两个有效 A/B 变体")
# 实验 UUID 决定匿名标签顺序;报告正文不输出映射,评分完成后再单独读取 answer_key.json。
if uuid.UUID(experiment_id).int % 2:
variants.reverse()
labels = {variant: chr(ord("A") + index) for index, variant in enumerate(variants)}
manifest = {"experiment_id": experiment_id, "samples": []}
answer_key = {"experiment_id": experiment_id, "labels": {label: variant for variant, label in labels.items()}}
for task in tasks:
payload = task.request_payload or {}
variant = str((payload.get("tryon_ab") or {}).get("variant") or "")
label = labels.get(variant)
asset = task.generated_assets.first()
if not label or asset is None:
continue
url = _asset_preview_url(asset)
if not url:
raise CommandError(f"任务 {task.id} 的成图没有可下载地址")
primary = asset.files.filter(is_primary=True).first() or asset.files.first()
content_type = str(getattr(primary, "content_type", "") or "")
suffix = ".jpg" if "jpeg" in content_type else (".webp" if "webp" in content_type else ".png")
shot_index = int(payload.get("index") or 0)
filename = f"{label}_{shot_index + 1}{suffix}"
response = requests.get(url, timeout=60)
response.raise_for_status()
(output / filename).write_bytes(response.content)
manifest["samples"].append(
{
"label": label,
"shot_index": shot_index,
"filename": filename,
"task_id": str(task.id),
"asset_id": str(asset.id),
}
)
(output / "blind_manifest.json").write_text(
json.dumps(manifest, ensure_ascii=False, indent=2), encoding="utf-8"
)
(output / "answer_key.json").write_text(
json.dumps(answer_key, ensure_ascii=False, indent=2), encoding="utf-8"
)
self.stdout.write(
self.style.SUCCESS(
f"BLIND_EXPORT samples={len(manifest['samples'])} dir={output} · 评分前不要读取 answer_key.json"
)
)