moss / MOSS-Transcribe-Diarize /aishell4_eval.py
czyhust's picture
Initial upload
969573c verified
Raw
History Blame Contribute Delete
7.07 kB
#!/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<mnt} cover={covered}", flush=True)
print(f"[GPU{gpu_id}] DONE", flush=True)
# ─── main ────────────────────────────────────────────────────────────────────
def main():
parser = argparse.ArgumentParser()
parser.add_argument("--ckpt", required=True)
parser.add_argument("--output_dir", required=True)
parser.add_argument("--gpus", default="0")
parser.add_argument("--workers_per_gpu", type=int, default=1)
parser.add_argument("--max_new_tokens", type=int, default=0)
args = parser.parse_args()
gpu_list = [g.strip() for g in args.gpus.split(",")]
n_workers = len(gpu_list) * args.workers_per_gpu
all_files = sorted([f.stem for f in Path(AUDIO_DIR).glob("*.wav")])
chunk_size = (len(all_files) + n_workers - 1) // n_workers
procs = []
worker_idx = 0
for gpu in gpu_list:
for _ in range(args.workers_per_gpu):
chunk = all_files[worker_idx * chunk_size : (worker_idx + 1) * chunk_size]
if not chunk:
break
p = mp.Process(target=run_worker, args=(gpu, args.ckpt, chunk, args.output_dir, args.max_new_tokens))
p.start()
procs.append(p)
worker_idx += 1
for p in procs:
p.join()
print("ALL DONE")
if __name__ == "__main__":
mp.set_start_method("spawn", force=True)
main()