import torch, soundfile as sffrom snac import SNACfrom transformers import AutoModelForCausalLM, AutoTokenizer REPO = "MosesJoshuaCoker/Novax_Krio_V3"device = "cuda" if torch.cuda.is_available() else "cpu" tokenizer = AutoTokenizer.from_pretrained(REPO)model = AutoModelForCausalLM.from_pretrained(REPO, torch_dtype=torch.float16).to(device).eval()snac_model = SNAC.from_pretrained("hubertsiuzdak/snac_24khz").eval().to(device) # Orpheus special tokens (fixed) + the voice name this model was trained with.SOH, EOT, EOH, SOA, SOS = 128259, 128009, 128260, 128261, 128257EOS_SPEECH, PAD_ID, AUDIO_BASE = 128258, 128263, 128266VOICE = "krio" @torch.inference_mode()def decode_frames(aud): """SNAC audio-token ids (already offset by AUDIO_BASE) -> 24 kHz waveform.""" l1, l2, l3 = [], [], [] for i in range(0, len(aud) - len(aud) % 7, 7): f = aud[i:i + 7] l1.append(f[0]) l2 += [f[1] - 4096, f[4] - 4 * 4096] l3 += [f[2] - 2 * 4096, f[3] - 3 * 4096, f[5] - 5 * 4096, f[6] - 6 * 4096] codes = [torch.tensor(c, device=device).unsqueeze(0).clamp(0, 4095) for c in (l1, l2, l3)] return snac_model.decode(codes)[0, 0].float().cpu().numpy() @torch.inference_mode()def say(text, max_seconds=20, temperature=0.6, top_p=0.95, repetition_penalty=1.3): prompt = torch.tensor([[SOH] + tokenizer.encode(f"{VOICE}: {text}") + [EOT, EOH, SOA, SOS]], device=device) out = model.generate( input_ids=prompt, attention_mask=torch.ones_like(prompt), max_new_tokens=int(max_seconds) * 84, # ~84 audio tokens per second do_sample=True, temperature=temperature, top_p=top_p, repetition_penalty=repetition_penalty, # keep >= 1.1 or Orpheus loops eos_token_id=EOS_SPEECH, pad_token_id=PAD_ID, )[0].tolist() audio_ids = [t - AUDIO_BASE for t in out[prompt.shape[1]:] if t >= AUDIO_BASE] return decode_frames(audio_ids) wav = say("Kushe! Aw di bodi?")sf.write("krio.wav", wav, 24000)print("wrote krio.wav")