feat: 优化模特上身图多图与裤长策略

This commit is contained in:
hh
2026-07-14 17:16:38 +08:00
parent 2b26c93031
commit 1a40a8cb4b
4 changed files with 425 additions and 66 deletions
+135 -4
View File
@@ -171,7 +171,7 @@ class NormalizeDraftTests(SimpleTestCase):
from io import BytesIO
from unittest.mock import patch
from unittest.mock import MagicMock, patch
from django.test import TestCase, override_settings
@@ -425,7 +425,7 @@ class StandaloneImageReferenceTests(TestCase):
used_prompts = [call.kwargs["prompt"] for call in prov.image_edit.call_args_list]
used_sizes = [call.kwargs["size"] for call in prov.image_edit.call_args_list]
self.assertTrue(all(user_prompt in prompt for prompt in used_prompts))
self.assertTrue(all("参考图角色(最高优先级" in prompt for prompt in used_prompts))
self.assertTrue(all("参考图角色(固定语义" in prompt for prompt in used_prompts))
self.assertTrue(all("严格使用 7:10 比例" in prompt for prompt in used_prompts))
self.assertTrue(all("半身近景" not in prompt and "手持或局部" not in prompt for prompt in used_prompts))
self.assertNotEqual(used_prompts[0], used_prompts[1])
@@ -434,7 +434,7 @@ class StandaloneImageReferenceTests(TestCase):
for index, task in enumerate(tasks):
trace = task.request_payload["tryon_prompt"]
self.assertTrue(trace["applied"])
self.assertEqual(trace["version"], "v2.2")
self.assertEqual(trace["version"], "v2.5")
self.assertEqual(trace["rollout_source"], "global")
self.assertEqual(trace["effective_prompt"], used_prompts[index])
self.assertEqual(trace["shot_index"], index)
@@ -446,6 +446,137 @@ class StandaloneImageReferenceTests(TestCase):
self.assertIsNone(trace["reference_roles"]["model_triview_number"])
self.assertEqual(task.request_payload["tryon_classification"]["kind"], "lower")
@override_settings(MODEL_TRYON_PROMPT_V2_ENABLED=True)
def test_model_tryon_v25_zero_cost_model_count_ratio_matrix(self):
"""Mock 全矩阵:两个模型都覆盖 1/2/4 张与预设/自定义比例,不调用真实供应商。"""
CreditAccount.objects.filter(team=self.team).update(balance="5000.0000")
self.product.title = "舒适柔软面料女裤"
self.product.category = "服饰内衣"
self.product.save(update_fields=["title", "category"])
model_asset = Asset.objects.create(
team=self.team,
created_by=self.user,
name="矩阵模特",
asset_type=Asset.Type.IMAGE,
source=Asset.Source.AI_GENERATED,
category=Asset.Category.PERSON,
)
AssetFile.objects.create(
asset=model_asset,
object_key="matrix-model.png",
bucket="b",
content_type="image/png",
preview_url="http://x/matrix-model.png",
is_primary=True,
)
volcano, _ = ModelProvider.objects.get_or_create(
name="volcengine",
defaults={"display_name": "火山"},
)
seedream_config = ModelConfig.objects.create(
provider=volcano,
name="seedream-v25-matrix",
display_name="Seedream V2.5 matrix",
capability=ModelConfig.Capability.IMAGE,
unit_price="1.0000",
)
gpt_config = ModelConfig.objects.filter(
capability=ModelConfig.Capability.IMAGE,
name__icontains="gpt-image",
).first()
self.assertIsNotNone(gpt_config)
seedream = MagicMock()
del seedream.image_edit
seedream.image_generation.return_value = {"data": [{"url": "http://x/seedream-matrix.png"}]}
seedream.extract_first_media_url.return_value = "http://x/seedream-matrix.png"
gpt = MagicMock()
gpt.image_edit.return_value = {"data": [{"url": "http://x/gpt-matrix.png"}]}
gpt.extract_first_media_url.return_value = "http://x/gpt-matrix.png"
def provider_for(model_config, *args, **kwargs):
del args, kwargs
return seedream if model_config.id == seedream_config.id else gpt
stored = MagicMock(object_key="matrix.png", bucket="b", content_type="image/png", size_bytes=3)
ratios = ("1:1", "3:4", "9:16", "7:10")
models = (
(seedream_config, seedream.image_generation, "image"),
(gpt_config, gpt.image_edit, "images"),
)
with (
patch("apps.ai.services.get_image_provider", side_effect=provider_for),
patch(
"apps.ai.services.VolcanoArkProvider.media_to_bytes",
return_value=(BytesIO(b"img"), "image/png"),
),
patch("apps.ai.services.TosStorage") as storage,
):
storage.return_value.upload_fileobj.return_value = stored
for model_config, provider_method, reference_key in models:
for count in (1, 2, 4):
for ratio in ratios:
with self.subTest(model=model_config.name, count=count, ratio=ratio):
provider_method.reset_mock()
tasks = enqueue_standalone_images(
team=self.team,
user=self.user,
prompt="模特上身展示,自然光,真实质感,电商主图",
mode="model",
count=count,
product_id=str(self.product.id),
model_id=str(model_asset.id),
ratio=ratio,
image_model=f"{model_config.provider.name}:{model_config.name}",
)
self.assertEqual(len(tasks), count)
self.assertEqual(len(provider_method.call_args_list), count)
for index, (task, call) in enumerate(zip(tasks, provider_method.call_args_list)):
task.refresh_from_db()
trace = task.request_payload["tryon_prompt"]
effective_prompt = call.kwargs["prompt"]
self.assertEqual(task.model_config_id, model_config.id)
self.assertEqual(trace["version"], "v2.5")
self.assertTrue(trace["applied"])
self.assertEqual(trace["rollout_source"], "global")
self.assertEqual(trace["shot_index"], index)
self.assertEqual(trace["batch_count"], count)
self.assertEqual(trace["requested_ratio"], ratio)
self.assertEqual(trace["resolved_ratio"], ratio)
self.assertEqual(trace["ratio_source"], "user")
self.assertEqual(trace["effective_prompt"], effective_prompt)
self.assertEqual(trace["reference_roles"]["product_numbers"], [1])
self.assertEqual(trace["reference_roles"]["model_portrait_number"], 2)
self.assertIn(f"严格使用 {ratio} 比例", effective_prompt)
self.assertIn("裤长人体落点", effective_prompt)
self.assertNotIn("背面", effective_prompt)
self.assertEqual(
call.kwargs[reference_key],
["http://x/cover.png", "http://x/matrix-model.png"],
)
for model_config, provider_method, _reference_key in models:
with self.subTest(model=model_config.name, explicit_back=True):
provider_method.reset_mock()
tasks = enqueue_standalone_images(
team=self.team,
user=self.user,
prompt="模特上身展示,完整展示后袋",
mode="model",
count=4,
product_id=str(self.product.id),
model_id=str(model_asset.id),
ratio="3:4",
image_model=f"{model_config.provider.name}:{model_config.name}",
)
prompts = [call.kwargs["prompt"] for call in provider_method.call_args_list]
self.assertEqual(len(tasks), 4)
self.assertEqual(len(prompts), 4)
self.assertTrue(all("背面全身站姿" not in prompt for prompt in prompts[:3]))
self.assertIn("背面全身站姿", prompts[3])
@override_settings(
MODEL_TRYON_PROMPT_V2_ENABLED=False,
MODEL_TRYON_PROMPT_V2_CANARY_TEAM_IDS=frozenset(),
@@ -492,7 +623,7 @@ class StandaloneImageReferenceTests(TestCase):
task = AITask.objects.get(id=submitted[0].id)
trace = task.request_payload["tryon_prompt"]
self.assertTrue(trace["applied"])
self.assertEqual(trace["version"], "v2.2")
self.assertEqual(trace["version"], "v2.5")
self.assertEqual(trace["rollout_source"], "canary")
@override_settings(