Views
No views yet

1sudo apt-get update && sudo apt-get install cbm ffmpeg git-lfs
2
3pip install unsloth
4pip install --no-deps bitsandbytes accelerate xformers==0.0.29.post3 peft trl==0.15.2 triton cut_cross_entropy unsloth_zoo
5pip install sentencepiece protobuf 'datasets>=3.4.1' huggingface_hub hf_transfer
6pip install --no-deps unsloth
7git clone https://github.com/SparkAudio/Spark-TTS
8pip install omegaconf einx
9
10pip uninstall torch torchaudio torchvision -y
11pip install torch torchaudio torchvision
12pip install tf-keras
13pip install soundfile soxr einops librosa
14
15git clone https://huggingface.co/svjack/Spark-TTS-0.5B-Wang-Leehom-Merged-Early
16git clone https://huggingface.co/unsloth/Spark-TTS-0.5B1import sys
2sys.path.append('Spark-TTS')
3
4import torch
5import re
6import numpy as np
7import soundfile as sf
8from IPython.display import Audio, display
9from unsloth import FastModel
10from transformers import AutoTokenizer
11from sparktts.models.audio_tokenizer import BiCodecTokenizer
12
13class SparkTTSLoRAInference:
14 def __init__(self, model_name="lora_model_merged_300/"):
15 """初始化模型和tokenizer"""
16 # 加载基础模型和LoRA适配器
17 self.model, self.tokenizer = FastModel.from_pretrained(
18 model_name=model_name,
19 max_seq_length=2048,
20 dtype=torch.float32,
21 load_in_4bit=False,
22 )
23 #self.model.load_adapter(lora_path) # 加载LoRA权重
24
25 # 初始化音频tokenizer
26 self.audio_tokenizer = BiCodecTokenizer("Spark-TTS-0.5B", "cuda")
27 FastModel.for_inference(self.model) # 启用优化推理模式
28
29 # 打印设备信息
30 print(f"Model loaded on device: {next(self.model.parameters()).device}")
31
32 def generate_speech_from_text(
33 self,
34 text: str,
35 temperature: float = 0.8,
36 top_k: int = 50,
37 top_p: float = 1,
38 max_new_audio_tokens: int = 2048,
39 device: torch.device = torch.device("cuda" if torch.cuda.is_available() else "cpu")
40 ) -> np.ndarray:
41 """
42 Generates speech audio from text using default voice control parameters.
43 Args:
44 text (str): The text input to be converted to speech.
45 temperature (float): Sampling temperature for generation.
46 top_k (int): Top-k sampling parameter.
47 top_p (float): Top-p (nucleus) sampling parameter.
48 max_new_audio_tokens (int): Max number of new tokens to generate (limits audio length).
49 device (torch.device): Device to run inference on.
50 Returns:
51 np.ndarray: Generated waveform as a NumPy array.
52 """
53 FastModel.for_inference(self.model) # Enable native 2x faster inference
54 prompt = "".join([
55 "<|task_tts|>",
56 "<|start_content|>",
57 text,
58 "<|end_content|>",
59 "<|start_global_token|>"
60 ])
61 model_inputs = self.tokenizer([prompt], return_tensors="pt").to(device)
62 print("Generating token sequence...")
63 generated_ids = self.model.generate(
64 **model_inputs,
65 max_new_tokens=max_new_audio_tokens, # Limit generation length
66 do_sample=True,
67 temperature=temperature,
68 top_k=top_k,
69 top_p=top_p,
70 eos_token_id=self.tokenizer.eos_token_id, # Stop token
71 pad_token_id=self.tokenizer.pad_token_id # Use models pad token id
72 )
73 print("Token sequence generated.")
74 generated_ids_trimmed = generated_ids[:, model_inputs.input_ids.shape[1]:]
75 predicts_text = self.tokenizer.batch_decode(generated_ids_trimmed, skip_special_tokens=False)[0]
76 # print(f"\nGenerated Text (for parsing):\n{predicts_text}\n") # Debugging
77 # Extract semantic token IDs using regex
78 semantic_matches = re.findall(r"<\|bicodec_semantic_(\d+)\|>", predicts_text)
79 if not semantic_matches:
80 print("Warning: No semantic tokens found in the generated output.")
81 return np.array([], dtype=np.float32)
82 pred_semantic_ids = torch.tensor([int(token) for token in semantic_matches]).long().unsqueeze(0) # Add batch dim
83 # Extract global token IDs using regex
84 global_matches = re.findall(r"<\|bicodec_global_(\d+)\|>", predicts_text)
85 if not global_matches:
86 print("Warning: No global tokens found in the generated output (controllable mode). Might use defaults or fail.")
87 pred_global_ids = torch.zeros((1, 1), dtype=torch.long)
88 else:
89 pred_global_ids = torch.tensor([int(token) for token in global_matches]).long().unsqueeze(0) # Add batch dim
90 pred_global_ids = pred_global_ids.unsqueeze(0) # Shape becomes (1, 1, N_global)
91 print(f"Found {pred_semantic_ids.shape[1]} semantic tokens.")
92 print(f"Found {pred_global_ids.shape[2]} global tokens.")
93 # Detokenize using BiCodecTokenizer
94 print("Detokenizing audio tokens...")
95 # Ensure audio_tokenizer and its internal model are on the correct device
96 self.audio_tokenizer.device = device
97 self.audio_tokenizer.model.to(device)
98 # Squeeze the extra dimension from global tokens as seen in SparkTTS example
99 wav_np = self.audio_tokenizer.detokenize(
100 pred_global_ids.to(device).squeeze(0), # Shape (1, N_global)
101 pred_semantic_ids.to(device) # Shape (1, N_semantic)
102 )
103 print("Detokenization complete.")
104 return wav_np
105
106tts = SparkTTSLoRAInference("Spark-TTS-0.5B-Wang-Leehom-Merged-Early")1generated_waveform = tts.generate_speech_from_text("音乐是灵魂的独白,在寂静中才能听见最真实的旋律。我选择用孤独淬炼创作,因为喧嚣的世界里,唯有孤独能让艺术扎根生长。", max_new_audio_tokens = 2048)
2if generated_waveform.size > 0:
3 output_filename = "infer1.wav"
4 sample_rate = tts.audio_tokenizer.config.get("sample_rate", 16000)
5 sf.write(output_filename, generated_waveform, sample_rate)
6 print(f"Audio saved to {output_filename}")
7 # Optional: Play audio
8 display(Audio(generated_waveform, rate=sample_rate))
1generated_waveform = tts.generate_speech_from_text("华流不是一道墙,而是一座桥。当东方韵律与西方节拍在音符间对话,我们会发现:所谓遥远,不过是心未抵达的距离。", max_new_audio_tokens = 2048)
2if generated_waveform.size > 0:
3 output_filename = "infer2.wav"
4 sample_rate = tts.audio_tokenizer.config.get("sample_rate", 16000)
5 sf.write(output_filename, generated_waveform, sample_rate)
6 print(f"Audio saved to {output_filename}")
7 # Optional: Play audio
8 display(Audio(generated_waveform, rate=sample_rate))
1generated_waveform = tts.generate_speech_from_text("地球的旋律需要所有人合奏。少一次浪费,多一次举手之劳,微光汇聚时,平凡也能成为改变世界的和弦。", max_new_audio_tokens = 2048)
2if generated_waveform.size > 0:
3 output_filename = "infer3.wav"
4 sample_rate = tts.audio_tokenizer.config.get("sample_rate", 16000)
5 sf.write(output_filename, generated_waveform, sample_rate)
6 print(f"Audio saved to {output_filename}")
7 # Optional: Play audio
8 display(Audio(generated_waveform, rate=sample_rate))