feat: 优化模特上身图多图与裤长策略
This commit is contained in:
@@ -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(
|
||||
|
||||
Reference in New Issue
Block a user