#!/usr/bin/env python3 """AISHELL-4 批量推理(多 GPU 并行,兼容 MOSS-Speaker-RoPE 和 MOSS-Transcribe-Diarize) 用法示例: python aishell4_eval.py \ --ckpt /wangshuai/moss/MOSS-Transcribe-Diarize/output_lr1e-4_spkw5/checkpoint-1207 \ --output_dir /wangshuai/moss/MOSS-Transcribe-Diarize/aishell4_eval_output_lr1e4_spkw5 \ --gpus 0 --workers_per_gpu 4 """ import os, sys, time, json, re, argparse, multiprocessing as mp from pathlib import Path module_root = Path("/wangshuai/moss/MOSS_Speaker-RoPE/moss_speaker_rope").parent if str(module_root) not in sys.path: sys.path.insert(0, str(module_root)) AUDIO_DIR = "/F00120240032/Aishell-4/test/test/wav" PROCESSOR_ID = "/wangshuai/moss/MOSS-Transcribe-Diarize/MOSS-Transcribe-Diarize" # ─── model detection ───────────────────────────────────────────────────────── def detect_model_type(ckpt: str) -> str: """Return "speaker_rope" or "moss". """ cj = json.loads((Path(ckpt) / "config.json").read_text()) mt = cj.get("model_type", "") if "speaker_rope" in mt: return "speaker_rope" return "moss" def load_model(ckpt: str, device, dtype, model_type: str): if model_type == "speaker_rope": sys.path.insert(0, "/taoye/lhy/czy/moss/MOSS_Speaker-RoPE") from moss_speaker_rope.configuration_moss_speaker_rope import MossSpeakerRopeConfig from moss_speaker_rope.modeling_moss_speaker_rope import MossSpeakerRopeForConditionalGeneration cj = json.loads((Path(ckpt) / "config.json").read_text()) for k in ("architectures", "auto_map", "model_type", "dtype", "transformers_version"): cj.pop(k, None) cfg = MossSpeakerRopeConfig(**cj) cfg._attn_implementation = "sdpa" cfg.text_config._attn_implementation = "sdpa" model = MossSpeakerRopeForConditionalGeneration.from_pretrained( ckpt, config=cfg, trust_remote_code=True, dtype=dtype).to(device).eval() model.model.speaker_encoder.float() return model, True # has_speaker=True else: from transformers import AutoModelForCausalLM model = AutoModelForCausalLM.from_pretrained( ckpt, trust_remote_code=True, dtype="auto").to(dtype=dtype).to(device).eval() return model, False # has_speaker=False # ─── worker ────────────────────────────────────────────────────────────────── def run_worker(gpu_id, ckpt, file_list, output_dir, max_new_tokens): os.environ["CUDA_VISIBLE_DEVICES"] = str(gpu_id) import torch, soundfile as sf, soxr from moss_speaker_rope.inference_utils import build_transcription_messages device = torch.device("cuda:0"); dtype = torch.bfloat16 model_type = detect_model_type(ckpt) model, has_speaker = load_model(ckpt, device, dtype, model_type) # Processor: use the one that matches the model type if model_type == "speaker_rope": from moss_speaker_rope.processing_moss_speaker_rope import MossSpeakerRopeProcessor processor = MossSpeakerRopeProcessor.from_pretrained(PROCESSOR_ID, trust_remote_code=True) else: from moss_transcribe_diarize.processing_moss_transcribe_diarize import MossTranscribeDiarizeProcessor processor = MossTranscribeDiarizeProcessor.from_pretrained(PROCESSOR_ID, trust_remote_code=True) out_dir = Path(output_dir); out_dir.mkdir(parents=True, exist_ok=True) for fname in file_list: fpath = Path(AUDIO_DIR) / f"{fname}.wav" dur = sf.info(str(fpath)).duration audio, sr = sf.read(str(fpath), dtype="float32", always_2d=True) audio = audio.mean(axis=1) sfr = int(processor.feature_extractor.sampling_rate) if sr != sfr: audio = soxr.resample(audio, sr, sfr) msgs = build_transcription_messages(str(fpath)) text = processor.apply_chat_template(msgs, tokenize=False, add_generation_prompt=True) batch = processor(text=text, audio=[audio], max_length=81920, return_tensors="pt") batch = {k: v.to(device) for k, v in batch.items()} prompt_len = batch["attention_mask"].sum().item() mnt = max_new_tokens if max_new_tokens > 0 else min(35000, int(dur * 13)) generate_kwargs = { "input_ids": batch["input_ids"], "attention_mask": batch["attention_mask"], "input_features": batch["input_features"], "audio_feature_lengths": batch["audio_feature_lengths"], "audio_chunk_mapping": batch["audio_chunk_mapping"], "max_new_tokens": mnt, "do_sample": False, "use_cache": True, } if has_speaker: generate_kwargs["speaker_input_values"] = batch["speaker_input_values"] generate_kwargs["speaker_chunk_mapping"] = batch["speaker_chunk_mapping"] t0 = time.time() with torch.inference_mode(): out = model.generate(**generate_kwargs) elapsed = time.time() - t0 n_gen = out.shape[1] - prompt_len txt = processor.tokenizer.decode(out[0][prompt_len:], skip_special_tokens=True) ts = re.findall(r"\[(\d+\.\d+)\]", txt) covered = f"{ts[0]}->{ts[-1]}" if ts else "?" (out_dir / f"{fname}.txt").write_text(txt, encoding="utf-8") print(f"[GPU{gpu_id}] {fname}: {elapsed:.0f}s {n_gen}tok eos={n_gen