"""节点 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 def _color_tag(cc: Dict[str, Any], idx: int) -> str: """种草图文件名里的颜色标识:优先 sku_code 的颜色段,回退颜色名/序号。""" sku = str(cc.get("sku_code") or "") if "-" in sku: tag = sku.split("-", 1)[1] else: tag = str(cc.get("color") or "") or f"c{idx}" return "".join(ch for ch in tag if ch.isalnum() or ch in "-_") or f"c{idx}" @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.backend(openai/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 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 "1504x2000") 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", ""), "") base_prefix = r.get("img_code") or r.get("oss_code") or "" 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 tag = _color_tag(cc, ci) pfx = f"{base_prefix}_{tag}" if base_prefix else f"seed_{tag}" generated = generate_seed_shots(ib, base, cn, material, n, str(shot_dir), r.get("composite_negative", ""), size=size, prefix=pfx) paths.extend(generated) if not paths: return None r["seed_shot_paths"] = paths urls: List[str] = [] for pth in paths: with seq_lock: if seq >= MAX_CODE: print(f"[seed_shot] 货号计数达上限 999,停止上传种草图") break code = f"{prefix}{seq:03d}" seq += 1 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}