diff --git a/alembic/versions/085_atom_clip_caption_embedding.py b/alembic/versions/085_atom_clip_caption_embedding.py new file mode 100644 index 000000000..f3348d565 --- /dev/null +++ b/alembic/versions/085_atom_clip_caption_embedding.py @@ -0,0 +1,33 @@ +"""asset_atom_clips 新增 caption/embedding 字段(#2035 语义标签增强) + +Revision ID: 085_atom_clip_caption_embedding +Revises: 084_lipsync_jobs_style +Create Date: 2026-09-25 +""" + +import sqlalchemy as sa + +from alembic import op + +revision = "085_atom_clip_caption_embedding" +down_revision = "084_lipsync_jobs_style" +branch_labels = None +depends_on = None + + +def upgrade() -> None: + # caption: 中文画面描述(10-30字) + op.add_column( + "asset_atom_clips", + sa.Column("caption", sa.Text(), nullable=True), + ) + # embedding: caption 对应的向量(豆包 embedding 接口返回,JSON 存 float 数组) + op.add_column( + "asset_atom_clips", + sa.Column("embedding", sa.JSON(), nullable=True), + ) + + +def downgrade() -> None: + op.drop_column("asset_atom_clips", "embedding") + op.drop_column("asset_atom_clips", "caption") diff --git a/apps/api/app/api/routes/generation_tasks.py b/apps/api/app/api/routes/generation_tasks.py index efe8679c8..ab295f591 100755 --- a/apps/api/app/api/routes/generation_tasks.py +++ b/apps/api/app/api/routes/generation_tasks.py @@ -134,10 +134,11 @@ def _ensure_library_has_ready_video_assets(assets) -> None: def _select_assets_from_library( assets: list, mode: str, - count: int, + count: int = 0, rng=None, script_tags: list | None = None, tag_names_by_id: dict | None = None, + db=None, ) -> list[str]: """根据选取模式从素材库中选取 ready 状态的视频素材 ID。 @@ -158,6 +159,42 @@ def _select_assets_from_library( if not ready_video_assets: return [] + # #2035:加载片段级 AI 标签,供叙事模式 AI 加权和 smart 模式语义匹配使用。 + # 失败降级为空(不影响选片主流程)。 + clip_ai_tags_by_asset: dict[str, list[dict]] = {} + ai_tags_by_asset: dict[str, dict] = {} # asset_id → 聚合后的 ai_tags dict(取首个有 has_text 的片段;合并 scene/objects/action 去重) + try: + if db is not None: + from packages.adapters.sqlalchemy_impl.models import AssetAtomClipModel + ready_ids = [a.id for a in ready_video_assets] + clip_rows = ( + db.query(AssetAtomClipModel.asset_id, AssetAtomClipModel.ai_tags) + .filter(AssetAtomClipModel.asset_id.in_(ready_ids)) + .filter(AssetAtomClipModel.ai_tags.isnot(None)) + .all() + ) + agg: dict[str, dict] = {} + for asset_id, ai_tags in clip_rows: + if not isinstance(ai_tags, dict): + continue + clip_ai_tags_by_asset.setdefault(asset_id, []).append(ai_tags) + # 聚合:合并 scene/objects/action 去重 + agg.setdefault(asset_id, {"scene": [], "objects": [], "action": [], "shot": "", "has_text": False}) + for key in ("scene", "objects", "action"): + for v in ai_tags.get(key) or []: + v = str(v).strip() + if v and v not in agg[asset_id][key]: + agg[asset_id][key].append(v) + if ai_tags.get("has_text") is True: + agg[asset_id]["has_text"] = True + if not agg[asset_id]["shot"] and ai_tags.get("shot"): + agg[asset_id]["shot"] = ai_tags["shot"] + ai_tags_by_asset = agg + except Exception: # noqa: BLE001 + logger.warning("[选片] 加载片段 AI 标签失败,降级不使用语义匹配", exc_info=True) + clip_ai_tags_by_asset = {} + ai_tags_by_asset = {} + # 叙事模式(#1970 PR3):文案标签命中池优先;无任何命中时完全降级为现有随机逻辑。 if script_tags: from packages.domain.narrative_match import pick_narrative_assets @@ -167,6 +204,7 @@ def _select_assets_from_library( ready_video_assets, script_tags=script_tags, tag_names_by_id=tag_names_by_id, + clip_ai_tags_by_asset=clip_ai_tags_by_asset, limit=limit, rng=rng, ) @@ -177,7 +215,16 @@ def _select_assets_from_library( # 评分维度:质量分(40%) + 时长适配(30%) + 新鲜度(20%) + 未使用加分(10%) # 排序注入随机噪声(#1743):同分素材每次选出不同组合,从素材组合层面降重 limit = count if count > 0 else None - results = smart_select_assets(ready_video_assets, limit=limit, kind="video", rng=rng) + # #2035:给 smart_select_assets 传入文案标签和 AI 标签映射,启用语义维度 + norm_script = {t.strip().lower() for t in (script_tags or []) if t and t.strip()} + results = smart_select_assets( + ready_video_assets, + limit=limit, + kind="video", + rng=rng, + script_tags=norm_script if norm_script else None, + ai_tags_by_asset=ai_tags_by_asset or None, + ) return [r.asset.id for r in results] # 默认 all 模式:返回全部 ready 视频素材 @@ -396,6 +443,7 @@ def create_generation_task( count=request.asset_select_count, script_tags=narrative_script_tags or None, tag_names_by_id=_tag_index, + db=db, ) elif project_id and not resolved_asset_ids and (request.asset_select_mode in ("smart",) or narrative_script_tags): # 项目级模式:未指定 asset_ids 且选择了 smart 模式(或叙事模式按标签匹配)时自动选取 @@ -410,6 +458,7 @@ def create_generation_task( count=request.asset_select_count, script_tags=narrative_script_tags or None, tag_names_by_id=_tag_index, + db=db, ) if not resolved_asset_ids: raise HTTPException( diff --git a/apps/worker/worker_app/celery_app.py b/apps/worker/worker_app/celery_app.py index c3c2eb2bb..85fab34b6 100755 --- a/apps/worker/worker_app/celery_app.py +++ b/apps/worker/worker_app/celery_app.py @@ -31,6 +31,7 @@ celery_app.conf.imports = ( # #1970 片段级 AI 标签:必须显式 import 注册,否则 worker 报 # "Received unregistered task of type 'worker.tag_atom_clip'" "worker_app.tasks.atom_clip_tagging", + "worker_app.tasks.asset_quality_scoring_task", "worker_app.tasks.backfill_atom_clip_tags", "worker_app.tasks.classification", "worker_app.tasks.generation", diff --git a/apps/worker/worker_app/tasks/asset_analyzer.py b/apps/worker/worker_app/tasks/asset_analyzer.py index 949d59819..952fab835 100755 --- a/apps/worker/worker_app/tasks/asset_analyzer.py +++ b/apps/worker/worker_app/tasks/asset_analyzer.py @@ -446,21 +446,3 @@ def classify_asset_real(video_path: str) -> tuple[str, float]: logger.warning(f"Classification failed, using fallback: {e}") return AssetClassification.OTHER.value, 0.3 - -def calculate_quality_score_real(video_path: str) -> float: - """ - 质量评分入口函数 - - Args: - video_path: 视频文件路径 - - Returns: - 质量评分 (0-100) - """ - try: - analyzer = AssetAnalyzer(video_path) - result = analyzer.calculate_quality_score() - return result.total - except Exception as e: - logger.warning(f"Quality scoring failed, using fallback: {e}") - return 50.0 diff --git a/apps/worker/worker_app/tasks/asset_quality_scoring_task.py b/apps/worker/worker_app/tasks/asset_quality_scoring_task.py new file mode 100644 index 000000000..7a15b4ca3 --- /dev/null +++ b/apps/worker/worker_app/tasks/asset_quality_scoring_task.py @@ -0,0 +1,92 @@ +"""素材质量评分 Celery 任务 — #2035. + +视频素材 READY 入库后异步触发:下载视频到临时文件,运行 FFmpeg+NumPy 质量分析, +将 0-100 总分写入 assets.quality_score 字段。失败不阻断主流程(保留 NULL 或旧值, +选片时按 50 分兜底)。 + +任务名:worker.calculate_asset_quality +""" + +from __future__ import annotations + +import os +import tempfile +from pathlib import Path + +from celery.utils.log import get_task_logger +from worker_app.celery_app import celery_app +from worker_app.db import SessionLocal + +from packages.adapters.sqlalchemy_impl.asset_repository import SQLAlchemyAssetRepository +from packages.shared.storage import get_shared_storage_service + +logger = get_task_logger(__name__) + + +@celery_app.task(name="worker.calculate_asset_quality", bind=True, max_retries=1, default_retry_delay=15) +def calculate_asset_quality_task(self, asset_id: str) -> dict: + """为单个视频素材计算质量评分并写回 assets.quality_score。 + + 流程: + 1. 下载视频到临时文件; + 2. 用 AssetAnalyzer(FFmpeg+NumPy) 提取分辨率/帧率/码率/清晰度/稳定性 5 维分数; + 3. 写回 assets.quality_score。 + + 失败/非视频/无文件等情况均静默降级,返回 status=skipped/failed 不抛异常。 + """ + db = SessionLocal() + tmp_dir = tempfile.mkdtemp(prefix="quality_score_") + try: + asset_repo = SQLAlchemyAssetRepository(db) + asset = asset_repo.find_by_id(asset_id) + if asset is None: + return {"status": "skipped", "reason": "asset not found", "asset_id": asset_id} + if not (getattr(asset, "mime_type", "") or "").startswith("video/"): + return {"status": "skipped", "reason": "not a video", "asset_id": asset_id} + # 已有质量分则幂等跳过(重新计算需显式置空) + if getattr(asset, "quality_score", None) is not None: + return {"status": "skipped", "reason": "already scored", "asset_id": asset_id} + + storage = get_shared_storage_service() + storage_key = getattr(asset, "storage_key", "") or "" + if not storage_key: + return {"status": "skipped", "reason": "no storage_key", "asset_id": asset_id} + + # 下载到临时文件 + safe_suffix = ".mp4" + local_path = Path(tmp_dir) / f"asset_{asset_id[:8]}{safe_suffix}" + ok = storage.download_asset(storage_key, local_path) + if not ok or not local_path.exists() or local_path.stat().st_size == 0: + return {"status": "failed", "reason": "download failed", "asset_id": asset_id} + + # 调用 AssetAnalyzer + from worker_app.tasks.asset_analyzer import AssetAnalyzer + try: + analyzer = AssetAnalyzer(str(local_path), temp_dir=tmp_dir) + result = analyzer.calculate_quality_score() + total = float(result.total) if result and 0 <= result.total <= 100 else 50.0 + except Exception as analyze_err: # noqa: BLE001 + logger.warning("[quality_score] 分析失败,使用默认50分: asset=%s err=%s", asset_id, analyze_err) + total = 50.0 + + # 写回数据库 + asset.quality_score = total + asset_repo.update(asset) + db.commit() + + logger.info("[quality_score] asset=%s score=%.1f", asset_id, total) + return {"status": "completed", "asset_id": asset_id, "quality_score": total} + except Exception as exc: # noqa: BLE001 + db.rollback() + logger.exception("[quality_score] asset=%s 失败: %s", asset_id, exc) + if self.request.retries < self.max_retries: + raise self.retry(exc=exc) from None + return {"status": "failed", "asset_id": asset_id, "error": str(exc)} + finally: + db.close() + # 清理临时文件 + try: + import shutil + shutil.rmtree(tmp_dir, ignore_errors=True) + except Exception: + pass diff --git a/apps/worker/worker_app/tasks/atom_clip_tagging.py b/apps/worker/worker_app/tasks/atom_clip_tagging.py index f3dc24ca0..2a6b89d38 100644 --- a/apps/worker/worker_app/tasks/atom_clip_tagging.py +++ b/apps/worker/worker_app/tasks/atom_clip_tagging.py @@ -78,14 +78,18 @@ def tag_atom_clip_task(self, atom_clip_id: str, force: bool = False) -> dict: atom_repo.update_ai_tags(atom_clip_id, ai_tags) logger.info( - "[atom_clip_tagging] clip_id=%s ai_tags=%s", + "[atom_clip_tagging] clip_id=%s ai_tags=%s caption=%r has_embedding=%s", atom_clip_id, {k: v for k, v in ai_tags.items() if k != "inherited_tags"}, + caption, + bool(embedding), ) return { "status": "completed", "clip_id": atom_clip_id, "has_ai_tags": any(v for k, v in ai_tags.items() if k != "inherited_tags" and v), + "caption": caption, + "embedding_dim": len(embedding) if embedding else 0, } except Exception as exc: db.rollback() diff --git a/apps/worker/worker_app/tasks/ingest.py b/apps/worker/worker_app/tasks/ingest.py index 9410d74f8..f96e8f5b6 100755 --- a/apps/worker/worker_app/tasks/ingest.py +++ b/apps/worker/worker_app/tasks/ingest.py @@ -808,17 +808,23 @@ def ingest_asset(job_id: str) -> dict: db.commit() - # ── #1970 素材原子切片:视频 READY 后异步触发,失败不阻断入库 ── + # ── #1970 素材原子切片 + #2035 质量评分:视频 READY 后异步触发,失败不阻断入库 ── # atom_clips 未就绪时选片逻辑有内存兜底(compute_fallback_clips)。 + # quality_score 未计算时选片按 50 分兜底。 try: if media_type == "video" and float(asset.duration or 0) > 0: celery_app.send_task( "worker.generate_atom_clips", args=[asset.id], ) + # #2035: 异步质量评分(不与 atom_clips 链式耦合,独立任务) + celery_app.send_task( + "worker.calculate_asset_quality", + args=[asset.id], + ) except Exception as atom_err: # noqa: BLE001 logger.warning( - "触发原子切片任务失败(不影响入库): asset_id=%s err=%s", + "触发原子切片/质量评分任务失败(不影响入库): asset_id=%s err=%s", asset.id, atom_err, ) diff --git a/packages/adapters/sqlalchemy_impl/asset_atom_clip_repository.py b/packages/adapters/sqlalchemy_impl/asset_atom_clip_repository.py index c1f09abdd..c62bb0bde 100644 --- a/packages/adapters/sqlalchemy_impl/asset_atom_clip_repository.py +++ b/packages/adapters/sqlalchemy_impl/asset_atom_clip_repository.py @@ -83,6 +83,23 @@ class SQLAlchemyAssetAtomClipRepository: models = query.all() return [self._to_domain(m) for m in models] + def update_caption_embedding(self, clip_id: str, caption: str | None, embedding: list[float] | None = None) -> bool: + """更新片段的 caption 和 embedding 字段。""" + upd: dict = {} + if caption is not None: + upd["caption"] = caption + if embedding is not None: + upd["embedding"] = embedding + if not upd: + return False + count = ( + self.session.query(AssetAtomClipModel) + .filter(AssetAtomClipModel.id == clip_id) + .update(upd) + ) + self.session.commit() + return count > 0 + def update_ai_tags(self, clip_id: str, ai_tags: dict) -> bool: """更新指定片段的 ai_tags 字段.""" count = ( @@ -118,6 +135,8 @@ class SQLAlchemyAssetAtomClipRepository: clip_index=clip.clip_index, tags=clip.tags, ai_tags=clip.ai_tags, + caption=clip.caption, + embedding=clip.embedding, scene_change_at=clip.scene_change_at, is_fallback=clip.is_fallback, created_at=clip.created_at or datetime.now(UTC), @@ -132,6 +151,9 @@ class SQLAlchemyAssetAtomClipRepository: duration=model.duration, clip_index=model.clip_index, tags=model.tags or [], + ai_tags=getattr(model, "ai_tags", None), + caption=getattr(model, "caption", None), + embedding=getattr(model, "embedding", None), scene_change_at=model.scene_change_at, is_fallback=model.is_fallback, created_at=model.created_at, diff --git a/packages/adapters/sqlalchemy_impl/models.py b/packages/adapters/sqlalchemy_impl/models.py index 131e311c8..bf6da4a57 100755 --- a/packages/adapters/sqlalchemy_impl/models.py +++ b/packages/adapters/sqlalchemy_impl/models.py @@ -2,6 +2,7 @@ from datetime import UTC, datetime from typing import Any from sqlalchemy import ( + JSON, Boolean, Column, @@ -841,6 +842,8 @@ class AssetAtomClipModel(Base): clip_index = Column(Integer, nullable=False) tags = Column(JSON, nullable=False, default=list) ai_tags = Column(JSON, nullable=True, default=None) + caption = Column(Text, nullable=True, default=None) + embedding = Column(JSON, nullable=True, default=None) scene_change_at = Column(Float, nullable=True) is_fallback = Column(Boolean, nullable=False, default=False) created_at = Column(DateTime, nullable=False, default=lambda: datetime.now(UTC)) diff --git a/packages/config/base.py b/packages/config/base.py index 63091a365..4f78b0516 100755 --- a/packages/config/base.py +++ b/packages/config/base.py @@ -70,6 +70,7 @@ class SharedSettings(BaseSettings): doubao_timeout: int = 30 doubao_max_retries: int = 2 doubao_vision_model: str = "doubao-1-5-vision-pro-250915" + doubao_embedding_model: str = "doubao-embedding-large-text-240915" # ── MediaKit (火山引擎 AI 媒体工具) ────────────────────────────────── mediakit_api_key: str = "" diff --git a/packages/domain/asset_atom_clip.py b/packages/domain/asset_atom_clip.py index 86cd164ec..b04032b1b 100644 --- a/packages/domain/asset_atom_clip.py +++ b/packages/domain/asset_atom_clip.py @@ -37,6 +37,8 @@ class AssetAtomClip: clip_index: int tags: list[str] = field(default_factory=list) ai_tags: dict | None = None + caption: str | None = None + embedding: list[float] | None = None scene_change_at: float | None = None is_fallback: bool = False created_at: datetime | None = None @@ -67,6 +69,8 @@ class AssetAtomClip: tags: list[str] | None = None, scene_change_at: float | None = None, is_fallback: bool = False, + caption: str | None = None, + embedding: list[float] | None = None, ) -> AssetAtomClip: """工厂方法:创建一个新的原子片段。""" return cls( @@ -79,4 +83,6 @@ class AssetAtomClip: tags=tags or [], scene_change_at=scene_change_at, is_fallback=is_fallback, + caption=caption, + embedding=embedding, ) diff --git a/packages/domain/atom_clip_tagger.py b/packages/domain/atom_clip_tagger.py index f40dbd1a5..f64e2ec69 100644 --- a/packages/domain/atom_clip_tagger.py +++ b/packages/domain/atom_clip_tagger.py @@ -35,6 +35,7 @@ def build_vision_prompt() -> str: - action: 动作类型列表(如 "演示", "说话", "操作") - shot: 景别("特写" / "中景" / "远景" 之一) - has_text: 画面中是否有显著文字(true/false) + - caption: 一句中文画面描述(10-30字),简洁概括这段视频的内容主体与场景 """ return """请分析这段视频片段的关键帧,识别内容并返回 JSON 格式标签。 @@ -44,7 +45,8 @@ def build_vision_prompt() -> str: "objects": ["物体1", "物体2"], "action": ["动作1"], "shot": "特写|中景|远景", - "has_text": true/false + "has_text": true/false, + "caption": "一句中文描述" } 规则: @@ -53,6 +55,7 @@ def build_vision_prompt() -> str: - action: 人物或物体正在进行的动作,如"演示"、"说话"、"操作"、"展示"等,1-3个 - shot: 景别判断,只能是"特写"、"中景"或"远景"之一 - has_text: 画面中是否有显著可读文字(标题、字幕、标语等) +- caption: 一句简洁的中文画面描述(10-30字),概括这段视频的主体内容、人物动作和场景,例如"一名女性在办公室中讲解产品展示" 请只返回 JSON,不要有其他说明文字。""" @@ -65,7 +68,7 @@ def parse_vision_response(text: str) -> dict: Returns: 结构化标签 dict,格式如: - {"scene": [...], "objects": [...], "action": [...], "shot": "...", "has_text": bool} + {"scene": [...], "objects": [...], "action": [...], "shot": "...", "has_text": bool, "caption": "..."} 解析失败时返回空 dict。 """ @@ -127,6 +130,15 @@ def parse_vision_response(text: str) -> dict: else: result["has_text"] = False + cap_val = data.get("caption", "") + if isinstance(cap_val, str): + cap_val = cap_val.strip() + if len(cap_val) > 60: + cap_val = cap_val[:60] + else: + cap_val = "" + result["caption"] = cap_val + return result @@ -237,14 +249,14 @@ def tag_atom_clip( Returns: 结构化标签 dict,格式如: {"scene": [...], "objects": [...], "action": [...], "shot": "...", - "has_text": bool, "inherited_tags": [...]} + "has_text": bool, "caption": "...", "inherited_tags": [...]} """ inherited = list(getattr(clip, "tags", []) or []) # 检查 DoubaoClient 是否可用 if not getattr(doubao_client, "is_available", False): logger.info("DoubaoClient 不可用,跳过 AI 标签: clip_id=%s", getattr(clip, "id", "")) - return {"inherited_tags": inherited} + return {"scene": [], "objects": [], "action": [], "shot": "", "has_text": False, "caption": "", "inherited_tags": inherited} # 提取帧图片 frame_urls: Optional[list[str]] = None @@ -261,7 +273,7 @@ def tag_atom_clip( if not frame_urls: logger.warning("帧提取失败,跳过 AI 标签: clip_id=%s", getattr(clip, "id", "")) - return {"inherited_tags": inherited} + return {"scene": [], "objects": [], "action": [], "shot": "", "has_text": False, "caption": "", "inherited_tags": inherited} # 调用视觉 API prompt = build_vision_prompt() @@ -275,17 +287,17 @@ def tag_atom_clip( ) except Exception as e: logger.warning("视觉 API 调用异常: clip_id=%s error=%s", getattr(clip, "id", ""), e) - return {"inherited_tags": inherited} + return {"scene": [], "objects": [], "action": [], "shot": "", "has_text": False, "caption": "", "inherited_tags": inherited} if not response_text: logger.warning("视觉 API 返回空: clip_id=%s", getattr(clip, "id", "")) - return {"inherited_tags": inherited} + return {"scene": [], "objects": [], "action": [], "shot": "", "has_text": False, "caption": "", "inherited_tags": inherited} # 解析标签 ai_tags = parse_vision_response(response_text) if not ai_tags: logger.warning("标签解析失败: clip_id=%s response=%s", getattr(clip, "id", ""), response_text[:200]) - return {"inherited_tags": inherited} + return {"scene": [], "objects": [], "action": [], "shot": "", "has_text": False, "caption": "", "inherited_tags": inherited} # 合并 inherited_tags ai_tags["inherited_tags"] = inherited diff --git a/packages/domain/narrative_match.py b/packages/domain/narrative_match.py index 9a0a5b978..2cd78c87c 100644 --- a/packages/domain/narrative_match.py +++ b/packages/domain/narrative_match.py @@ -245,16 +245,39 @@ def pick_narrative_assets( clip_ai_tags_by_asset=clip_ai_tags_by_asset, ) + # #2035:把文案标签与聚合的素材级 ai_tags 透传给 smart_select_assets, + # 让 smart 评分维度(ai_semantic)在叙事模式内部兜底/补位时同样生效。 + wanted_norm = _normalize_tags(script_tags) + asset_ai_tags: dict[str, dict] = {} + if clip_ai_tags_by_asset: + for aid, clips in clip_ai_tags_by_asset.items(): + agg: dict = {"scene": [], "objects": [], "action": []} + for clip_tags in clips or []: + if not isinstance(clip_tags, dict): + continue + for key in ("scene", "objects", "action"): + for v in clip_tags.get(key) or []: + v = str(v).strip() + if v and v not in agg[key]: + agg[key].append(v) + asset_ai_tags[aid] = agg + + smart_kwargs = dict( + kind="video", + rng=rng, + script_tags=wanted_norm if wanted_norm else None, + ai_tags_by_asset=asset_ai_tags if asset_ai_tags else None, + ) + need = limit if (limit is not None and limit > 0) else None if not matched: - # 完全降级:与改造前随机混剪同一逻辑 - return [r.asset for r in smart_select_assets(assets, kind="video", limit=need, rng=rng)] + return [r.asset for r in smart_select_assets(assets, limit=need, **smart_kwargs)] - picked = [r.asset for r in smart_select_assets(matched, kind="video", limit=need, rng=rng)] + picked = [r.asset for r in smart_select_assets(matched, limit=need, **smart_kwargs)] if need is not None and len(picked) < need and unmatched: rest_need = need - len(picked) - picked.extend(r.asset for r in smart_select_assets(unmatched, kind="video", limit=rest_need, rng=rng)) + picked.extend(r.asset for r in smart_select_assets(unmatched, limit=rest_need, **smart_kwargs)) elif need is None: - picked.extend(r.asset for r in smart_select_assets(unmatched, kind="video", rng=rng)) + picked.extend(r.asset for r in smart_select_assets(unmatched, **smart_kwargs)) return picked diff --git a/packages/domain/smart_match.py b/packages/domain/smart_match.py index 85ba709d3..d4ca3324d 100755 --- a/packages/domain/smart_match.py +++ b/packages/domain/smart_match.py @@ -40,6 +40,14 @@ def _get_enum_value(obj: Any, attr: str) -> str: return val.value if hasattr(val, "value") else str(val) +def normalize_tag(tag) -> str: + """标准化标签:去两端空白、小写;非 str 转 str。返回空串表示应丢弃。""" + if tag is None: + return "" + s = str(tag).strip().lower() + return s + + def _duration_bucket(duration: float | None) -> str: """将素材时长分为 3 档:short(<10s) / medium(10-30s) / long(>30s)。""" if duration is None or duration <= 0: @@ -54,14 +62,21 @@ def _duration_bucket(duration: float | None) -> str: def score_asset( asset: Any, now: datetime | None = None, + script_tags: set | None = None, + ai_tags_by_asset: dict | None = None, ) -> tuple[float, dict[str, float]]: """为单个素材计算综合得分(0-100)。 - 维度权重: - - quality_score (40%):素材质量分(0-100),无质量分按 50 计 - - duration_fitness (30%):时长适配度,5-30s 为最优区间 - - recency (20%):新鲜度,30 天内衰减 - - unused_bonus (10%):未被使用过的素材加分 + 维度权重(#2035 加入 AI 语义匹配维度): + - quality_score (30%):素材质量分(0-100),无质量分按 50 计 + - duration_fitness (25%):时长适配度,5-30s 为最优区间 + - recency (15%):新鲜度,30 天内衰减 + - unused (10%):未被/少被使用过的素材加分 + - ai_semantic (20%):AI 标签(scene/objects/action)与文案标签重合度;无数据给 50 中性分 + + Args: + script_tags: 标准化后的文案标签集合,用于 AI 语义匹配维度打分。 + ai_tags_by_asset: asset_id → ai_tags dict 映射,ai_tags 含 scene/objects/action 字段。 Returns: (total_score, breakdown_dict) @@ -73,7 +88,7 @@ def score_asset( # 1. 质量分 (0-100) → 权重 40% raw_quality = asset.quality_score if asset.quality_score is not None else 50.0 - quality_component = raw_quality * 0.4 + quality_component = raw_quality * 0.30 breakdown["quality"] = round(quality_component, 2) # 2. 时长适配度 (0-100) → 权重 30% @@ -90,7 +105,7 @@ def score_asset( # >30s: 指数衰减,60s 时约 50 分 duration_fitness = 100.0 * math.exp(-0.02 * (duration - 30)) duration_fitness = max(duration_fitness, 10.0) - duration_component = duration_fitness * 0.3 + duration_component = duration_fitness * 0.25 breakdown["duration"] = round(duration_component, 2) # 3. 新鲜度 (0-100) → 权重 20% @@ -103,7 +118,7 @@ def score_asset( created_at = created_at.replace(tzinfo=UTC) age_days = max(0, (now - created_at).total_seconds() / 86400) recency = 100.0 * math.exp(-0.05 * age_days) # ~14天半衰期 - recency_component = recency * 0.2 + recency_component = recency * 0.15 breakdown["recency"] = round(recency_component, 2) # 4. 未使用偏好 (0-100) → 权重 10% @@ -121,7 +136,34 @@ def score_asset( unused_component = unused_score * 0.1 breakdown["unused"] = round(unused_component, 2) - total = quality_component + duration_component + recency_component + unused_component + # 5. AI 语义匹配 (0-100) → 权重 20% + if script_tags and ai_tags_by_asset: + asset_ai = ai_tags_by_asset.get(getattr(asset, "id", "")) or {} + ai_terms: set = set() + for key in ("scene", "objects", "action"): + vals = asset_ai.get(key) or [] + if isinstance(vals, list): + for v in vals: + norm = normalize_tag(v) + if norm: + ai_terms.add(norm) + if ai_terms: + norm_script = {normalize_tag(t) for t in script_tags if normalize_tag(t)} + overlap = ai_terms & norm_script + union = ai_terms | norm_script + ratio = (len(overlap) / len(union)) if union else 0.0 + if overlap: + ai_score = 50.0 + 50.0 * ratio + else: + ai_score = 20.0 + else: + ai_score = 50.0 + else: + ai_score = 50.0 + ai_component = ai_score * 0.20 + breakdown["ai_semantic"] = round(ai_component, 2) + + total = quality_component + duration_component + recency_component + unused_component + ai_component return round(total, 2), breakdown @@ -132,6 +174,8 @@ def smart_select_assets( kind: str | None = None, now: datetime | None = None, rng: random.Random | None = None, + script_tags: set | None = None, + ai_tags_by_asset: dict | None = None, ) -> list[SmartMatchResult]: """从素材列表中智能选取素材。 @@ -159,7 +203,7 @@ def smart_select_assets( # Step 3: 评分 scored: list[SmartMatchResult] = [] for a in ready_assets: - total, breakdown = score_asset(a, now=now) + total, breakdown = score_asset(a, now=now, script_tags=script_tags, ai_tags_by_asset=ai_tags_by_asset) scored.append(SmartMatchResult(asset=a, score=total, breakdown=breakdown)) # Step 4: 按「得分 + 随机噪声」降序排序 diff --git a/packages/shared/ai_client.py b/packages/shared/ai_client.py index 5467013f3..4304948bb 100755 --- a/packages/shared/ai_client.py +++ b/packages/shared/ai_client.py @@ -40,6 +40,46 @@ class DoubaoClient: self.vision_model: str = settings.doubao_vision_model @property + + def embed_text(self, text: str, timeout: int | None = None) -> list[float] | None: + """调用豆包文本 Embedding API,返回浮点向量;失败返回 None。""" + if not self.is_available or not text or not text.strip(): + return None + + url = f"{self.base_url}/embeddings" + headers = { + "Authorization": f"Bearer {self.api_key}", + "Content-Type": "application/json", + } + payload: dict[str, Any] = { + "model": getattr(self, "embedding_model", None) or "doubao-embedding-large-text-240915", + "input": text.strip(), + "encoding_format": "float", + } + + req_timeout = timeout or self.timeout + last_error: Exception | None = None + for attempt in range(self.max_retries + 1): + try: + resp = httpx.post(url, headers=headers, json=payload, timeout=req_timeout) + resp.raise_for_status() + data = resp.json() + emb_list = data.get("data") or [] + if emb_list and isinstance(emb_list, list): + vec = emb_list[0].get("embedding") + if isinstance(vec, list) and vec: + return [float(x) for x in vec] + logger.warning("embedding 返回结构异常: %s", str(data)[:200]) + return None + except Exception as e: + last_error = e + if attempt < self.max_retries: + wait = 0.5 * (2**attempt) + logger.warning("豆包 Embedding 调用失败,%.1fs 后重试 (%d/%d): %s", wait, attempt+1, self.max_retries+1, e) + time.sleep(wait) + logger.error("豆包 Embedding 调用最终失败: %s", last_error) + return None + def is_available(self) -> bool: """是否可用(配置了 API Key).""" return bool(self.api_key) diff --git a/tests/unit/test_1970_atom_clip_tagger.py b/tests/unit/test_1970_atom_clip_tagger.py index 5b95107b9..b3808119c 100644 --- a/tests/unit/test_1970_atom_clip_tagger.py +++ b/tests/unit/test_1970_atom_clip_tagger.py @@ -210,7 +210,12 @@ class TestTagAtomClip: doubao_client=fake_doubao, ) - assert result == {"inherited_tags": ["tag1", "tag2"]} + assert result["inherited_tags"] == ["tag1", "tag2"] + assert result.get("caption", "") == "" + assert result.get("scene", []) == [] + assert result.get("objects", []) == [] + assert result.get("action", []) == [] + assert "inherited_tags" in result assert len(fake_doubao.vision_calls) == 0 def test_mediakit_unavailable_no_ffmpeg(self): @@ -227,7 +232,12 @@ class TestTagAtomClip: ) # 没有 ffmpeg 的情况下,帧提取失败 - assert result == {"inherited_tags": ["tag1", "tag2"]} + assert result["inherited_tags"] == ["tag1", "tag2"] + assert result.get("caption", "") == "" + assert result.get("scene", []) == [] + assert result.get("objects", []) == [] + assert result.get("action", []) == [] + assert "inherited_tags" in result def test_vision_api_error_returns_inherited(self): """视觉 API 抛异常 → 降级 inherited_tags.""" @@ -242,7 +252,12 @@ class TestTagAtomClip: mediakit_client=fake_mediakit, ) - assert result == {"inherited_tags": ["tag1", "tag2"]} + assert result["inherited_tags"] == ["tag1", "tag2"] + assert result.get("caption", "") == "" + assert result.get("scene", []) == [] + assert result.get("objects", []) == [] + assert result.get("action", []) == [] + assert "inherited_tags" in result def test_vision_api_empty_response(self): """视觉 API 返回空 → 降级 inherited_tags.""" @@ -257,7 +272,12 @@ class TestTagAtomClip: mediakit_client=fake_mediakit, ) - assert result == {"inherited_tags": ["tag1", "tag2"]} + assert result["inherited_tags"] == ["tag1", "tag2"] + assert result.get("caption", "") == "" + assert result.get("scene", []) == [] + assert result.get("objects", []) == [] + assert result.get("action", []) == [] + assert "inherited_tags" in result def test_vision_api_invalid_json_response(self): """视觉 API 返回无效 JSON → 降级 inherited_tags.""" @@ -272,7 +292,12 @@ class TestTagAtomClip: mediakit_client=fake_mediakit, ) - assert result == {"inherited_tags": ["tag1", "tag2"]} + assert result["inherited_tags"] == ["tag1", "tag2"] + assert result.get("caption", "") == "" + assert result.get("scene", []) == [] + assert result.get("objects", []) == [] + assert result.get("action", []) == [] + assert "inherited_tags" in result def test_clip_with_empty_tags(self): """空素材标签 → inherited_tags 为空列表.""" @@ -285,7 +310,8 @@ class TestTagAtomClip: doubao_client=fake_doubao, ) - assert result == {"inherited_tags": []} + assert result["inherited_tags"] == [] + assert result.get("caption", "") == "" if __name__ == "__main__": diff --git a/tests/unit/test_2035_semantic_tags.py b/tests/unit/test_2035_semantic_tags.py new file mode 100644 index 000000000..4909a588b --- /dev/null +++ b/tests/unit/test_2035_semantic_tags.py @@ -0,0 +1,257 @@ +"""#2035 语义标签增强 / 质量评分 / AI选片 单测。""" +from __future__ import annotations + +import random +from dataclasses import dataclass, field +from datetime import UTC, datetime, timedelta + +import pytest + +from packages.domain.asset_atom_clip import AssetAtomClip +from packages.domain.atom_clip_tagger import parse_vision_response +from packages.domain.narrative_match import ( + _compute_ai_score, + _extract_ai_tag_names, + match_assets_by_script_tags, + pick_narrative_assets, +) +from packages.domain.smart_match import score_asset, smart_select_assets + + +# ── helpers ────────────────────────────────────────────────────── + + +@dataclass +class FakeAsset: + id: str + tag_ids: list[str] = field(default_factory=list) + tags: list[str] = field(default_factory=list) + status: object = None + file_type: str = "video" + duration: float = 10.0 + quality_score: float | None = 50.0 + created_at: datetime | None = None + metadata: dict = field(default_factory=dict) + usage_count: int = 0 + + def __post_init__(self): + if self.status is None: + class _S: + value = "ready" + self.status = _S() + if self.created_at is None: + self.created_at = datetime.now(UTC) - timedelta(days=1) + + +# ── parse_vision_response: caption 提取 ───────────────────────── + + +class TestParseVisionResponseCaption: + def test_extracts_caption(self): + text = '{"scene":["办公室"],"objects":["电脑","人"],"action":["说话"],"shot":"中景","has_text":false,"caption":"职场女性在办公室讲解产品功能"}' + result = parse_vision_response(text) + assert result["caption"] == "职场女性在办公室讲解产品功能" + assert result["has_text"] is False + assert result["scene"] == ["办公室"] + + def test_caption_truncated_at_60(self): + long = "A" * 100 + text = '{"scene":[],"objects":[],"action":[],"shot":"中景","has_text":false,"caption":"' + long + '"}' + result = parse_vision_response(text) + assert len(result["caption"]) == 60 + + def test_missing_caption_defaults_empty(self): + text = '{"scene":[],"objects":[],"action":[],"shot":"特写","has_text":true}' + result = parse_vision_response(text) + assert result["caption"] == "" + + def test_empty_input_returns_empty_dict(self): + assert parse_vision_response("") == {} + assert parse_vision_response(None) == {} + + +# ── AI 标签提取 ────────────────────────────────────────────────── + + +class TestExtractAiTagNames: + def test_extracts_scene_objects_action(self): + tags = {"scene": ["办公室"], "objects": ["电脑", "杯子"], "action": ["说话"], "shot": "中景", "has_text": False} + names = _extract_ai_tag_names(tags) + assert "办公室" in names + assert "电脑" in names + assert "杯子" in names + assert "说话" in names + assert "中景" not in names # shot 不参与匹配 + + def test_empty_tags(self): + assert _extract_ai_tag_names({}) == set() + assert _extract_ai_tag_names({"scene": []}) == set() + + +# ── _compute_ai_score ─────────────────────────────────────────── + + +class TestComputeAiScore: + def test_basic_hit(self): + clip_map = {"a1": [{"scene": ["工厂"], "objects": ["产品"], "action": ["演示"]}]} + score = _compute_ai_score("a1", {"工厂", "演示"}, clip_map) + # 2 hits * weight 2.0 = 4.0 + assert score == 4.0 + + def test_no_hit(self): + clip_map = {"a1": [{"scene": ["户外"], "objects": [], "action": []}]} + assert _compute_ai_score("a1", {"办公室"}, clip_map) == 0.0 + + def test_no_clip_map(self): + assert _compute_ai_score("a1", {"工厂"}, None) == 0.0 + assert _compute_ai_score("a1", set(), {"a1": [{"scene": ["x"]}]}) == 0.0 + + def test_best_clip_score_not_sum(self): + """多片段取最高得分,不是累加。""" + clip_map = { + "a1": [ + {"scene": ["工厂"], "objects": [], "action": []}, # 1 hit + {"scene": ["工厂"], "objects": ["产品"], "action": ["演示"]}, # 3 hits + {"scene": ["户外"], "objects": [], "action": []}, # 0 + ] + } + score = _compute_ai_score("a1", {"工厂", "产品", "演示"}, clip_map) + assert score == 3 * 2.0 # best = 6.0, not (1+3+0)*2 = 8.0 + + +# ── match_assets_by_script_tags 接受 clip_ai_tags_by_asset ────── + + +class TestMatchSplitAiTags: + def test_ai_hit_only_puts_in_matched(self): + """素材无人工标签,但 AI 标签命中 → 命中池。""" + assets = [FakeAsset("a1", tags=[]), FakeAsset("a2", tags=["旅游"])] + clip_map = {"a1": [{"scene": ["工厂"], "objects": [], "action": []}]} + matched, unmatched = match_assets_by_script_tags( + assets, script_tags=["工厂"], clip_ai_tags_by_asset=clip_map + ) + assert [a.id for a in matched] == ["a1"] + assert [a.id for a in unmatched] == ["a2"] + + def test_ai_and_manual_both_hit(self): + assets = [FakeAsset("a1", tags=["工厂"]), FakeAsset("a2", tags=[])] + clip_map = {"a1": [{"objects": ["产品"]}]} + matched, unmatched = match_assets_by_script_tags( + assets, script_tags=["工厂", "产品"], clip_ai_tags_by_asset=clip_map + ) + assert [a.id for a in matched] == ["a1"] + # a2 无人标签也无AI命中 → unmatched + assert [a.id for a in unmatched] == ["a2"] + + +# ── score_asset AI 语义维度 ──────────────────────────────────── + + +class TestScoreAssetAiSemantic: + def test_no_ai_data_gives_neutral_ai_component(self): + a = FakeAsset("a1", quality_score=80) + total, breakdown = score_asset(a, now=datetime.now(UTC)) + # ai_semantic 中性分 50 * 0.20 = 10 + assert breakdown["ai_semantic"] == 10.0 + + def test_ai_hit_boosts_score(self): + a = FakeAsset("a1", quality_score=50) + ai_map = {"a1": {"scene": ["工厂"], "objects": ["产品"], "action": ["演示"]}} + total_hit, _ = score_asset( + a, now=datetime.now(UTC), script_tags={"工厂", "产品"}, ai_tags_by_asset=ai_map + ) + total_miss, _ = score_asset( + a, now=datetime.now(UTC), script_tags={"旅游"}, ai_tags_by_asset=ai_map + ) + total_neutral, _ = score_asset(a, now=datetime.now(UTC)) + assert total_hit > total_neutral + assert total_neutral > total_miss + + def test_weights_sum_to_100(self): + a = FakeAsset("a1", quality_score=100, duration=15) + a.created_at = datetime.now(UTC) + a.metadata = {"generation_use_count": 0} + _, bd = score_asset(a, now=datetime.now(UTC)) + # 满分素材:quality=30, duration=25, recency=~15 (new), unused=10, ai=10(neutral) + # 总和应该 ~90 + assert 85 <= sum(bd.values()) <= 100.5 + + +# ── smart_select_assets 接受 script_tags/ai_tags_by_asset ────── + + +class TestSmartSelectAi: + def test_ai_hit_ranks_higher(self): + a1 = FakeAsset("a1", quality_score=50, duration=15) + a2 = FakeAsset("a2", quality_score=50, duration=15) + a3 = FakeAsset("a3", quality_score=50, duration=15) + ai_map = { + "a1": {"scene": ["工厂"], "objects": ["产品"], "action": ["演示"]}, + "a2": {"scene": ["户外"], "objects": [], "action": []}, + "a3": {}, + } + rng = random.Random(42) + results = smart_select_assets( + [a1, a2, a3], + kind="video", + rng=rng, + script_tags={"工厂", "产品", "演示"}, + ai_tags_by_asset=ai_map, + ) + assert results[0].asset.id == "a1" # AI 命中应排第一 + + def test_without_ai_params_works_as_before(self): + a1 = FakeAsset("a1", quality_score=80, duration=15) + a2 = FakeAsset("a2", quality_score=40, duration=15) + rng = random.Random(0) + results = smart_select_assets([a1, a2], kind="video", rng=rng) + assert results[0].asset.id == "a1" + + +# ── pick_narrative_assets 接受 clip_ai_tags_by_asset ──────────── + + +class TestPickNarrativeAi: + def test_ai_tagged_assets_selected_first(self): + a1 = FakeAsset("a1", tags=[]) + a2 = FakeAsset("a2", tags=[]) + a3 = FakeAsset("a3", tags=["无关"]) + clip_map = { + "a1": [{"scene": ["工厂"], "objects": ["产品"], "action": ["演示"]}], + "a2": [{"scene": ["户外"], "objects": [], "action": []}], + } + rng = random.Random(0) + picked = pick_narrative_assets( + [a1, a2, a3], + script_tags=["工厂", "产品"], + tag_names_by_id={}, + clip_ai_tags_by_asset=clip_map, + rng=rng, + limit=2, + ) + assert picked[0].id == "a1" # a1 命中 AI 标签应在首位 + assert {a.id for a in picked} == {"a1", "a3"} or {a.id for a in picked} == {"a1", "a2"} + + +# ── AssetAtomClip 字段扩展 ───────────────────────────────────── + + +class TestAssetAtomClipNewFields: + def test_caption_embedding_fields(self): + clip = AssetAtomClip.create( + asset_id="a1", + start_time=0, + end_time=5, + clip_index=0, + tags=[], + caption="测试画面描述", + embedding=[0.1, 0.2, 0.3], + ) + assert clip.caption == "测试画面描述" + assert clip.embedding == [0.1, 0.2, 0.3] + assert clip.ai_tags is None + + def test_default_fields_none(self): + clip = AssetAtomClip.create("a1", 0, 5, 0) + assert clip.caption is None + assert clip.embedding is None diff --git a/tests/unit/test_smart_match.py b/tests/unit/test_smart_match.py index 6c28f2e92..d0da1a4ac 100755 --- a/tests/unit/test_smart_match.py +++ b/tests/unit/test_smart_match.py @@ -93,31 +93,31 @@ class TestScoreAsset: def test_no_quality_score_defaults_to_50(self): asset = FakeAsset(id="a1", quality_score=None, duration=15) score, breakdown = score_asset(asset, now=NOW) - # quality component should be 50 * 0.4 = 20 - assert breakdown["quality"] == pytest.approx(20.0, abs=0.1) + # quality component should be 50 * 0.30 = 15 + assert breakdown["quality"] == pytest.approx(15.0, abs=0.1) def test_optimal_duration_5_to_30_gets_full_score(self): for dur in [5, 10, 20, 30]: asset = FakeAsset(id="a1", quality_score=50, duration=dur) _, breakdown = score_asset(asset, now=NOW) - # duration component should be 100 * 0.3 = 30 - assert breakdown["duration"] == pytest.approx(30.0, abs=0.1) + # duration component should be 100 * 0.25 = 25 + assert breakdown["duration"] == pytest.approx(25.0, abs=0.1) def test_short_duration_below_5s_penalized(self): asset = FakeAsset(id="a1", quality_score=50, duration=2) _, breakdown = score_asset(asset, now=NOW) - assert breakdown["duration"] < 30.0 + assert breakdown["duration"] < 25.0 # below max duration score def test_long_duration_above_30s_penalized(self): asset = FakeAsset(id="a1", quality_score=50, duration=120) _, breakdown = score_asset(asset, now=NOW) - assert breakdown["duration"] < 30.0 + assert breakdown["duration"] < 25.0 # below max duration score def test_zero_duration_gives_moderate_score(self): asset = FakeAsset(id="a1", quality_score=50, duration=0) _, breakdown = score_asset(asset, now=NOW) # duration_fitness = 30.0, component = 30 * 0.3 = 9 - assert breakdown["duration"] == pytest.approx(9.0, abs=0.1) + assert breakdown["duration"] == pytest.approx(7.5, abs=0.1) def test_unused_asset_gets_full_bonus(self): asset = FakeAsset(id="a1", quality_score=50, duration=15, metadata={}) @@ -133,17 +133,17 @@ class TestScoreAsset: """int() conversion of non-numeric metadata should not raise, should default to 0.""" asset = FakeAsset(id="a1", quality_score=50, duration=15, metadata={"generation_use_count": "high"}) _, breakdown = score_asset(asset, now=NOW) - assert breakdown["unused"] == pytest.approx(10.0, abs=0.1) # use_count=0 → unused_score=100 → 100*0.1=10 + assert breakdown["unused"] == pytest.approx(10.0, abs=0.1) # use_count=0 → unused_score=100 → 100*0.10=10 def test_recent_asset_scores_higher_recency(self): asset = FakeAsset(id="a1", quality_score=50, duration=15, created_at=NOW - timedelta(days=1)) _, breakdown = score_asset(asset, now=NOW) - assert breakdown["recency"] > 15 # > 75% of max 20 + assert breakdown["recency"] > 11 # > 75% of max 15 def test_old_asset_scores_lower_recency(self): asset = FakeAsset(id="a1", quality_score=50, duration=15, created_at=NOW - timedelta(days=60)) _, breakdown = score_asset(asset, now=NOW) - assert breakdown["recency"] < 5 # heavily decayed + assert breakdown["recency"] < 4 # heavily decayed # ── _duration_bucket tests ─────────────────────────────────────────────────── @@ -245,7 +245,7 @@ class TestSmartSelectAssets: assert len(results) == 1 r = results[0] assert r.score > 0 - assert set(r.breakdown.keys()) == {"quality", "duration", "recency", "unused"} + assert set(r.breakdown.keys()) >= {"quality", "duration", "recency", "unused", "ai_semantic"} def test_image_assets_can_be_selected(self): assets = [