Views
No views yet
1"""MossSpeech inference demo aligned with Hugging Face Transformers guidelines."""
2import os
3from dataclasses import astuple
4
5import torch
6import torchaudio
7
8from transformers import (
9 AutoModel,
10 AutoProcessor,
11 GenerationConfig,
12 StoppingCriteria,
13 StoppingCriteriaList,
14)
15
16
17prompt = "Hello!"
18prompt_audio = "<your path to prompt>"
19model_path = "fnlp/MOSS-Speech"
20codec_path = "fnlp/MOSS-Speech-Codec"
21output_path = "outputs"
22output_modality = "audio" # or text
23
24generation_config = GenerationConfig(
25 temperature=0.7,
26 top_p=0.95,
27 top_k=20,
28 repetition_penalty=1.0,
29 max_new_tokens=1000,
30 min_new_tokens=10,
31 do_sample=True,
32 use_cache=True,
33)
34
35
36class StopOnToken(StoppingCriteria):
37 """Stop generation once the final token equals the provided stop ID."""
38
39 def __init__(self, stop_id: int) -> None:
40 super().__init__()
41 self.stop_id = stop_id
42
43 def __call__(self, input_ids: torch.LongTensor, scores) -> bool: # type: ignore[override]
44 return input_ids[0, -1].item() == self.stop_id
45
46
47def prepare_stopping_criteria(processor):
48 tokenizer = processor.tokenizer
49 stop_tokens = [
50 tokenizer.pad_token_id,
51 tokenizer.convert_tokens_to_ids("<|im_end|>"),
52 ]
53 return StoppingCriteriaList([StopOnToken(token_id) for token_id in stop_tokens])
54
55
56messages = [
57 [
58 {
59 "role": "system",
60 "content": "You are a helpful voice assistant. Answer the user's questions with spoken responses."},
61 # "content": "You are a helpful assistant. Answer the user's questions with text."}, # if output_modality = "text"
62 {
63 "role": "user",
64 "content": prompt
65 }
66 ]
67]
68
69
70processor = AutoProcessor.from_pretrained(model_path, codec_path=codec_path, device="cuda", trust_remote_code=True)
71stopping_criteria = prepare_stopping_criteria(processor)
72encoded_inputs = processor(messages, output_modality)
73
74model = AutoModel.from_pretrained(model_path, trust_remote_code=True, device_map="cuda").eval()
75
76with torch.inference_mode():
77 token_ids = model.generate(
78 input_ids=encoded_inputs["input_ids"].to("cuda"),
79 attention_mask=encoded_inputs["attention_mask"].to("cuda"),
80 generation_config=generation_config,
81 stopping_criteria=stopping_criteria,
82 )
83
84results = processor.decode(token_ids, output_modality, decoder_audio_prompt_path=prompt_audio)
85
86os.makedirs(output_path, exist_ok=True)
87for index, (result, modality) in enumerate(zip(results, output_modality)):
88 audio, text, sample_rate = astuple(result)
89 if modality == "audio":
90 torchaudio.save(f"{output_path}/audio_{index}.wav", audio, sample_rate)
91 else:
92 print(text)