Files
pod_trend_agent/graph/nodes/seed_shot_node.py
T
3218485270 d68cc3b9e3 v106-v109 童装支持 + 标题模板外部化 + 模板导出增强
- 新增男童/女童检测(gender_from_category 优先判童装)与童装场景图生成
  (configs/kids_features.yaml:模特/场景/服装风格,同一商品固定同一组)
- 童装 SPU 字段映射:kids_type→SPU商品属性-类型、kids_age→适用年龄段、
  target_audience 按性别映射、kids_type_map 女童「上衣」→「针织上衣」
- 标题生成提示词外部化:prompts/title_prompt_{1,2,3}.md + config.yaml 路由表
- 模板多站点匹配:经营站点可配多个,命中任意一个即匹配
- 种草图生成失败自动重试(seed_shot.retries,换场景/模特/风格)
- 模板导出新增 Preview 文件夹(成功产品 _composite.oss.jpg + result.xlsx)
- 修复 SKU 尺码未按从小到大排序(_size_rank 支持单一年龄码 6Y/10Y)
2026-09-01 17:24:23 +08:00

190 lines
9.1 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()))
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}
# 并发:每个产品一个独立线程(默认);config.seed_shot.concurrency 可覆盖
seed_concurrency = int((config.get("seed_shot") or {}).get("concurrency") or 0) or len(products) or 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}