Files
xiaoxia-saas/apps/worker/video_processing/dedup.py
T
灵应 ee8c60a143 perf(duplication): 查重模块代码优化 — 修复6个代码质量问题
1. 修复 hamming_distance 不等长哈希处理(hex() 前导零丢失)
2. 集成颜色直方图到 check_duplicate(此前计算但未使用,浪费 CPU)
3. check_duplicate 改为返回最佳匹配而非首个匹配
4. 修复仓库删除顺序(先删片段再删记录,防止孤儿数据)
5. 域模型添加 can_retry()/reset_for_retry(),仅 failed 状态允许重试
6. 列表接口暴露 offset/limit 分页参数
2026-07-01 14:48:15 +08:00

316 lines
11 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.
"""Video deduplication module - compute fingerprints and detect duplicates."""
import hashlib
import json
import logging
import os
import subprocess
import tempfile
from dataclasses import dataclass
from typing import Optional
import cv2
import numpy as np
from celery import Task
from sqlalchemy.orm import Session
from worker_app.celery_app import celery_app
from worker_app.db import SessionLocal
from packages.adapters.sqlalchemy_impl.generated_video_repository import SQLAlchemyGeneratedVideoRepository
from packages.shared.storage import get_storage_service
logger = logging.getLogger(__name__)
def compute_phash(image: np.ndarray, hash_size: int = 8) -> str:
"""Compute perceptual hash of an image using DCT."""
# Resize to 32x32 for DCT
resized = cv2.resize(image, (hash_size * 4, hash_size * 4))
gray = cv2.cvtColor(resized, cv2.COLOR_BGR2GRAY).astype(np.float32)
# Apply 2D DCT
dct = cv2.dct(gray)
# Take top-left 8x8 low-frequency components
dct_low = dct[:hash_size, :hash_size]
# Compute median (excluding DC component at [0,0])
dct_low[0, 0] = 0
median = np.median(dct_low)
# Generate hash based on comparison with median
diff = (dct_low > median).astype(int)
hash_str = "".join(str(b) for row in diff for b in row)
return hex(int(hash_str, 2))[2:]
def hamming_distance(hash1: str, hash2: str) -> int:
"""
计算两个十六进制哈希之间的汉明距离。
自动处理不等长哈希:短哈希左侧补零对齐,避免因 hex() 去掉前导零
而导致距离计算错误。
Args:
hash1: 第一个十六进制哈希字符串
hash2: 第二个十六进制哈希字符串
Returns:
汉明距离(不同位的数量)
"""
# 对齐长度:短哈希左侧补零,防止 hex() 截断前导零导致误判
max_len = max(len(hash1), len(hash2))
hash1 = hash1.zfill(max_len)
hash2 = hash2.zfill(max_len)
# 逐字符比较十六进制位,统计差异数
return sum(c1 != c2 for c1, c2 in zip(hash1, hash2))
def compute_color_histogram(image: np.ndarray, bins: int = 32) -> list[float]:
"""Compute color histogram for an image."""
hist = []
for i in range(3):
h = cv2.calcHist([image], [i], None, [bins], [0, 256])
h = cv2.normalize(h, h).flatten()
hist.extend(h)
return hist
@dataclass
class VideoFingerprint:
"""Video fingerprint containing multiple similarity metrics."""
md5: str
keyframe_phashes: list[str]
color_histograms: list[list[float]]
duration: float
resolution: tuple[int, int]
def to_dict(self) -> dict:
return {
"md5": self.md5,
"keyframe_phashes": self.keyframe_phashes,
"color_histograms": self.color_histograms,
"duration": self.duration,
"resolution": list(self.resolution),
}
class VideoDeduplicator:
"""Video deduplication using multiple fingerprint methods."""
PHASH_THRESHOLD = 10
HISTOGRAM_THRESHOLD = 0.85
def compute_fingerprint(self, video_path: str) -> VideoFingerprint:
"""Compute video fingerprint using MD5, pHash, and color histogram."""
cap = cv2.VideoCapture(video_path)
if not cap.isOpened():
raise RuntimeError(f"Cannot open video: {video_path}")
fps = cap.get(cv2.CAP_PROP_FPS)
frame_count = int(cap.get(cv2.CAP_PROP_FRAME_COUNT))
duration = frame_count / fps if fps > 0 else 0
width = int(cap.get(cv2.CAP_PROP_FRAME_WIDTH))
height = int(cap.get(cv2.CAP_PROP_FRAME_HEIGHT))
md5_hash = hashlib.md5(usedforsecurity=False)
keyframe_phashes = []
color_histograms = []
frame_interval = max(1, frame_count // 10)
for i in range(0, frame_count, frame_interval):
cap.set(cv2.CAP_PROP_POS_FRAMES, i)
ret, frame = cap.read()
if not ret:
continue
_, buffer = cv2.imencode(".jpg", frame)
md5_hash.update(buffer)
keyframe_phashes.append(compute_phash(frame))
color_histograms.append(compute_color_histogram(frame))
cap.release()
return VideoFingerprint(
md5=md5_hash.hexdigest(),
keyframe_phashes=keyframe_phashes,
color_histograms=color_histograms,
duration=duration,
resolution=(width, height),
)
def check_duplicate(self, fingerprint: VideoFingerprint, project_id: str, session: Session) -> Optional[dict]:
"""
检查视频是否与项目中已有视频重复。
采用多指标融合策略:
1. 精确匹配:MD5 完全一致 → 直接判定重复(similarity=1.0)
2. 感知相似:pHash 平均汉明距离 < PHASH_THRESHOLD
3. 颜色相似:直方图余弦相似度 > HISTOGRAM_THRESHOLD(辅助验证)
返回相似度最高的匹配结果,而非第一个匹配。
Args:
fingerprint: 待检测视频的指纹
project_id: 项目 ID,仅在同一项目内搜索
session: 数据库会话
Returns:
重复信息字典(含 duplicate, duplicate_of, reason, similarity),
或 None 表示未找到重复。
"""
video_repo = SQLAlchemyGeneratedVideoRepository(session)
existing_videos = video_repo.list_by_project(project_id)
best_match: Optional[dict] = None
for existing in existing_videos:
if not existing.video_fingerprint:
continue
ef = existing.video_fingerprint
# 精确匹配:MD5 完全一致
if fingerprint.md5 == ef.get("md5"):
return {"duplicate": True, "duplicate_of": existing.id, "reason": "exact_md5_match", "similarity": 1.0}
# 感知哈希相似度
existing_phashes = ef.get("keyframe_phashes", [])
if not existing_phashes:
continue
# 计算每个新关键帧到已有关键帧的最小汉明距离,取平均
min_distances = []
for phash in fingerprint.keyframe_phashes:
distances = [hamming_distance(phash, ep) for ep in existing_phashes]
min_distances.append(min(distances))
avg_distance = sum(min_distances) / len(min_distances) if min_distances else 100
if avg_distance >= self.PHASH_THRESHOLD:
continue
phash_similarity = 1.0 - (avg_distance / 64)
# 颜色直方图辅助验证(如果可用)
existing_histograms = ef.get("color_histograms", [])
final_similarity = phash_similarity
reason = "phash_similar"
if existing_histograms and fingerprint.color_histograms:
hist_sim = self._average_histogram_similarity(
fingerprint.color_histograms, existing_histograms
)
if hist_sim >= self.HISTOGRAM_THRESHOLD:
# 双指标加权:pHash 60% + 直方图 40%
final_similarity = 0.6 * phash_similarity + 0.4 * hist_sim
reason = "phash+histogram"
else:
# 直方图不达标,降低置信度但仍以 pHash 为主
final_similarity = phash_similarity * 0.8
reason = "phash_only"
# 保留最佳匹配
if best_match is None or final_similarity > best_match["similarity"]:
best_match = {
"duplicate": True,
"duplicate_of": existing.id,
"reason": reason,
"similarity": round(final_similarity, 4),
}
return best_match
@staticmethod
def _average_histogram_similarity(
histograms_a: list[list[float]], histograms_b: list[list[float]]
) -> float:
"""
计算两组颜色直方图之间的平均余弦相似度。
对每组直方图对取最小长度对齐,计算余弦相似度后取平均。
Args:
histograms_a: 第一组直方图(每帧一个 list)
histograms_b: 第二组直方图
Returns:
平均余弦相似度,范围 [0, 1]
"""
if not histograms_a or not histograms_b:
return 0.0
similarities = []
for ha in histograms_a:
best = 0.0
vec_a = np.array(ha, dtype=np.float64)
norm_a = np.linalg.norm(vec_a)
if norm_a == 0:
continue
for hb in histograms_b:
vec_b = np.array(hb, dtype=np.float64)
# 对齐长度
min_len = min(len(vec_a), len(vec_b))
va, vb = vec_a[:min_len], vec_b[:min_len]
norm_b = np.linalg.norm(vb)
if norm_b == 0:
continue
sim = float(np.dot(va, vb) / (norm_a * norm_b))
best = max(best, sim)
similarities.append(best)
return sum(similarities) / len(similarities) if similarities else 0.0
@celery_app.task(bind=True, max_retries=3, name="worker.check_duplicate")
def check_duplicate_task(self: Task, generated_video_id: str) -> dict:
"""Celery task to check if generated video is a duplicate."""
session = SessionLocal()
temp_dir = tempfile.mkdtemp()
try:
video_repo = SQLAlchemyGeneratedVideoRepository(session)
storage_service = get_storage_service()
deduplicator = VideoDeduplicator()
video = video_repo.get(generated_video_id)
if video is None:
raise ValueError(f"Generated video {generated_video_id} not found")
local_path = os.path.join(temp_dir, f"{generated_video_id}.mp4")
storage_key = video.file_url.split("/")[-1]
storage_service.download_file(
f"projects/{video.project_id}/generated/{generated_video_id}/{generated_video_id}.mp4", local_path
)
fingerprint = deduplicator.compute_fingerprint(local_path)
duplicate_result = deduplicator.check_duplicate(fingerprint, video.project_id, session)
video.video_fingerprint = fingerprint.to_dict()
if duplicate_result:
video.is_duplicate = True
video.duplicate_of = duplicate_result["duplicate_of"]
else:
video.is_duplicate = False
video.duplicate_of = None
video_repo.update(video)
session.commit()
logger.info(f"Duplicate check completed for video {generated_video_id}: is_duplicate={video.is_duplicate}")
return {
"ok": True,
"video_id": generated_video_id,
"is_duplicate": video.is_duplicate,
"duplicate_of": video.duplicate_of,
"fingerprint": fingerprint.to_dict(),
}
except Exception as e:
logger.error(f"Duplicate check failed for {generated_video_id}: {str(e)}")
session.rollback()
raise self.retry(exc=e, countdown=60)
finally:
session.close()
import shutil
shutil.rmtree(temp_dir, ignore_errors=True)