#!/usr/bin/env python3 """dubbing_{2b,9b}_fulldir_cv3 断网自包含验证: 完整加载 + 6 方向各 dub 1 条 + whisper 回测。 用法: HF_HUB_OFFLINE=1 TRANSFORMERS_OFFLINE=1 MODELSCOPE_OFFLINE=1 \ python verify_export.py [--model-dir DIR] [--directions zh2en,zh2es,zh2ja,en2zh,en2es,en2ja] 验证点: - 断网完整加载 (HF_HUB_OFFLINE / TRANSFORMERS_OFFLINE / MODELSCOPE_OFFLINE); - 每方向出一条非静音 wav, 时长合理 (1-15s, 超出仅 WARN); - whisper-large-v3 回测每方向输出, 对送入 CV3 的译文 (info tgt_cv/tgt_raw) 算 content 口径误差: en/es=WER, zh=CER (NFKC+去标点), ja=kata CER (pykakasi); 量级应 ~0.0-0.1 (对照 docs/RESULTS_2B.md D 线终版表); - 打印每方向译文与 samples/refs.json 参考 (若有) 供人工核对。 """ import argparse import json import os import string import sys import unicodedata os.environ.setdefault('HF_HUB_OFFLINE', '1') os.environ.setdefault('TRANSFORMERS_OFFLINE', '1') os.environ.setdefault('MODELSCOPE_OFFLINE', '1') _HERE = os.path.dirname(os.path.abspath(__file__)) sys.path.insert(0, _HERE) WHISPER = os.environ.get('DUBBING_WHISPER', '/mnt/pfs-write/users/lusheng/Index-MT/models/whisper-large-v3') _PUNCT = set(string.punctuation) | set(',。!?、;:""''()《》¿¡،؛。?!') _k = None def ja2kata(t): global _k if _k is None: import pykakasi # 随包 vendored (code/ja_ext), 由 pipeline import 时入 sys.path _k = pykakasi.kakasi() return ''.join(x['kana'] if x['kana'] else x['orig'] for x in _k.convert(t)) def edit_dist(r, h): prev = list(range(len(h) + 1)) for i, rc in enumerate(r, 1): cur = [i] for j, hc in enumerate(h, 1): cur.append(min(prev[j] + 1, cur[j - 1] + 1, prev[j - 1] + (rc != hc))) prev = cur return prev[-1] def wer(ref, hyp): norm = lambda s: [w.strip(string.punctuation) for w in s.lower().split() if w.strip(string.punctuation)] r, h = norm(ref), norm(hyp) return edit_dist(r, h) / len(r) if r else float('nan') def cer_zh(ref, hyp): norm = lambda s: [c for c in unicodedata.normalize('NFKC', s) if not c.isspace() and c not in _PUNCT] r, h = norm(ref), norm(hyp) return edit_dist(r, h) / len(r) if r else float('nan') def cer_ja_kata(ref, hyp): return cer_zh(ja2kata(ref), ja2kata(hyp)) def content_err(lang, tgt_text, asr_text): if lang in ('en', 'es'): return 'wer', wer(tgt_text, asr_text) if lang == 'zh': return 'cer', cer_zh(tgt_text, asr_text) return 'cer_kata', cer_ja_kata(tgt_text, asr_text) def main(): ap = argparse.ArgumentParser() ap.add_argument('--model-dir', default=_HERE) ap.add_argument('--directions', default='zh2en,zh2es,zh2ja,en2zh,en2es,en2ja') args = ap.parse_args() directions = args.directions.split(',') assert not os.environ.get('DUBBING_HOME') or \ os.path.abspath(os.environ['DUBBING_HOME']) == os.path.abspath(args.model_dir) from modeling_dubbing import DubbingBridgeModel model = DubbingBridgeModel.from_pretrained(args.model_dir) print('[verify] model loaded, single process', flush=True) refs = json.load(open(os.path.join(args.model_dir, 'samples', 'refs.json'))) src_wav = { 'zh': os.path.join(args.model_dir, 'samples', 'input_zh.wav'), 'en': os.path.join(args.model_dir, 'samples', 'input_en.wav'), } from transformers import pipeline as hf_pipeline asr = hf_pipeline('automatic-speech-recognition', model=WHISPER, device=0, generate_kwargs={'task': 'transcribe'}) print('[verify] whisper loaded', flush=True) report = [] for d in directions: src, lang = d.split('2') out = os.path.join(args.model_dir, 'samples', f'dub_{d}.wav') wav, sr, info = model.dub(src_wav[src], lang=lang, out_wav=out, return_info=True) dur = wav.shape[1] / sr rms = float(wav.float().pow(2).mean().sqrt()) assert rms > 1e-4, f'{d}: silent output (rms={rms})' status = 'OK' if 1.0 <= dur <= 15.0 else 'WARN dur out of 1-15s' t = asr(out, generate_kwargs={'language': lang, 'task': 'transcribe'})['text'] # content 口径: ASR vs 送入 CV3 的文本 (ja 对 tgt_raw 做 kata; zh 对 tgt_cv) tgt_for_metric = info['tgt_raw'] if lang == 'ja' else info['tgt_cv'] mname, mval = content_err(lang, tgt_for_metric, t) ref = (refs.get(f'{src}_src', {}).get('refs') or {}).get(lang) rec = {'direction': d, 'src_lang_detected': info.get('src_lang'), 'dur': round(dur, 2), 'rms': round(rms, 4), 'src_text': info['zh'], 'tgt_raw': info['tgt_raw'], 'tgt_cv': info['tgt_cv'], 'asr': t, 'ref': ref, f'content_{mname}': round(mval, 3), 'status': status} report.append(rec) print(f'[verify:{d}] dur={dur:.2f}s rms={rms:.4f} ' f'content_{mname}={mval:.3f} {status}', flush=True) print(f'[verify:{d}] src({info.get("src_lang")}): {info["zh"]}', flush=True) print(f'[verify:{d}] hyp_raw : {info["tgt_raw"]}', flush=True) print(f'[verify:{d}] cv_text : {info["tgt_cv"]}' + (' (ja kata)' if lang == 'ja' else ''), flush=True) print(f'[verify:{d}] asr : {t.strip()}', flush=True) if ref: print(f'[verify:{d}] ref : {ref}', flush=True) json.dump(report, open(os.path.join(args.model_dir, 'samples', 'verify_report.json'), 'w'), indent=1, ensure_ascii=False) print('VERIFY EXPORT DONE', flush=True) if __name__ == '__main__': main()