Files
pod_trend_agent/graph/nodes/seed_shot_node.py
T
3218485270 5ab5cf6586 v110-v112 自定义模式完善 + 模板导出增强 + 多模态兼容优化
- 自定义模式:分析模型输出 delta 唯一改动指令,生图模板 custom_image_prompt.md({delta} 占位符),不再使用负向提示词;generate_design 按 custom_mode 分支,Pinterest 模式保留原创化指令,两模式互不影响
- 多模态分析 response_format 三级回退(json_schema → json_object → none),兼容 DeepSeek
- 模板导出:details 扩展列(细节1/2/3)、target_audience 扩展列(适用人群1)、固定值风格1=休闲/风格2=运动
- 童装特征库更新 + 标题模板外部化 + 图源映射增强
2026-09-03 18:28:39 +08:00

192 lines
9.3 KiB
Python
Raw Blame History

This file contains ambiguous Unicode characters
This file contains Unicode characters that might be confused with other characters. If you think that this is intentional, you can safely ignore this warning. Use the Escape button to reveal them.
"""节点 8/8:种草图生成(seed_shot)——在 oss_upload 之后。
种草图用 img2img 真正生成(不是复用主图):从 product 已生成的三合一主图
color_composites,每颜色一张)中按分配规则选参考图,以该图提取衣服颜色并作为
参考图,按 seed_shot_templates.yaml 模板 + model_features.yaml 随机模特特征生成新图:
- [商品名称] ← product 的 cn_title(上一节点多模态生成)
- [材质] ← 数据库 SPU.material 字段
- [模特特征] ← model_features.yaml 随机一条
分配规则(config.seed_shot.count):
- count <= 颜色数:随机取 count 个不同颜色,各生成 1 张
- count > 颜色数:每个颜色至少 1 张,剩余随机补足(可重复)
种草图同样压缩上传到 OSS(货号计数与 oss_upload 共用 state["oss_seq"] 续接),
URL 写入 r["seed_shot_urls"],供 template_export 插入模板详情图文列。
未配置图像后端 / 无三合一主图 / count=0 时跳过,不中断。
"""
import random
import time
from pathlib import Path
from typing import Any, Dict, List
from graph.validate import with_fallback
def _plan_seed_shots(comps: List[Dict[str, Any]], count: int) -> List[tuple]:
"""按颜色分配种草图数量:count <= 颜色数 → 随机取 count 个不同颜色各 1 张;
count > 颜色数 → 每色 1 张 + 随机补足(可重复)。返回 [(composite, n)]。"""
if count <= 0 or not comps:
return []
if len(comps) >= count:
picked = random.sample(comps, count)
return [(cc, 1) for cc in picked]
plan: List[tuple] = [(cc, 1) for cc in comps] # 每色至少 1 张
for _ in range(count - len(comps)):
cc = random.choice(comps) # 随机补足(可重复)
for i, (c, n) in enumerate(plan):
if c is cc:
plan[i] = (c, n + 1)
break
return plan
@with_fallback("seed_shot")
def seed_shot_node(state: Dict[str, Any]) -> Dict[str, Any]:
products: List[Dict[str, Any]] = state.get("product") or []
config = state["config"] or {}
country = state.get("country", "")
output_dir = Path(state["output_dir"])
ss_cfg = config.get("seed_shot") or {}
count = int(ss_cfg.get("count", 1))
if not bool(ss_cfg.get("enabled", True)) or count <= 0 or not products:
return {"seed_shots": [], "stats": state.get("stats") or {}}
# 图像后端(复用 compose 配置)
compose_cfg = config.get("compose") or {}
ib = None
if compose_cfg.get("backend"):
from graph.backends import get_image_backend
try:
ib = get_image_backend(compose_cfg["backend"])
if ib is not None:
ib.bind_config(compose_cfg)
except Exception as e: # noqa: BLE001
print(f"[seed_shot] 图像后端不可用: {e}")
if ib is None:
print("[seed_shot] 未配置 compose.backendopenai/mock),跳过种草图生成")
return {"seed_shots": [], "stats": state.get("stats") or {}}
# 材质映射:db SPU.material(清洗换行)
material_map: Dict[str, str] = {}
try:
from graph.product import list_spus
dbp = (config.get("product") or {}).get("db_path", "db/spu_sku.db")
p = Path(dbp)
if not p.is_absolute():
from graph.paths import project_root, runtime_root
for root in (runtime_root(), project_root()):
if (root / p).exists():
p = root / p
break
for s in list_spus(str(p)):
m = " ".join(str(s.get("material", "")).replace("\r", " ").replace("\n", " ").split())
material_map[s["code"]] = m
except Exception as e: # noqa: BLE001
print(f"[seed_shot] 材质读取失败(用空): {e}")
from graph.seed_shot import generate_seed_shots, read_template_category, gender_from_category
from graph.oss_upload import build_oss_key, compress_for_oss, upload_to_oss
from graph.nodes.oss_upload_node import _gen_rand4, MAX_CODE
ts = str(state.get("task_timestamp") or time.strftime("%Y%m%d%H%M%S"))
prefix = str(((config.get("product") or {}).get("code_prefix")) or "DG").strip()
seq = int(state.get("oss_seq") or 0)
oss_cfg = config.get("oss") or {}
oss_enabled = bool(oss_cfg.get("enabled", True)) and bool(oss_cfg.get("oss_bucket"))
size = str(ss_cfg.get("size") or "1536x2048")
# 类目 → 性别:模版「类目」表头值含「男」→ 男模;含「女」→ 女模;都不含 → 全部随机
gender = None
tp = str(((config.get("product") or {}).get("template_path")) or "").strip()
if tp:
category = read_template_category(tp)
gender = gender_from_category(category)
if gender:
label = {"male": "男", "female": "女", "boy_kids": "男童", "girl_kids": "女童"}.get(gender, gender)
print(f"[seed_shot] 类目「{category[:30]}…」检测到 {label} → 固定 {gender} 模特")
elif category:
print(f"[seed_shot] 类目「{category[:30]}…」无男/女 → 男女模特随机")
all_shots: List[Dict[str, Any]] = []
shot_dir = output_dir / "seed_shots"
shot_dir.mkdir(parents=True, exist_ok=True)
import concurrent.futures
import threading as _th
seq_lock = _th.Lock() # seq(货号计数)跨线程共享,需加锁
def _shot_one(r: Dict[str, Any]):
"""单个产品的种草图生成+上传(每产品独立线程)。"""
nonlocal seq
# 三合一主图:多色用 color_composites;单色回退 composite_path
comps = r.get("color_composites") or []
if not comps and r.get("composite_path") and Path(r["composite_path"]).exists():
comps = [{"sku_code": r.get("sku_code"), "color": r.get("color", ""),
"composite_path": r["composite_path"]}]
if not comps:
print(f"[seed_shot] {r.get('spu_code', '')} 无三合一主图,跳过种草图")
return None
plan = _plan_seed_shots(comps, count)
cn = (r.get("cn_title") or "").strip() or r.get("topic", "")
material = material_map.get(r.get("spu_code", ""), "")
# 按对应货号命名(img_code=货号,如 DG000);无货号时回退 seed
base_prefix = r.get("img_code") or r.get("oss_code") or ""
pfx = base_prefix or "seed"
paths: List[str] = []
for ci, (cc, n) in enumerate(plan, start=1):
base = cc.get("composite_path")
if not base or not Path(base).exists():
print(f"[seed_shot] {r.get('spu_code', '')} 参考图缺失({base}),跳过该颜色种草图")
continue
generated = generate_seed_shots(ib, base, cn, material, n, str(shot_dir),
r.get("composite_negative", ""),
size=size, prefix=pfx, gender=gender,
retries=int(ss_cfg.get("retries", 3)))
paths.extend(generated)
if not paths:
return None
r["seed_shot_paths"] = paths
urls: List[str] = []
for pth in paths:
# OSS key 用对应货号(img_code),不再自增;无货号时回退自增计数
with seq_lock:
if not base_prefix:
if seq >= MAX_CODE:
print(f"[seed_shot] 货号计数达上限 999,停止上传种草图")
break
code = f"{prefix}{seq:03d}"
seq += 1
else:
code = base_prefix
if oss_enabled:
try:
compressed = compress_for_oss(pth, str(Path(pth).with_suffix(".oss.jpg")))
url = upload_to_oss(oss_cfg, compressed,
build_oss_key(country, ts, code, _gen_rand4(), local=compressed))
if url:
urls.append(url)
r["seed_shot_urls"] = urls
except Exception as e: # noqa: BLE001
print(f"[seed_shot] 种草图上传失败 {pth}: {e}")
else:
print(f"[seed_shot] oss 未启用,仅本地保存: {pth}")
return {"spu_code": r.get("spu_code"), "sku_code": r.get("sku_code"),
"paths": paths, "urls": urls}
# 并发:每个产品一个独立线程;默认上限 5(与 product 一致,避免多产品压垮图像网关),
# config.seed_shot.concurrency 可显式覆盖(含 0/留空→默认 5)
seed_concurrency = int((config.get("seed_shot") or {}).get("concurrency") or 0) or 5
seed_concurrency = min(seed_concurrency, len(products)) if products else 1
if len(products) > 1:
print(f"[seed_shot] 并发 {seed_concurrency} 生成种草图({len(products)} 个产品)")
with concurrent.futures.ThreadPoolExecutor(max_workers=seed_concurrency) as _ex:
for item in _ex.map(_shot_one, products):
if item:
all_shots.append(item)
stats = dict(state.get("stats") or {})
stats["seed_shot"] = {"count": len(all_shots), "seq": seq}
return {"seed_shots": all_shots, "product": products, "oss_seq": seq, "stats": stats}