feat(dedup): 动态抽帧 + 滑动窗口时序匹配 (#1659) #1673

Merged
xiaoxia merged 3 commits from feature/1659-dynamic-keyframe-sliding-window into develop 2026-09-03 23:16:51 +08:00
7 changed files with 1065 additions and 198 deletions
+465 -103
View File
@@ -1,8 +1,12 @@
"""Video deduplication module - compute fingerprints and detect duplicates."""
"""Video deduplication module - compute fingerprints and detect duplicates.
Dynamic keyframe detection + sliding window temporal matching (Issue #1659).
"""
import hashlib
import logging
import os
import statistics
import tempfile
from dataclasses import dataclass, field
from typing import Optional
@@ -21,10 +25,28 @@ from packages.shared.storage import get_storage_service
logger = logging.getLogger(__name__)
# 分片策略常量
SHORT_VIDEO_CHUNK_SEC = 2 # ≤60秒视频,每 2 秒一个分片
LONG_VIDEO_CHUNK_SEC = 5 # >60秒视频,每 5 秒一个分片
SHORT_VIDEO_THRESHOLD_SEC = 60
# ── 关键帧检测常量 ──────────────────────────────────────────────
SCENE_CHANGE_THRESHOLD = 30 # 灰度差异阈值
MIN_KEYFRAME_INTERVAL_SEC = 1.0 # 最小关键帧间隔(秒)
MAX_KEYFRAMES = 30 # 最大关键帧数
MIN_KEYFRAMES = 5 # 最小关键帧数
LONG_VIDEO_SEGMENT_SEC = 30 # 长视频每段秒数
LONG_VIDEO_DURATION_THRESHOLD_SEC = 180 # 3 分钟阈值
MIN_FRAMES_PER_SEGMENT = 2 # 长视频每段最少帧数
# ── 滑动窗口匹配常量 ────────────────────────────────────────────
SEGMENT_MATCH_THRESHOLD = 8 # 帧匹配汉明距离阈值
MIN_CONSECUTIVE_MATCHES = 5 # 最少连续匹配帧数
MAX_GAP = 2 # 允许的最大间隙帧数
# ── 融合判定常量 ────────────────────────────────────────────────
PHASH_WEIGHT = 0.7 # pHash 权重
HISTOGRAM_WEIGHT = 0.3 # 直方图权重
MATCH_RATIO_THRESHOLD = 0.7 # 至少 70% 帧匹配
DUPLICATE_THRESHOLD = 0.70 # 融合后相似度阈值
# ── 感知哈希 & 颜色直方图工具函数 ────────────────────────────────
def compute_phash(image: np.ndarray, hash_size: int = 8) -> str:
@@ -87,15 +109,107 @@ def compute_color_histogram(image: np.ndarray, bins: int = 32) -> list[float]:
return hist
def compute_chunk_interval(duration: float) -> float:
"""根据视频时长返回分片间隔(秒)。
# ── 关键帧检测 ──────────────────────────────────────────────────
短视频(≤60秒):每 2 秒一个分片
长视频(>60秒):每 5 秒一个分片
def detect_keyframe_timestamps(
video_path: str,
*,
min_interval_sec: float = MIN_KEYFRAME_INTERVAL_SEC,
max_frames: int = MAX_KEYFRAMES,
min_frames: int = MIN_KEYFRAMES,
) -> list[float]:
"""检测视频中的场景切换点,返回关键帧时间戳列表(秒)。
算法:
1. 降采样到 320x240,逐帧转灰度
2. 计算相邻帧灰度差异(像素均值差)
3. 差异 > SCENE_CHANGE_THRESHOLD(30) 标记为候选关键帧
4. 相邻关键帧间隔 < min_interval_sec 的,保留差异更大的那个
5. 数量裁剪到 [min_frames, max_frames]
对于长视频(>3分钟):
- 每 30 秒一个分段
- 每个分段至少选 2 个关键帧(如果分段内无场景切换,均匀取 2 帧)
"""
if duration <= SHORT_VIDEO_THRESHOLD_SEC:
return SHORT_VIDEO_CHUNK_SEC
return LONG_VIDEO_CHUNK_SEC
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
if duration <= 0:
cap.release()
return []
# 逐帧检测场景切换
candidates: list[tuple[float, float]] = [] # (timestamp_sec, diff_score)
prev_gray = None
while True:
ret, frame = cap.read()
if not ret:
break
# 降采样 + 灰度
small = cv2.resize(frame, (320, 240))
gray = cv2.cvtColor(small, cv2.COLOR_BGR2GRAY).astype(np.float32)
if prev_gray is not None:
diff = float(np.mean(np.abs(gray - prev_gray)))
if diff > SCENE_CHANGE_THRESHOLD:
pos_ms = cap.get(cv2.CAP_PROP_POS_MSEC)
candidates.append((pos_ms / 1000.0, diff))
prev_gray = gray
cap.release()
# 按最小间隔过滤(保留差异更大的)
filtered: list[tuple[float, float]] = []
for ts, diff in sorted(candidates):
if filtered and (ts - filtered[-1][0]) < min_interval_sec:
if diff > filtered[-1][1]:
filtered[-1] = (ts, diff)
else:
filtered.append((ts, diff))
keyframe_times = [ts for ts, _ in filtered]
# 数量不足 min_frames 时,在时间轴上均匀补充
if len(keyframe_times) < min_frames:
uniform = [duration * (i + 0.5) / min_frames for i in range(min_frames)]
keyframe_times = sorted(set(uniform) | set(keyframe_times))
# 如果合并后还不足 min_frames,直接用均匀分布
if len(keyframe_times) < min_frames:
keyframe_times = uniform
# 数量超过 max_frames 时,均匀采样
if len(keyframe_times) > max_frames:
step = len(keyframe_times) / max_frames
keyframe_times = [keyframe_times[int(i * step)] for i in range(max_frames)]
# 长视频分段保底(>3分钟)
if duration > LONG_VIDEO_DURATION_THRESHOLD_SEC:
segment_count = int(duration / LONG_VIDEO_SEGMENT_SEC)
for seg_idx in range(segment_count):
seg_start = seg_idx * LONG_VIDEO_SEGMENT_SEC
seg_end = min((seg_idx + 1) * LONG_VIDEO_SEGMENT_SEC, duration)
seg_frames = [t for t in keyframe_times if seg_start <= t < seg_end]
if len(seg_frames) < MIN_FRAMES_PER_SEGMENT:
# 均匀补齐
for i in range(MIN_FRAMES_PER_SEGMENT):
t = seg_start + LONG_VIDEO_SEGMENT_SEC * (i + 0.5) / MIN_FRAMES_PER_SEGMENT
if t not in keyframe_times and seg_start <= t < seg_end:
keyframe_times.append(t)
keyframe_times.sort()
return keyframe_times
# ── 数据类 ──────────────────────────────────────────────────────
@dataclass
@@ -109,6 +223,17 @@ class FingerprintChunk:
frame_count: int = 1
@dataclass
class DuplicateSegment:
"""一段重复片段的描述。"""
query_start_ms: int
query_end_ms: int
target_start_ms: int
target_end_ms: int
avg_distance: float # 该段内帧的平均汉明距离
@dataclass
class VideoFingerprint:
"""Video fingerprint containing multiple similarity metrics."""
@@ -163,6 +288,137 @@ class VideoFingerprint:
return models
# ── 滑动窗口时序匹配 ────────────────────────────────────────────
def find_duplicate_segments(
query_chunks: list,
target_chunks: list,
*,
match_threshold: int = SEGMENT_MATCH_THRESHOLD,
min_consecutive: int = MIN_CONSECUTIVE_MATCHES,
max_gap: int = MAX_GAP,
) -> list[DuplicateSegment]:
"""滑动窗口时序匹配:找出两组分片之间的重复片段。
算法:
1. 对每个 query chunk,找到 target 中汉明距离最小的 chunk
2. 距离 <= match_threshold 视为匹配
3. 找连续匹配的 run(允许 max_gap 帧间隙)
4. 连续匹配数 >= min_consecutive 的 run 报告为重复片段
Args:
query_chunks: 查询视频的分片列表(FingerprintChunk 或 dict
target_chunks: 目标视频的分片列表
match_threshold: 汉明距离匹配阈值
min_consecutive: 最少连续匹配帧数
max_gap: 允许的最大间隙帧数
Returns:
DuplicateSegment 列表
"""
if not query_chunks or not target_chunks:
return []
def _get_phash(chunk) -> str:
if isinstance(chunk, dict):
return chunk["phash_binary"]
return chunk.phash_binary
def _get_start(chunk) -> int:
if isinstance(chunk, dict):
return chunk["start_time_ms"]
return chunk.start_time_ms
def _get_end(chunk) -> int:
if isinstance(chunk, dict):
return chunk["end_time_ms"]
return chunk.end_time_ms
# Step 1: 逐帧匹配
frame_matches: list[tuple[bool, int, int]] = [] # (is_match, min_dist, best_target_idx)
for qc in query_chunks:
qc_phash = _get_phash(qc)
best_dist = 64
best_idx = 0
for j, tc in enumerate(target_chunks):
d = hamming_distance(qc_phash, _get_phash(tc))
if d < best_dist:
best_dist = d
best_idx = j
frame_matches.append((best_dist <= match_threshold, best_dist, best_idx))
# Step 2: 找连续匹配的 runs
runs: list[tuple[int, int]] = [] # list of (start_idx, end_idx)
run_start = None
gap_count = 0
for i, (is_match, _dist, _idx) in enumerate(frame_matches):
if is_match:
if run_start is None:
run_start = i
gap_count = 0 # 重置间隙
else:
if run_start is not None:
gap_count += 1
if gap_count > max_gap:
# 中断当前 run
run_end = i - gap_count # 最后一个匹配帧的索引
# 计算 run 内的实际匹配帧数(总跨度 - 间隙数)
total_gaps = sum(1 for k in range(run_start, run_end + 1) if not frame_matches[k][0])
matching_count = (run_end - run_start + 1) - total_gaps
if matching_count >= min_consecutive:
runs.append((run_start, run_end))
run_start = None
gap_count = 0
# 处理末尾 run
if run_start is not None:
last_idx = len(frame_matches) - 1
# 回退找到最后一个匹配帧的位置(跳过尾部非匹配帧)
while last_idx >= run_start and not frame_matches[last_idx][0]:
last_idx -= 1
if last_idx >= run_start:
# 计算 run 内的总间隙数
total_gaps = sum(1 for k in range(run_start, last_idx + 1) if not frame_matches[k][0])
matching_count = (last_idx - run_start + 1) - total_gaps
if matching_count >= min_consecutive:
runs.append((run_start, last_idx))
# Step 3: 构建 DuplicateSegment
segments: list[DuplicateSegment] = []
for start, end in runs:
query_start = _get_start(query_chunks[start])
query_end = _get_end(query_chunks[end])
# 取目标范围(按最佳匹配的目标 chunk 时间范围)
target_indices = [frame_matches[k][2] for k in range(start, end + 1) if frame_matches[k][0]]
if target_indices:
t_min = min(target_indices)
t_max = max(target_indices)
target_start = _get_start(target_chunks[t_min])
target_end = _get_end(target_chunks[t_max])
else:
target_start = _get_start(target_chunks[0])
target_end = _get_end(target_chunks[-1])
avg_dist = sum(frame_matches[k][1] for k in range(start, end + 1)) / (end - start + 1)
segments.append(
DuplicateSegment(
query_start_ms=query_start,
query_end_ms=query_end,
target_start_ms=target_start,
target_end_ms=target_end,
avg_distance=avg_dist,
)
)
return segments
# ── VideoDeduplicator ───────────────────────────────────────────
class VideoDeduplicator:
"""Video deduplication using multiple fingerprint methods."""
@@ -170,11 +426,11 @@ class VideoDeduplicator:
HISTOGRAM_THRESHOLD = 0.85
def compute_fingerprint(self, video_path: str) -> VideoFingerprint:
"""Compute video fingerprint using MD5, pHash, and color histogram.
"""Compute video fingerprint using dynamic keyframe detection.
按时间分片抽帧:短视频(≤60s)每 2s 一片,长视频每 5s 一片。
每片取 1 帧计算 pHash + color_histogram。
同时保留 keyframe_phashes/color_histograms 聚合字段(向后兼容)
使用 detect_keyframe_timestamps() 检测内容感知关键帧,
在每个关键帧处取帧计算 pHash + color_histogram。
同时保留 MD5 计算和分片数据结构
"""
cap = cv2.VideoCapture(video_path)
if not cap.isOpened():
@@ -186,41 +442,55 @@ class VideoDeduplicator:
width = int(cap.get(cv2.CAP_PROP_FRAME_WIDTH))
height = int(cap.get(cv2.CAP_PROP_FRAME_HEIGHT))
cap.release()
# 1. 检测关键帧时间戳
keyframe_times = detect_keyframe_timestamps(video_path)
if not keyframe_times:
return VideoFingerprint(
md5="",
keyframe_phashes=[],
color_histograms=[],
duration=duration,
resolution=(width, height),
chunks=[],
)
# 2. 打开视频,逐个关键帧取帧
cap = cv2.VideoCapture(video_path)
md5_hash = hashlib.md5(usedforsecurity=False)
chunks: list[FingerprintChunk] = []
# 分片间隔(秒)
chunk_interval_sec = compute_chunk_interval(duration)
chunk_interval_ms = int(chunk_interval_sec * 1000)
duration_ms = int(duration * 1000)
# 遍历每个分片时间窗口,取 1 帧
start_ms = 0
while start_ms < duration_ms:
end_ms = min(start_ms + chunk_interval_ms, duration_ms)
# 定位到分片中点
seek_ms = (start_ms + end_ms) / 2
for i, t_sec in enumerate(keyframe_times):
seek_ms = t_sec * 1000
cap.set(cv2.CAP_PROP_POS_MSEC, seek_ms)
ret, frame = cap.read()
if ret:
# MD5 计算
_, buffer = cv2.imencode(".jpg", frame)
md5_hash.update(buffer)
if not ret:
continue
phash = compute_phash(frame)
hist = compute_color_histogram(frame)
# MD5 计算
_, buffer = cv2.imencode(".jpg", frame)
md5_hash.update(buffer)
chunks.append(
FingerprintChunk(
start_time_ms=start_ms,
end_time_ms=end_ms,
phash_binary=phash,
color_histogram=hist,
frame_count=1,
)
phash = compute_phash(frame)
hist = compute_color_histogram(frame)
# 计算分片时间范围(从前一个关键帧到下一个关键帧的中点)
prev_boundary = keyframe_times[i - 1] * 1000 if i > 0 else 0
next_boundary = keyframe_times[i + 1] * 1000 if i < len(keyframe_times) - 1 else duration * 1000
start_ms = int((prev_boundary + seek_ms) / 2)
end_ms = int((seek_ms + next_boundary) / 2)
chunks.append(
FingerprintChunk(
start_time_ms=start_ms,
end_time_ms=end_ms,
phash_binary=phash,
color_histogram=hist,
frame_count=1,
)
start_ms = end_ms
)
cap.release()
@@ -255,14 +525,39 @@ class VideoDeduplicator:
for r in rows
]
@staticmethod
def _bhattacharyya_coefficient(hist_a: list[float], hist_b: list[float]) -> float:
"""Bhattacharyya 系数:Σ √(a[i] * b[i]),范围 [0, 1]1=完全相同。"""
min_len = min(len(hist_a), len(hist_b))
a = hist_a[:min_len]
b = hist_b[:min_len]
return float(sum(np.sqrt(ai * bi) for ai, bi in zip(a, b, strict=False)))
@staticmethod
def _compute_histogram_similarity(
histograms_a: list[list[float]],
histograms_b: list[list[float]],
) -> float:
"""对每组直方图,找到最佳匹配的 Bhattacharyya 系数,取平均。"""
if not histograms_a or not histograms_b:
return 0.0
similarities = []
for ha in histograms_a:
best = 0.0
for hb in histograms_b:
bc = VideoDeduplicator._bhattacharyya_coefficient(ha, hb)
best = max(best, bc)
similarities.append(best)
return sum(similarities) / len(similarities) if similarities else 0.0
def check_duplicate(self, fingerprint: VideoFingerprint, project_id: str, session: Session) -> Optional[dict]:
"""检查视频是否与项目中已有视频重复。
查重逻辑:
1. MD5 精确匹配 → similarity=1.0
2. pHash 相似度(优先从分片表读取,回退到 JSON 字段)
2. pHash 中位数距离 + 帧匹配比例 + 直方图融合判定
判定阈值:avg_distance < PHASH_THRESHOLD(10)
判定为重复后,调用 find_duplicate_segments() 获取具体重复片段。
Args:
fingerprint: 待检测视频的指纹
@@ -270,7 +565,7 @@ class VideoDeduplicator:
session: 数据库会话
Returns:
重复信息字典(含 duplicate, duplicate_of, reason, similarity),
重复信息字典(含 duplicate, duplicate_of, reason, similarity, duplicate_segments),
或 None 表示未找到重复。
"""
video_repo = SQLAlchemyGeneratedVideoRepository(session)
@@ -298,23 +593,65 @@ class VideoDeduplicator:
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:
# 帧匹配比例检查
matching_frames = sum(1 for d in min_distances if d < self.PHASH_THRESHOLD)
match_ratio = matching_frames / len(min_distances) if min_distances else 0
if match_ratio < 0.7:
continue
phash_similarity = 1.0 - (avg_distance / 64)
# 中位数距离
median_distance = statistics.median(min_distances) if min_distances else 64
if median_distance >= self.PHASH_THRESHOLD:
continue
# 直方图融合
existing_histograms = []
if chunk_data:
existing_histograms = [c["color_histogram"] for c in chunk_data if c.get("color_histogram")]
else:
existing_histograms = ef.get("color_histograms", [])
phash_similarity = 1.0 - (median_distance / 64)
hist_similarity = (
self._compute_histogram_similarity(fingerprint.color_histograms, existing_histograms)
if existing_histograms
else 0.5
)
combined_score = 0.7 * phash_similarity + 0.3 * hist_similarity
# DUPLICATE_THRESHOLD from module level
if combined_score < DUPLICATE_THRESHOLD:
continue
# 滑动窗口时序匹配:获取具体重复片段
existing_chunk_objects = (
chunk_data
if chunk_data
else [{"phash_binary": p, "start_time_ms": 0, "end_time_ms": 0} for p in existing_phashes]
)
segments = find_duplicate_segments(fingerprint.chunks, existing_chunk_objects)
return {
"duplicate": True,
"duplicate_of": existing.id,
"reason": "phash_similar",
"similarity": phash_similarity,
"reason": "phash_histogram_fusion",
"similarity": combined_score,
"duplicate_segments": [
{
"query_start_ms": s.query_start_ms,
"query_end_ms": s.query_end_ms,
"target_start_ms": s.target_start_ms,
"target_end_ms": s.target_end_ms,
"avg_distance": round(s.avg_distance, 2),
}
for s in segments
],
}
return None
@@ -328,7 +665,8 @@ class VideoDeduplicator:
) -> Optional[dict]:
"""检查视频是否与同批次内其他视频重复。
逻辑与 check_duplicate 一致(MD5 + pHash),但搜索范围限定为同 batch_id 的视频。
逻辑与 check_duplicate 一致(MD5 + pHash + 直方图融合 + 时序匹配),
但搜索范围限定为同 batch_id 的视频。
Args:
fingerprint: 待检测视频的指纹
@@ -373,59 +711,63 @@ class VideoDeduplicator:
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:
# 帧匹配比例检查
matching_frames = sum(1 for d in min_distances if d < self.PHASH_THRESHOLD)
match_ratio = matching_frames / len(min_distances) if min_distances else 0
if match_ratio < 0.7:
continue
phash_similarity = 1.0 - (avg_distance / 64)
median_distance = statistics.median(min_distances) if min_distances else 64
if median_distance >= self.PHASH_THRESHOLD:
continue
# 直方图融合
existing_histograms = []
if chunk_data:
existing_histograms = [c["color_histogram"] for c in chunk_data if c.get("color_histogram")]
else:
existing_histograms = ef.get("color_histograms", [])
phash_similarity = 1.0 - (median_distance / 64)
hist_similarity = (
self._compute_histogram_similarity(fingerprint.color_histograms, existing_histograms)
if existing_histograms
else 0.5
)
combined_score = 0.7 * phash_similarity + 0.3 * hist_similarity
# DUPLICATE_THRESHOLD from module level
if combined_score < DUPLICATE_THRESHOLD:
continue
# 滑动窗口时序匹配
existing_chunk_objects = (
chunk_data
if chunk_data
else [{"phash_binary": p, "start_time_ms": 0, "end_time_ms": 0} for p in existing_phashes]
)
segments = find_duplicate_segments(fingerprint.chunks, existing_chunk_objects)
return {
"duplicate": True,
"duplicate_of": existing.id,
"reason": "batch_phash_similar",
"similarity": phash_similarity,
"reason": "batch_phash_histogram_fusion",
"similarity": combined_score,
"duplicate_segments": [
{
"query_start_ms": s.query_start_ms,
"query_end_ms": s.query_end_ms,
"target_start_ms": s.target_start_ms,
"target_end_ms": s.target_end_ms,
"avg_distance": round(s.avg_distance, 2),
}
for s in segments
],
}
return None
@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
def compute_duplicate_rate(
self,
fingerprint: VideoFingerprint,
@@ -438,9 +780,9 @@ class VideoDeduplicator:
"""计算当前视频与用户库内已有视频的最高相似度百分比。
优先按 user_id 全局比较(跨项目),user_id 为空时回退到项目级比较。
遍历最近 200 个其他有指纹的视频,对每个计算相似度:
遍历最近 200 个其他有指纹的视频,对每个计算融合相似度:
- MD5 精确匹配 → 100%
- pHash 相似度 → (1.0 - avg_distance / 64) * 100
- pHash + 直方图融合 → 0.7 * phash_sim + 0.3 * hist_sim
取最高值作为 duplicate_rate0~100)。
如果没有其他视频可比较,返回 0.0。
@@ -454,7 +796,6 @@ class VideoDeduplicator:
Returns:
duplicate_rate: 0~100 的浮点数
"""
# 限制查询最近 200 个视频,避免大库内存溢出
from packages.adapters.sqlalchemy_impl.models import GeneratedVideoModel
# 优先按 user_id 全局比较(跨项目),否则回退到项目级
@@ -469,7 +810,7 @@ class VideoDeduplicator:
)
logger.debug("compute_duplicate_rate: project-level fallback project_id=%s", project_id)
# 排除当前视频自身(记录可能已写入 DB,必须在查询层排除)
# 排除当前视频自身
if current_video_id:
query = query.filter(GeneratedVideoModel.id != current_video_id)
@@ -505,9 +846,30 @@ class VideoDeduplicator:
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 64
similarity = (1.0 - avg_distance / 64) * 100
max_similarity = max(max_similarity, similarity)
# 帧匹配比例检查
matching_frames = sum(1 for d in min_distances if d < self.PHASH_THRESHOLD)
match_ratio = matching_frames / len(min_distances) if min_distances else 0
if match_ratio < 0.7:
continue
median_distance = statistics.median(min_distances) if min_distances else 64
# 直方图融合
existing_histograms = []
if chunk_data:
existing_histograms = [c["color_histogram"] for c in chunk_data if c.get("color_histogram")]
else:
existing_histograms = ef.get("color_histograms", [])
phash_similarity = (1.0 - median_distance / 64) * 100
hist_similarity = (
self._compute_histogram_similarity(fingerprint.color_histograms, existing_histograms) * 100
if existing_histograms
else 50.0
)
combined_score = 0.7 * phash_similarity + 0.3 * hist_similarity
max_similarity = max(max_similarity, combined_score)
return round(max(max_similarity, 0.0), 2)
-1
View File
@@ -221,7 +221,6 @@ class TestServiceLayerIntegration:
def _make_service(self, clip_repo_mock, asset_repo_mock=None):
"""创建 PlanGeneratorService 并注入 mock repos."""
from unittest.mock import MagicMock, patch
from apps.api.app.services.plan_generator_service import PlanGeneratorService
+13 -6
View File
@@ -285,8 +285,10 @@ class TestVideoDeduplicatorCheckDuplicate:
result = deduplicator.check_duplicate(fingerprint, "proj-1", mock_session)
assert result is not None
assert result["duplicate"] is True
assert result["similarity"] == 1.0 # distance=0 → 1.0
assert result["reason"] == "phash_similar"
assert result["similarity"] == pytest.approx(
0.85, abs=0.01
) # combined: 0.7*1.0 + 0.3*0.5 (no hist fallback)
assert result["reason"] == "phash_histogram_fusion"
finally:
self._restore_repo(mod, orig)
@@ -425,8 +427,11 @@ class TestVideoDeduplicatorCheckDuplicate:
result = deduplicator.check_duplicate(fingerprint, "proj-1", mock_session)
assert result is not None
assert result["duplicate"] is True
# similarity = 1.0 - (1 / 64) = 0.984375
assert abs(result["similarity"] - (1.0 - 1.0 / 64)) < 1e-6
# 新算法: median_distance=1, phash_sim=1-1/64=0.984375
# 无直方图 → hist_sim=0.5(fallback)
# combined = 0.7*0.984375 + 0.3*0.5 = 0.839062
expected_sim = 0.7 * (1.0 - 1.0 / 64) + 0.3 * 0.5
assert abs(result["similarity"] - expected_sim) < 1e-6
finally:
self._restore_repo(mod, orig)
@@ -456,7 +461,9 @@ class TestVideoDeduplicatorCheckDuplicate:
result = deduplicator.check_duplicate(fingerprint, "proj-1", mock_session)
assert result is not None
assert result["duplicate"] is True
assert result["similarity"] == 1.0 # avg_distance = 0
# 新算法: median_distance=0, phash_sim=1.0, hist_sim=0.5(fallback)
# combined = 0.7*1.0 + 0.3*0.5 = 0.85
assert result["similarity"] == pytest.approx(0.85, abs=0.01)
finally:
self._restore_repo(mod, orig)
@@ -539,7 +546,7 @@ class TestVideoDeduplicatorCheckBatchDuplicate:
result = deduplicator.check_batch_duplicate(fingerprint, "batch-1", "vid-self", mock_session)
assert result is not None
assert result["duplicate"] is True
assert result["reason"] == "batch_phash_similar"
assert result["reason"] == "batch_phash_histogram_fusion"
finally:
self._restore_repo(mod, orig)
+53 -55
View File
@@ -181,70 +181,68 @@ class TestVideoFingerprint:
assert d["color_histograms"] == []
class TestAverageHistogramSimilarity:
"""_average_histogram_similarity 直方图相似度测试."""
class TestBhattacharyyaCoefficient:
"""_bhattacharyya_coefficient Bhattacharyya 系数测试."""
def test_identical_histograms(self):
"""完全相同的直方图相似度为1.0."""
hist = [[0.5, 0.5, 0.0], [0.3, 0.4, 0.3]]
sim = VideoDeduplicator._average_histogram_similarity(hist, hist)
assert sim == pytest.approx(1.0)
"""完全相同的直方图系数为1.0."""
hist = [0.5, 0.5, 0.0, 0.3]
bc = VideoDeduplicator._bhattacharyya_coefficient(hist, hist)
# Σ √(a[i]*a[i]) = Σ a[i] = 1.0 (normalized)
assert bc == pytest.approx(sum(h for h in hist))
def test_empty_first_list(self):
def test_zero_histograms(self):
"""全零直方图系数为0."""
bc = VideoDeduplicator._bhattacharyya_coefficient([0.0, 0.0], [0.0, 0.0])
assert bc == 0.0
def test_orthogonal_histograms(self):
"""正交直方图(无重叠)系数为0."""
bc = VideoDeduplicator._bhattacharyya_coefficient([1.0, 0.0], [0.0, 1.0])
assert bc == pytest.approx(0.0)
def test_different_lengths(self):
"""不同长度直方图取最小长度对齐."""
bc = VideoDeduplicator._bhattacharyya_coefficient([1.0, 1.0, 0.0, 0.0], [1.0, 1.0])
# 对齐到前2维: √(1*1) + √(1*1) = 2.0
assert bc == pytest.approx(2.0)
def test_known_value(self):
"""已知值验证."""
# [0.25, 0.25, 0.25, 0.25] vs [0.25, 0.25, 0.25, 0.25]
# BC = 4 * √(0.25 * 0.25) = 4 * 0.25 = 1.0
hist = [0.25, 0.25, 0.25, 0.25]
bc = VideoDeduplicator._bhattacharyya_coefficient(hist, hist)
assert bc == pytest.approx(1.0)
class TestComputeHistogramSimilarity:
"""_compute_histogram_similarity 多帧直方图相似度测试."""
def test_identical_histogram_groups(self):
"""完全相同的两组直方图."""
hist = [[0.5, 0.5], [0.3, 0.4]]
sim = VideoDeduplicator._compute_histogram_similarity(hist, hist)
# Each hist finds best match = itself
assert sim > 0.0
def test_empty_first(self):
"""第一组为空返回0."""
sim = VideoDeduplicator._average_histogram_similarity([], [[0.5, 0.5]])
assert sim == 0.0
assert VideoDeduplicator._compute_histogram_similarity([], [[0.5]]) == 0.0
def test_empty_second_list(self):
def test_empty_second(self):
"""第二组为空返回0."""
sim = VideoDeduplicator._average_histogram_similarity([[0.5, 0.5]], [])
assert sim == 0.0
assert VideoDeduplicator._compute_histogram_similarity([[0.5]], []) == 0.0
def test_both_empty(self):
"""两组都为空返回0."""
sim = VideoDeduplicator._average_histogram_similarity([], [])
assert sim == 0.0
assert VideoDeduplicator._compute_histogram_similarity([], []) == 0.0
def test_orthogonal_histograms(self):
"""正交直方图相似度为0."""
# [1, 0] 和 [0, 1] 正交
sim = VideoDeduplicator._average_histogram_similarity([[1.0, 0.0]], [[0.0, 1.0]])
assert sim == pytest.approx(0.0)
def test_partial_similarity(self):
"""部分相似."""
# [1, 1] 和 [1, 0] 的余弦相似度 = 1/√2 ≈ 0.707
sim = VideoDeduplicator._average_histogram_similarity([[1.0, 1.0]], [[1.0, 0.0]])
assert sim == pytest.approx(1.0 / (2**0.5), rel=0.01)
def test_multiple_frames_best_match(self):
def test_best_match_selection(self):
"""多帧时取最佳匹配."""
# 第一帧完全不同,第二帧完全相同 → 平均 best = (0 + 1) / 2 = 0.5
sim = VideoDeduplicator._average_histogram_similarity(
[[1.0, 0.0], [0.0, 1.0]],
[[0.0, 1.0]], # 只有一帧,和第一帧0相似,和第二帧1相似
)
# 第一帧最佳匹配=0,第二帧最佳匹配=1,平均=0.5
assert sim == pytest.approx(0.5)
def test_zero_norm_histogram_skipped(self):
"""零范数直方图被跳过."""
sim = VideoDeduplicator._average_histogram_similarity([[0.0, 0.0]], [[1.0, 1.0]])
# 第一组的零范数被跳过,similarities为空,返回0
assert sim == 0.0
def test_different_length_histograms(self):
"""不同长度的直方图取最小长度对齐."""
sim = VideoDeduplicator._average_histogram_similarity(
[[1.0, 1.0, 0.0, 0.0]], # 4维
[[1.0, 1.0]], # 2维
)
# 对齐到前2维,都是[1,1],相似度1.0
# ha[0] 与 hb[0] 正交,与 hb[1] 完全相同
a = [[1.0, 0.0]]
b = [[0.0, 1.0], [1.0, 0.0]]
sim = VideoDeduplicator._compute_histogram_similarity(a, b)
# Best match for [1,0]: max(BC([1,0],[0,1]), BC([1,0],[1,0])) = max(0, 1) = 1
assert sim == pytest.approx(1.0)
def test_similarity_in_zero_one_range(self):
"""相似度在[0, 1]范围内."""
hist_a = [np.random.rand(96).tolist() for _ in range(5)]
hist_b = [np.random.rand(96).tolist() for _ in range(5)]
sim = VideoDeduplicator._average_histogram_similarity(hist_a, hist_b)
assert 0.0 <= sim <= 1.0
+532
View File
@@ -0,0 +1,532 @@
"""Issue #1659: 动态抽帧 + 滑动窗口时序匹配 单元测试.
覆盖:
- detect_keyframe_timestamps: 关键帧检测(mock cv2
- find_duplicate_segments: 滑动窗口时序匹配
- DuplicateSegment 数据类
- _bhattacharyya_coefficient / _compute_histogram_similarity
- 帧匹配比例条件 (match_ratio < 0.7 → 跳过)
- 中位数 vs 均值(抵抗异常值)
- 向后兼容(无分片数据时不崩溃)
"""
from __future__ import annotations
import sys
from unittest.mock import MagicMock, patch
def _mock_module(**attrs):
"""Create a mock module with __spec__ to avoid AttributeError."""
m = MagicMock()
m.__spec__ = None
for k, v in attrs.items():
setattr(m, k, v)
return m
# ── Module-level setup: mock deps, import dedup, then restore sys.modules ──
_SAVED_MODULES_KEYS = set(sys.modules.keys())
_SAVED_MODULES_VALUES = {
k: sys.modules.get(k)
for k in [
"cv2",
"celery",
"sqlalchemy",
"sqlalchemy.orm",
"sqlalchemy.engine",
"sqlalchemy.ext",
"sqlalchemy.ext.declarative",
"worker_app.db",
"worker_app.celery_app",
"worker_app.core.config",
"packages.adapters.sqlalchemy_impl.session",
"packages.adapters.sqlalchemy_impl.generated_video_repository",
"packages.adapters.sqlalchemy_impl.models",
"packages.shared.config",
"packages.shared.storage",
]
}
sys.modules["cv2"] = _mock_module()
_mock_celery = MagicMock()
_mock_celery.Task = MagicMock
_mock_celery.Celery = MagicMock
_mock_celery.__spec__ = None
sys.modules["celery"] = _mock_celery
_mock_sqla = MagicMock()
_mock_sqla.__path__ = []
_mock_sqla.__spec__ = None
sys.modules["sqlalchemy"] = _mock_sqla
_mock_sqla_orm = MagicMock()
_mock_sqla_orm.__path__ = []
_mock_sqla_orm.__spec__ = None
_mock_sqla_orm.Session = MagicMock
sys.modules["sqlalchemy.orm"] = _mock_sqla_orm
sys.modules["sqlalchemy.engine"] = _mock_module()
sys.modules["sqlalchemy.ext"] = _mock_module()
sys.modules["sqlalchemy.ext.declarative"] = _mock_module()
sys.modules["worker_app.db"] = _mock_module(SessionLocal=MagicMock())
sys.modules["worker_app.celery_app"] = _mock_module(celery_app=MagicMock())
sys.modules["worker_app.core.config"] = _mock_module(get_settings=MagicMock(return_value=MagicMock()))
sys.modules["packages.adapters.sqlalchemy_impl.session"] = _mock_module(
Base=MagicMock(),
build_engine=MagicMock(),
build_session_factory=MagicMock(),
ensure_database_exists=MagicMock(),
initialize_database=MagicMock(),
)
sys.modules["packages.adapters.sqlalchemy_impl.generated_video_repository"] = _mock_module(
SQLAlchemyGeneratedVideoRepository=MagicMock
)
sys.modules["packages.adapters.sqlalchemy_impl.models"] = _mock_module(
VideoFingerprintChunkModel=MagicMock,
GeneratedVideoModel=MagicMock,
)
sys.modules["packages.shared.config"] = _mock_module(get_shared_settings=MagicMock(return_value=MagicMock()))
sys.modules["packages.shared.storage"] = _mock_module()
# Save a reference to the dedup module for use in tests (after sys.modules restore)
import video_processing.dedup as _dedup_mod
from video_processing.dedup import ( # noqa: E402
DUPLICATE_THRESHOLD,
HISTOGRAM_WEIGHT,
LONG_VIDEO_DURATION_THRESHOLD_SEC,
MATCH_RATIO_THRESHOLD,
MAX_GAP,
MAX_KEYFRAMES,
MIN_CONSECUTIVE_MATCHES,
MIN_KEYFRAME_INTERVAL_SEC,
MIN_KEYFRAMES,
PHASH_WEIGHT,
SCENE_CHANGE_THRESHOLD,
SEGMENT_MATCH_THRESHOLD,
DuplicateSegment,
FingerprintChunk,
VideoDeduplicator,
VideoFingerprint,
detect_keyframe_timestamps,
find_duplicate_segments,
hamming_distance,
)
# ── Restore sys.modules immediately after import ──
for _key in list(sys.modules.keys()):
if _key not in _SAVED_MODULES_KEYS:
del sys.modules[_key]
for _key, _value in _SAVED_MODULES_VALUES.items():
if _value is not None:
sys.modules[_key] = _value
elif _key in sys.modules:
del sys.modules[_key]
del _SAVED_MODULES_KEYS, _SAVED_MODULES_VALUES, _key, _value
# ── Helper ──────────────────────────────────────────────────────
def _make_chunk(start_ms: int, end_ms: int, phash: str, hist: list[float] | None = None) -> FingerprintChunk:
"""创建测试用 FingerprintChunk."""
return FingerprintChunk(
start_time_ms=start_ms,
end_time_ms=end_ms,
phash_binary=phash,
color_histogram=hist or [0.1] * 96,
frame_count=1,
)
# ── TestDuplicateSegment ────────────────────────────────────────
class TestDuplicateSegment:
"""DuplicateSegment 数据类测试."""
def test_creation(self):
"""正常创建."""
seg = DuplicateSegment(
query_start_ms=1000,
query_end_ms=5000,
target_start_ms=2000,
target_end_ms=6000,
avg_distance=3.5,
)
assert seg.query_start_ms == 1000
assert seg.avg_distance == 3.5
def test_fields(self):
"""所有字段可访问."""
seg = DuplicateSegment(0, 1000, 500, 1500, 2.0)
assert seg.query_end_ms == 1000
assert seg.target_start_ms == 500
assert seg.target_end_ms == 1500
# ── TestDetectKeyframeTimestamps ────────────────────────────────
class TestDetectKeyframeTimestamps:
"""detect_keyframe_timestamps 关键帧检测测试.
由于 cv2 在单元测试环境中是 mock,这里只测试边界条件。
完整的视频处理测试在集成测试中进行。
"""
def test_cannot_open_video_raises(self):
"""无法打开视频时抛出 RuntimeError."""
cv2_mock = _dedup_mod.cv2
mock_cap = MagicMock()
mock_cap.isOpened.return_value = False
cv2_mock.VideoCapture.return_value = mock_cap
import pytest
with pytest.raises(RuntimeError, match="Cannot open video"):
detect_keyframe_timestamps("/fake/path.mp4")
def test_zero_duration_returns_empty(self):
"""视频时长为 0 时返回空列表."""
cv2_mock = _dedup_mod.cv2
mock_cap = MagicMock()
mock_cap.isOpened.return_value = True
# cv2.CAP_PROP_FPS etc. are Mock objects; configure get() to return 0 for frame_count
mock_cap.get.return_value = 0
mock_cap.read.return_value = (False, None)
cv2_mock.VideoCapture.return_value = mock_cap
result = detect_keyframe_timestamps("/fake/zero.mp4")
assert result == []
def test_function_signature(self):
"""验证函数签名和默认参数."""
import inspect
sig = inspect.signature(detect_keyframe_timestamps)
params = sig.parameters
assert "video_path" in params
assert "min_interval_sec" in params
assert "max_frames" in params
assert "min_frames" in params
# 默认值
assert params["min_interval_sec"].default == 1.0
assert params["max_frames"].default == 30
assert params["min_frames"].default == 5
# ── TestFindDuplicateSegments ───────────────────────────────────
class TestFindDuplicateSegments:
"""find_duplicate_segments 滑动窗口时序匹配测试."""
def test_identical_chunks_full_match(self):
"""两组完全相同的 chunks → 整段匹配."""
chunks_a = [_make_chunk(i * 1000, (i + 1) * 1000, "aaaaaaaaaaaaaaaa") for i in range(10)]
chunks_b = [_make_chunk(i * 1000, (i + 1) * 1000, "aaaaaaaaaaaaaaaa") for i in range(10)]
segments = find_duplicate_segments(chunks_a, chunks_b)
assert len(segments) >= 1
# 应该覆盖大部分范围
total_query_range = segments[-1].query_end_ms - segments[0].query_start_ms
assert total_query_range > 5000 # 至少覆盖 5 秒
def test_completely_different_chunks(self):
"""两组完全不同的 chunks → 空列表."""
# 距离都 > 阈值
chunks_a = [_make_chunk(i * 1000, (i + 1) * 1000, "0000000000000000") for i in range(10)]
chunks_b = [_make_chunk(i * 1000, (i + 1) * 1000, "ffffffffffffffff") for i in range(10)]
segments = find_duplicate_segments(chunks_a, chunks_b)
assert segments == []
def test_partial_overlap(self):
"""部分重叠 → 只返回重叠段."""
# 前 5 帧相同,后 5 帧不同
same_hash = "aaaaaaaaaaaaaaaa"
diff_hash_a = "0000000000000000"
diff_hash_b = "ffffffffffffffff"
chunks_a = [_make_chunk(i * 1000, (i + 1) * 1000, same_hash) for i in range(5)] + [
_make_chunk(i * 1000, (i + 1) * 1000, diff_hash_a) for i in range(5, 10)
]
chunks_b = [_make_chunk(i * 1000, (i + 1) * 1000, same_hash) for i in range(5)] + [
_make_chunk(i * 1000, (i + 1) * 1000, diff_hash_b) for i in range(5, 10)
]
segments = find_duplicate_segments(chunks_a, chunks_b)
# 应该只有前 5 帧的匹配段
if segments:
assert segments[0].query_end_ms <= 5000
def test_min_consecutive_not_met(self):
"""连续 4 帧匹配(< min_consecutive=5)→ 不报重复.
注意:使用不同的 hash 对,确保后半部分帧距离 > 阈值。
"""
same_hash = "aaaaaaaaaaaaaaaa"
# 4 帧匹配,后面 6 帧各自不同(在 query 和 target 中使用不同 hash
chunks_a = [_make_chunk(i * 1000, (i + 1) * 1000, same_hash) for i in range(4)] + [
_make_chunk(i * 1000, (i + 1) * 1000, "bbbbbbbbbbbbbbbb") for i in range(4, 10)
]
chunks_b = [_make_chunk(i * 1000, (i + 1) * 1000, same_hash) for i in range(4)] + [
_make_chunk(i * 1000, (i + 1) * 1000, "cccccccccccccccc") for i in range(4, 10)
]
# hamming("bbbb...", "cccc...") should be > 8 (SEGMENT_MATCH_THRESHOLD)
# b=1011, c=1100 → 4 bits differ per hex digit × 16 digits = 64 bits total? No...
# Actually: hamming_distance("bbbbbbbbbbbbbbbb", "cccccccccccccccc")
# b=0xb=1011, c=0xc=1100 → XOR=0111=0x7 → 3 bits per digit × 16 = 48
# That's > 8 so won't match
segments = find_duplicate_segments(chunks_a, chunks_b)
# 只有 4 帧匹配(< min_consecutive=5),所以不报告
assert segments == []
def test_max_gap_behavior(self):
"""5 帧匹配 + 1 帧间隙 + 3 帧匹配 → 验证 max_gap 行为.
关键:间隙帧必须在 query 和 target 中使用不同 hash,使其真正不匹配。
"""
match_hash = "aaaaaaaaaaaaaaaa"
gap_hash_a = "bbbbbbbbbbbbbbbb" # query 端
gap_hash_b = "cccccccccccccccc" # target 端(与 query 端距离 > 8
tail_hash_a = "dddddddddddddddd"
tail_hash_b = "eeeeeeeeeeeeeeee"
# 5 帧匹配, 1 帧间隙, 3 帧匹配, 5 帧不匹配
hashes_a = [match_hash] * 5 + [gap_hash_a] + [match_hash] * 3 + [tail_hash_a] * 5
hashes_b = [match_hash] * 5 + [gap_hash_b] + [match_hash] * 3 + [tail_hash_b] * 5
chunks_a = [_make_chunk(i * 1000, (i + 1) * 1000, h) for i, h in enumerate(hashes_a)]
chunks_b = [_make_chunk(i * 1000, (i + 1) * 1000, h) for i, h in enumerate(hashes_b)]
# max_gap=2, 所以 1 帧间隙会被合并
segments = find_duplicate_segments(chunks_a, chunks_b, max_gap=2)
# 5 match + 1 gap + 3 match = run of 9(间隙被桥接)
assert len(segments) == 1
# run 覆盖 indices 0-85 match + 1 gap + 3 match),但 gap 帧不计入 match
# query_start = chunks_a[0].start = 0
# query_end = chunks_a[8].end = 9000
assert segments[0].query_start_ms == 0
assert segments[0].query_end_ms == 9000
def test_max_gap_exceeded(self):
"""间隙超过 max_gap → 分成两段."""
match_hash = "aaaaaaaaaaaaaaaa"
gap_hash_a = "bbbbbbbbbbbbbbbb"
gap_hash_b = "cccccccccccccccc"
tail_hash_a = "dddddddddddddddd"
tail_hash_b = "eeeeeeeeeeeeeeee"
# 5 帧匹配, 3 帧间隙 (> max_gap=2), 5 帧匹配, 5 帧不匹配
hashes_a = [match_hash] * 5 + [gap_hash_a] * 3 + [match_hash] * 5 + [tail_hash_a] * 5
hashes_b = [match_hash] * 5 + [gap_hash_b] * 3 + [match_hash] * 5 + [tail_hash_b] * 5
chunks_a = [_make_chunk(i * 1000, (i + 1) * 1000, h) for i, h in enumerate(hashes_a)]
chunks_b = [_make_chunk(i * 1000, (i + 1) * 1000, h) for i, h in enumerate(hashes_b)]
segments = find_duplicate_segments(chunks_a, chunks_b, max_gap=2)
# 3 帧间隙 > max_gap=2 → 分成两段(每段 5 帧匹配)
assert len(segments) == 2
def test_empty_chunks(self):
"""空 chunks 返回空列表."""
assert find_duplicate_segments([], [_make_chunk(0, 1000, "aa")]) == []
assert find_duplicate_segments([_make_chunk(0, 1000, "aa")], []) == []
assert find_duplicate_segments([], []) == []
def test_dict_chunks_compatibility(self):
"""dict 格式的 chunks 也能正常工作."""
chunks_a = [
{"phash_binary": "aaaaaaaaaaaaaaaa", "start_time_ms": i * 1000, "end_time_ms": (i + 1) * 1000}
for i in range(10)
]
chunks_b = [
{"phash_binary": "aaaaaaaaaaaaaaaa", "start_time_ms": i * 1000, "end_time_ms": (i + 1) * 1000}
for i in range(10)
]
segments = find_duplicate_segments(chunks_a, chunks_b)
assert len(segments) >= 1
def test_segment_time_ranges(self):
"""返回的 segment 时间范围正确.
每个 query chunk 匹配到 target 中对应的 chunk(相同 hash),
确保 target 时间范围正确映射。
"""
# 给每个 chunk 唯一的 hash(但保证 query[i] == target[i]
def _unique_hash(i: int) -> str:
return format(i, "016x")
chunks_a = [_make_chunk(i * 2000, (i + 1) * 2000, _unique_hash(i)) for i in range(7)]
chunks_b = [_make_chunk(i * 2000, (i + 1) * 2000, _unique_hash(i)) for i in range(7)]
segments = find_duplicate_segments(chunks_a, chunks_b)
assert len(segments) >= 1
seg = segments[0]
assert seg.query_start_ms == 0
assert seg.query_end_ms == 14000
# target 应该映射到正确的范围
assert seg.target_start_ms == 0
assert seg.target_end_ms == 14000
assert seg.avg_distance == 0.0 # 完全相同
# ── TestMedianVsMean ────────────────────────────────────────────
class TestMedianVsMean:
"""中位数 vs 均值:验证中位数抵抗异常值."""
def test_median_resists_outlier(self):
"""距离 [3,3,3,3,30]:均值=8.4,中位数=3.
中位数 < PHASH_THRESHOLD(10),均值也 < 10。
但更极端的:[3,3,3,3,60]:均值=14.4,中位数=3.
"""
import statistics
distances = [3, 3, 3, 3, 60]
assert statistics.median(distances) == 3
assert sum(distances) / len(distances) == 14.4
# 中位数 < 10 → 通过阈值
assert statistics.median(distances) < 10
# ── TestMatchRatioCondition ─────────────────────────────────────
class TestMatchRatioCondition:
"""帧匹配比例条件测试."""
def test_ratio_below_threshold_skips(self):
"""10 帧中只有 5 帧距离 < 10 → match_ratio=0.5 < 0.7 → 跳过."""
distances = [3, 5, 7, 8, 9, 15, 20, 25, 30, 40]
threshold = 10
matching = sum(1 for d in distances if d < threshold)
ratio = matching / len(distances)
assert ratio == 0.5
assert ratio < 0.7 # 应该被跳过
def test_ratio_above_threshold_passes(self):
"""10 帧中 8 帧距离 < 10 → match_ratio=0.8 >= 0.7 → 通过."""
distances = [3, 5, 7, 8, 9, 3, 5, 7, 20, 30]
threshold = 10
matching = sum(1 for d in distances if d < threshold)
ratio = matching / len(distances)
assert ratio == 0.8
assert ratio >= 0.7 # 应该通过
# ── TestBhattacharyyaFusion ─────────────────────────────────────
class TestBhattacharyyaFusion:
"""直方图融合逻辑测试."""
def test_high_phash_high_hist_is_duplicate(self):
"""pHash 高相似 + 直方图高相似 → combined_score 高."""
phash_similarity = 0.95 # median_distance ≈ 3
hist_similarity = 0.90
combined = 0.7 * phash_similarity + 0.3 * hist_similarity
assert combined > 0.70 # DUPLICATE_THRESHOLD
def test_high_phash_low_hist_maybe_not(self):
"""pHash 高相似 + 直方图低相似 → combined_score 取决于权重."""
phash_similarity = 0.85 # median_distance ≈ 10
hist_similarity = 0.10
combined = 0.7 * phash_similarity + 0.3 * hist_similarity
# 0.7 * 0.85 + 0.3 * 0.10 = 0.595 + 0.03 = 0.625 < 0.70
assert combined < 0.70
def test_no_histogram_fallback(self):
"""无直方图数据时 hist_similarity 回退到 0.5."""
phash_similarity = 0.90
hist_similarity = 0.5 # fallback
combined = 0.7 * phash_similarity + 0.3 * hist_similarity
# 0.7 * 0.90 + 0.3 * 0.5 = 0.63 + 0.15 = 0.78 > 0.70
assert combined > 0.70
# ── TestBackwardCompatibility ───────────────────────────────────
class TestBackwardCompatibility:
"""向后兼容测试."""
def test_no_chunks_no_crash(self):
"""已有视频无分片数据 → find_duplicate_segments 返回空列表."""
# 模拟:fingerprint 有 chunks,但 existing 只有 JSON phashes
query_chunks = [_make_chunk(i * 1000, (i + 1) * 1000, "aaaaaaaaaaaaaaaa") for i in range(10)]
# 没有 start_time_ms/end_time_ms 的简化 dict
target_as_dicts = [{"phash_binary": "aaaaaaaaaaaaaaaa"} for _ in range(10)]
# find_duplicate_segments 需要 start_time_ms/end_time_ms
# 在没有的情况下应该不崩溃(用默认值)
# 实际上我们的实现用 _get_start/_get_end 访问,缺 key 会 KeyError
# 所以 check_duplicate 传入时会补上默认值
target_with_defaults = [
{"phash_binary": "aaaaaaaaaaaaaaaa", "start_time_ms": 0, "end_time_ms": 0} for _ in range(10)
]
segments = find_duplicate_segments(query_chunks, target_with_defaults)
# 不会崩溃
assert isinstance(segments, list)
def test_few_chunks_no_crash(self):
"""少量 chunk 不崩溃."""
chunks_a = [_make_chunk(0, 5000, "aaaaaaaaaaaaaaaa")]
chunks_b = [{"phash_binary": "aaaaaaaaaaaaaaaa", "start_time_ms": 0, "end_time_ms": 5000}]
segments = find_duplicate_segments(chunks_a, chunks_b)
# 1 帧 < min_consecutive=5,不会报重复
assert segments == []
# ── TestConstants ───────────────────────────────────────────────
class TestConstants:
"""常量值验证 — 使用已在模块顶部导入的常量,避免重新 import."""
def test_segment_match_threshold(self):
# 从已导入的 find_duplicate_segments 默认参数间接验证
assert SEGMENT_MATCH_THRESHOLD == 8
def test_min_consecutive_matches(self):
assert MIN_CONSECUTIVE_MATCHES == 5
def test_max_gap(self):
assert MAX_GAP == 2
def test_scene_change_threshold(self):
assert SCENE_CHANGE_THRESHOLD == 30
def test_min_keyframe_interval(self):
assert MIN_KEYFRAME_INTERVAL_SEC == 1.0
def test_max_keyframes(self):
assert MAX_KEYFRAMES == 30
def test_min_keyframes(self):
assert MIN_KEYFRAMES == 5
def test_long_video_threshold(self):
assert LONG_VIDEO_DURATION_THRESHOLD_SEC == 180
def test_duplicate_threshold(self):
assert DUPLICATE_THRESHOLD == 0.70
def test_phash_weight(self):
assert PHASH_WEIGHT == 0.7
def test_histogram_weight(self):
assert HISTOGRAM_WEIGHT == 0.3
def test_match_ratio_threshold(self):
assert MATCH_RATIO_THRESHOLD == 0.7
+2 -2
View File
@@ -122,7 +122,7 @@ class TestComputeDuplicateRate:
rate = deduplicator.compute_duplicate_rate(fingerprint, "proj1", "vid1", session)
# hamming distance = 2, similarity = (1 - 2/64) * 100 = 96.875
assert rate == pytest.approx(96.88, abs=0.1)
assert rate == pytest.approx(82.81, abs=0.1) # 新算法: 0.7*(1-2/64)*100 + 0.3*50
def test_excludes_self_video(self):
from video_processing.dedup import VideoDeduplicator
@@ -186,7 +186,7 @@ class TestComputeDuplicateRate:
rate = deduplicator.compute_duplicate_rate(fingerprint, "proj1", "vid1", session)
# max similarity: e2 distance=1, (1-1/64)*100 = 98.4375
assert rate == pytest.approx(98.44, abs=0.1)
assert rate == pytest.approx(83.91, abs=0.1) # 新算法: 0.7*(1-1/64)*100 + 0.3*50
def test_user_id_scope_cross_project(self):
"""传 user_id 时应跨项目查询,而非仅当前项目."""
-31
View File
@@ -110,7 +110,6 @@ from video_processing.dedup import ( # noqa: E402
FingerprintChunk,
VideoFingerprint,
_save_fingerprint_chunks,
compute_chunk_interval,
)
# ── Restore sys.modules immediately after import ──
@@ -125,36 +124,6 @@ for _key, _value in _SAVED_MODULES_VALUES.items():
del _SAVED_MODULES_KEYS, _SAVED_MODULES_VALUES, _key, _value
class TestChunkInterval:
"""测试分片间隔策略。"""
def test_short_video_interval(self):
"""短视频(≤60秒)每 2 秒一个分片。"""
assert compute_chunk_interval(0) == 2
assert compute_chunk_interval(30) == 2
assert compute_chunk_interval(60) == 2
def test_long_video_interval(self):
"""长视频(>60秒)每 5 秒一个分片。"""
assert compute_chunk_interval(61) == 5
assert compute_chunk_interval(120) == 5
assert compute_chunk_interval(300) == 5
def test_chunk_count_60s_video(self):
"""60秒视频 → 30 片(60/2=30)。"""
duration = 60
interval = compute_chunk_interval(duration)
expected_chunks = int(duration / interval)
assert expected_chunks == 30
def test_chunk_count_120s_video(self):
"""120秒视频 → 24 片(120/5=24)。"""
duration = 120
interval = compute_chunk_interval(duration)
expected_chunks = int(duration / interval)
assert expected_chunks == 24
class TestVideoFingerprintToChunkModels:
"""测试 VideoFingerprint.to_chunk_models() 输出。"""