1import numpy as np, soundfile as sf, sys, torch, torchaudio
2from safetensors.torch import load_file
3
4REPO = "/path/to/openechotts"
5sys.path.insert(0, f"{REPO}/training_torchtitan")
6sys.path.insert(0, f"{REPO}/training_torchtitan/eval")
7sys.path.insert(0, f"{REPO}/ablation_textenc")
8
9from model import BlockDiTT5, BlockDiTT5Config
10from custom_encoder import CustomTextEncoder
11from configs import CustomEncoderConfig
12from eval_checkpoint import (
13 load_dacvae, decode_latent, euler_sample, pad_tokens, pad_latent,
14)
15
16device = torch.device("cuda")
17
18# Build + load DiT
19import json
20cfg = json.load(open("ckpt_hf/config.json"))
21model_cfg = BlockDiTT5Config(**cfg["model_config"])
22model = BlockDiTT5(model_cfg).to(device).eval()
23model.load_state_dict({k: v.float() for k, v in load_file("ckpt_hf/model.safetensors").items()})
24
25# Build + load byte text encoder
26enc_cfg = CustomEncoderConfig(vocab_size=256, dim=768, intermediate_size=2048,
27 n_layers=8, n_heads=6, norm_eps=1e-5)
28encoder = CustomTextEncoder(enc_cfg, attn_type="standard", ffn_type="standard").to(device).eval()
29encoder.load_state_dict(load_file("ckpt_hf/encoder.safetensors"))
30
31# Encode reference audio via DACVAE (20 kHz mono any length; trimmed to 512 frames / 20.5s)
32dacvae = load_dacvae(device)
33wav, sr = torchaudio.load("reference.wav")
34if sr != 48000:
35 wav = torchaudio.functional.resample(wav, sr, 48000)
36if wav.shape[0] > 1:
37 wav = wav.mean(0, keepdim=True)
38with torch.no_grad():
39 z = dacvae.encode(wav.unsqueeze(0).to(device))
40 if isinstance(z, tuple): z = z[0]
41 ref = z.squeeze(0).transpose(0, 1).cpu().numpy()
42ref = ref[: (min(ref.shape[0], 512) // 4) * 4]
43
44spk_lat, spk_mask = pad_latent(ref, 512)
45spk_lat_t = torch.from_numpy(spk_lat).unsqueeze(0).to(device)
46spk_mask_t = torch.from_numpy(spk_mask).unsqueeze(0).to(device)
47
48# Text → bytes → embed
49text = "Hello world, this is a zero-shot voice clone demo."
50ids, tmask = pad_tokens(list(text.encode("utf-8")), 512)
51ids_t = torch.from_numpy(ids).unsqueeze(0).to(device)
52tmask_t = torch.from_numpy(tmask).unsqueeze(0).to(device)
53with torch.amp.autocast("cuda", dtype=torch.bfloat16):
54 text_emb = encoder(ids_t, tmask_t)
55
56# Sample + decode
57latent = euler_sample(model, text_emb, tmask_t, spk_lat_t, spk_mask_t,
58 num_steps=30, cfg_scale=3.0, output_length=500,
59 seed=42, device=str(device))
60audio = decode_latent(dacvae, latent, device)
61sf.write("out.wav", audio, 48000)