(example, rank, audio_dir, delay_token_nums=10)
| 27 | process_globals = {} |
| 28 | |
| 29 | def convert_format(example, rank, audio_dir, delay_token_nums=10): |
| 30 | global process_globals |
| 31 | num_gpus = torch.cuda.device_count() |
| 32 | |
| 33 | # 每个进程只加载一次 ort_session |
| 34 | if "ort_session" not in process_globals: |
| 35 | onnx_path = "../pretrained_models/Fun-CosyVoice3-0.5B-2512/speech_tokenizer_v3.onnx" |
| 36 | # 根据 rank 分配到不同的 GPU,使用取模运算实现轮询分配 |
| 37 | num_gpus = torch.cuda.device_count() |
| 38 | device_id = rank % num_gpus |
| 39 | print(f"Loading audio tokenizer for process rank {rank} on GPU {device_id}") |
| 40 | process_globals["ort_session"] = get_audio_tokenizer(onnx_path, device_id=device_id) |
| 41 | |
| 42 | ort_session = process_globals["ort_session"] |
| 43 | |
| 44 | messages = [ |
| 45 | {"role": "user", "content": AUDIO_TEMPLATE}, |
| 46 | {"role": "assistant", "content": AUDIO_TEMPLATE} |
| 47 | ] |
| 48 | |
| 49 | # Save input audio to local |
| 50 | input_audio_data = example['input_audio'] |
| 51 | input_audio_path = f"{audio_dir}/{input_audio_data['path']}" |
| 52 | sf.write(input_audio_path, input_audio_data['array'], input_audio_data['sampling_rate']) |
| 53 | |
| 54 | # Save output audio to local |
| 55 | output_audio_data = example['output_audio'] |
| 56 | output_audio_path = f"{audio_dir}/{output_audio_data['path']}" |
| 57 | sf.write(output_audio_path, output_audio_data['array'], output_audio_data['sampling_rate']) |
| 58 | |
| 59 | assistant_tokens = extract_speech_token(ort_session, output_audio_path) |
| 60 | |
| 61 | audios = [ |
| 62 | { |
| 63 | "path": os.path.realpath(input_audio_path), |
| 64 | "text": "", |
| 65 | "token": AUDIO_PAD_TOKEN * int(input_audio_data['array'].shape[0] / input_audio_data['sampling_rate'] * TOKEN_FPS), |
| 66 | "ref_path": "", |
| 67 | "ref_text": example['speech_input'], |
| 68 | }, |
| 69 | { |
| 70 | "path": "", |
| 71 | "text": example['output'], |
| 72 | "token": AUDIO_BOS_TOKEN * delay_token_nums + ''.join([f'[AU{token:04d}]' for token in assistant_tokens]), |
| 73 | "ref_path": os.path.realpath(output_audio_path), |
| 74 | "ref_text": "", |
| 75 | } |
| 76 | ] |
| 77 | audios = [json.dumps(audio, ensure_ascii=False, sort_keys=True) for audio in audios] |
| 78 | |
| 79 | return {"system": DEFAULT_S2M_PROMPT, "messages": messages, "audio": audios} |
| 80 | |
| 81 | def main(): |
| 82 | parser = argparse.ArgumentParser(description="Process audio dataset for training") |
nothing calls this directly
no test coverage detected