feat(#2035): semantic tags + quality score + AI caption+embedding
CI/CD Pipeline / Check push changed paths (pull_request) Has been skipped
CI/CD Pipeline / Dedup Check - skip PR tests when covered by push pipeline (pull_request) Successful in 2s
CI/CD Pipeline / Check if frontend-only change (pull_request) Successful in 2s
CI/CD Pipeline / Frontend Lint (pull_request) Has been skipped
CI/CD Pipeline / Frontend Unit Tests (pull_request) Has been skipped
CI/CD Pipeline / PR Build Web Image (pull_request) Has been skipped
CI/CD Pipeline / Build Staging API Image (pull_request) Has been skipped
CI/CD Pipeline / Build Staging Web Image (pull_request) Has been skipped
CI/CD Pipeline / Build Staging Worker Image (pull_request) Has been skipped
CI/CD Pipeline / Retag skipped Staging API Image (pull_request) Has been skipped
CI/CD Pipeline / Retag skipped Staging Web Image (pull_request) Has been skipped
CI/CD Pipeline / Retag skipped Staging Worker Image (pull_request) Has been skipped
CI/CD Pipeline / Deploy Staging (Watchtower auto-deploy) (pull_request) Has been skipped
CI/CD Pipeline / Staging E2E Tests (pull_request) Has been skipped
CI/CD Pipeline / Staging API Integration Tests (pull_request) Has been skipped
CI/CD Pipeline / ACR Image Cleanup (pull_request) Has been skipped
CI/CD Pipeline / PR Build Worker Image (pull_request) Successful in 55s
CI/CD Pipeline / PR Build API Image (pull_request) Successful in 59s
Preview Deploy / Deploy Preview Environment (pull_request) Successful in 1m38s
CI/CD Pipeline / Integration Tests (pull_request) Successful in 2m26s
CI/CD Pipeline / Validate - Python (mypy + alembic) (pull_request) Successful in 2m35s
PR Automation / Auto Approve on CI Green (pull_request) Successful in 3m21s
CI/CD Pipeline / Build Production API Image (pull_request) Has been cancelled
CI/CD Pipeline / Build Production Web Image (pull_request) Has been cancelled
CI/CD Pipeline / Build Production Worker Image (pull_request) Has been cancelled
CI/CD Pipeline / Deploy Production (pull_request) Has been cancelled
CI/CD Pipeline / Production Browser E2E (pull_request) Has been cancelled
CI/CD Pipeline / Validate - Style (pull_request) Has been cancelled
CI/CD Pipeline / Validate - Security (pull_request) Has been cancelled
CI/CD Pipeline / Unit Tests (pull_request) Has been cancelled
CI/CD Pipeline / Canary Release to Production (pull_request) Has been cancelled
CI/CD Pipeline / CI Gate (pull_request) Has been cancelled
AI Code Review / AI Code Review (pull_request) Has been cancelled
PR Automation / Auto Merge on CI Green + Approved (pull_request) Has been cancelled

- Fix: generation_tasks.py passes clip_ai_tags_by_asset to pick_narrative_assets
  so AI tags (weight 2.0) actually participate in narrative mode selection
- Feat: quality_score auto-computation via AssetAnalyzer on ingest
  (worker.calculate_asset_price celery task; fallback 50.0 on failure)
- Feat: smart_match adds ai_semantic dimension (20% weight) using Jaccard
  similarity between asset AI tags (scene/objects/action) and script tags
- Feat: atom_clip caption (10-30 Chinese chars) via Doubao Vision, saved to
  asset_atom_clips.caption (Text column, migration 085)
- Feat: atom_clip embedding vector via Doubao embeddings API, saved to
  asset_atom_clips.embedding (JSON column)
- Chore: remove dead calculate_quality_score_real wrapper
- Tests: 20 new unit tests covering caption parsing, ai_semantic scoring,
  narrative AI tag propagation, score weight changes; update existing tests
  for new fallback dict shape and reweighted dimensions
- Fail-open: tagging/embedding/quality failures never block main flow
This commit is contained in:
saas-backend
2026-09-25 11:06:30 +08:00
parent 1be658f72d
commit 01afc2cf69
18 changed files with 664 additions and 63 deletions
@@ -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")
+51 -2
View File
@@ -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(
+1
View File
@@ -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",
@@ -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
@@ -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
@@ -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()
+8 -2
View File
@@ -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,
)
@@ -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,
@@ -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))
+1
View File
@@ -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 = ""
+6
View File
@@ -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,
)
+20 -8
View File
@@ -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
+28 -5
View File
@@ -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
+54 -10
View File
@@ -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: 按「得分 + 随机噪声」降序排序
+40
View File
@@ -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)
+32 -6
View File
@@ -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__":
+257
View File
@@ -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
+11 -11
View File
@@ -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 = [