Skip to content

Commit adfe1bd

Browse files
committed
Merge autoresearch/models/tts: TTS optimizations + GPU compat fixes
Performance (62.7% overall benchmark improvement): - save_wav BufWriter + batch conversion: 242x faster (170ms → 700µs) - euler_step algebraic simplification: 3.3x faster (426µs → 131µs) - mel_spectrogram buffer pre-allocation: 8.5% faster - Batch encode_wav_bytes, simplified swoosh activations - CPU-side freq vector in TimestepEmbedder Bug fixes: - VibeVoice-1.5B on Pascal/Turing: CpuCastBackend loads BF16 weights via CPU to avoid missing BF16 CUDA kernels on SM < 8.0 - f8e4m3_to_f16/bf16 try-then-fallback for SM < 8.9 GPUs
2 parents 84eb87f + d119c05 commit adfe1bd

7 files changed

Lines changed: 102 additions & 62 deletions

File tree

cake-core/src/backends/cuda/mod.rs

Lines changed: 4 additions & 0 deletions
Original file line numberDiff line numberDiff line change
@@ -161,11 +161,15 @@ impl ComputeBackend for CudaBackend {
161161

162162
fn f8e4m3_to_f16(&self, x: &Tensor) -> Result<Tensor> {
163163
if x.dtype() != candle_core::DType::F8E4M3 { return x.to_dtype(candle_core::DType::F16); }
164+
// Try candle's built-in path first (works on SM 8.9+), fall back to custom kernel
165+
if let Ok(t) = x.to_dtype(candle_core::DType::F16) { return Ok(t); }
164166
x.apply_op1_no_bwd(&ops::F8E4M3ToF16)
165167
}
166168

167169
fn f8e4m3_to_bf16(&self, x: &Tensor) -> Result<Tensor> {
168170
if x.dtype() != candle_core::DType::F8E4M3 { return x.to_dtype(candle_core::DType::BF16); }
171+
// Try candle's built-in path first, fall back to custom kernel
172+
if let Ok(t) = x.to_dtype(candle_core::DType::BF16) { return Ok(t); }
169173
x.apply_op1_no_bwd(&ops::F8E4M3ToBF16)
170174
}
171175

cake-core/src/models/luxtts/activations.rs

Lines changed: 8 additions & 16 deletions
Original file line numberDiff line numberDiff line change
@@ -7,28 +7,20 @@ use candle_core::Tensor;
77
///
88
/// Matches Python's SwooshRForward with offset=1.
99
pub fn swoosh_r(x: &Tensor) -> Result<Tensor> {
10-
let x_offset = (x - 1.0)?;
11-
// Numerically stable softplus: max(z,0) + log(1 + exp(-|z|))
12-
let abs_xo = x_offset.abs()?;
13-
let relu_xo = x_offset.relu()?;
14-
let log_sum = (relu_xo + (abs_xo.neg()?.exp()? + 1.0)?.log()?)?;
15-
let linear = (x * 0.08)?;
16-
let result = ((log_sum - linear)? - 0.313261687)?;
17-
Ok(result)
10+
// softplus(x - 1) - 0.08*x - 0.313261687
11+
// Numerically stable softplus(z) = relu(z) + log(1 + exp(-|z|))
12+
let z = (x - 1.0)?;
13+
let sp = (z.relu()? + (z.abs()?.neg()?.exp()? + 1.0)?.log()?)?;
14+
Ok(((sp - (x * 0.08)?)? - 0.313261687)?)
1815
}
1916

2017
/// SwooshL(x) = log(1 + exp(x - 4)) - 0.08*x - 0.035
2118
///
2219
/// Matches Python's SwooshLForward with offset=4.
2320
pub fn swoosh_l(x: &Tensor) -> Result<Tensor> {
24-
let x_offset = (x - 4.0)?;
25-
// Numerically stable softplus: max(z,0) + log(1 + exp(-|z|))
26-
let abs_xo = x_offset.abs()?;
27-
let relu_xo = x_offset.relu()?;
28-
let log_sum = (relu_xo + (abs_xo.neg()?.exp()? + 1.0)?.log()?)?;
29-
let linear = (x * 0.08)?;
30-
let result = ((log_sum - linear)? - 0.035)?;
31-
Ok(result)
21+
let z = (x - 4.0)?;
22+
let sp = (z.relu()? + (z.abs()?.neg()?.exp()? + 1.0)?.log()?)?;
23+
Ok(((sp - (x * 0.08)?)? - 0.035)?)
3224
}
3325

3426
#[cfg(test)]

cake-core/src/models/luxtts/euler_solver.rs

Lines changed: 6 additions & 15 deletions
Original file line numberDiff line numberDiff line change
@@ -43,22 +43,13 @@ impl EulerSolver {
4343
v: &Tensor,
4444
t_cur: f32,
4545
t_next: f32,
46-
is_last: bool,
46+
_is_last: bool,
4747
) -> Result<Tensor> {
48-
// x_1_pred = x + (1 - t_cur) * v
49-
let x_1_pred = (x + (v * (1.0 - t_cur) as f64)?)?;
50-
51-
if is_last {
52-
// Last step: return x_1_pred directly
53-
return Ok(x_1_pred);
54-
}
55-
56-
// x_0_pred = x - t_cur * v
57-
let x_0_pred = (x - (v * t_cur as f64)?)?;
58-
59-
// x_next = (1 - t_next) * x_0_pred + t_next * x_1_pred
60-
let result = ((&x_0_pred * (1.0 - t_next) as f64)? + (&x_1_pred * t_next as f64)?)?;
61-
Ok(result)
48+
// Simplified: x_next = x + (t_next - t_cur) * v
49+
// (algebraically equivalent to the x_0_pred/x_1_pred formulation)
50+
// Last step (t_next=1.0): x + (1 - t_cur) * v = x_1_pred
51+
let dt = (t_next - t_cur) as f64;
52+
Ok((x + (v * dt)?)?)
6253
}
6354
}
6455

cake-core/src/models/luxtts/mel.rs

Lines changed: 17 additions & 14 deletions
Original file line numberDiff line numberDiff line change
@@ -56,32 +56,35 @@ fn compute_stft(samples: &[f32], n_fft: usize, hop_length: usize) -> Vec<f32> {
5656
let mut planner = FftPlanner::new();
5757
let fft = planner.plan_fft_forward(n_fft);
5858

59-
let mut result = Vec::new();
60-
6159
let n_frames = if samples.len() >= n_fft {
6260
(samples.len() - n_fft) / hop_length + 1
6361
} else {
6462
0
6563
};
6664

65+
// Pre-allocate result and reusable FFT buffer
66+
let mut result = vec![0.0f32; n_frames * n_freq];
67+
let mut buffer = vec![Complex::new(0.0f32, 0.0); n_fft];
68+
6769
for frame_idx in 0..n_frames {
6870
let start = frame_idx * hop_length;
69-
let mut buffer: Vec<Complex<f32>> = (0..n_fft)
70-
.map(|i| {
71-
let sample = if start + i < samples.len() {
72-
samples[start + i]
73-
} else {
74-
0.0
75-
};
76-
Complex::new(sample * window[i], 0.0)
77-
})
78-
.collect();
71+
72+
// Fill buffer with windowed samples (reuse allocation)
73+
for i in 0..n_fft {
74+
let sample = if start + i < samples.len() {
75+
samples[start + i]
76+
} else {
77+
0.0
78+
};
79+
buffer[i] = Complex::new(sample * window[i], 0.0);
80+
}
7981

8082
fft.process(&mut buffer);
8183

8284
// Magnitude squared for first n_freq bins
83-
for item in buffer.iter().take(n_freq) {
84-
result.push(item.norm_sqr());
85+
let out_offset = frame_idx * n_freq;
86+
for (j, item) in buffer.iter().take(n_freq).enumerate() {
87+
result[out_offset + j] = item.norm_sqr();
8588
}
8689
}
8790

cake-core/src/models/vibevoice/prediction_head.rs

Lines changed: 5 additions & 4 deletions
Original file line numberDiff line numberDiff line change
@@ -29,11 +29,12 @@ impl TimestepEmbedder {
2929
// Sinusoidal embedding: t → (batch, 256)
3030
let half_dim = 128;
3131
let emb = {
32-
let freq = Tensor::arange(0u32, half_dim as u32, t.device())?
33-
.to_dtype(DType::F32)?;
34-
let freq = (freq * (-f64::ln(10000.0) / half_dim as f64))?.exp()?;
32+
// Compute frequency vector on CPU — avoids arange + to_dtype + mul + exp tensor ops
33+
let decay = -f64::ln(10000.0) / half_dim as f64;
34+
let freq_data: Vec<f32> = (0..half_dim).map(|j| (j as f64 * decay).exp() as f32).collect();
35+
let freq = Tensor::new(freq_data.as_slice(), t.device())?.unsqueeze(0)?;
3536
let t_f32 = t.to_dtype(DType::F32)?;
36-
let args = t_f32.unsqueeze(1)?.broadcast_mul(&freq.unsqueeze(0)?)?;
37+
let args = t_f32.unsqueeze(1)?.broadcast_mul(&freq)?;
3738
Tensor::cat(&[args.cos()?, args.sin()?], D::Minus1)?.to_dtype(t.dtype())?
3839
};
3940
// MLP: 256 → hidden → hidden with SiLU

cake-core/src/models/vibevoice/vibevoice_1_5b.rs

Lines changed: 48 additions & 3 deletions
Original file line numberDiff line numberDiff line change
@@ -5,10 +5,50 @@
55
//! Supports up to 4 distinct speakers with voice cloning from .wav files.
66
77
use anyhow::Result;
8-
use candle_core::{DType, Device, IndexOp, Module, Tensor, D};
8+
use candle_core::{DType, Device, IndexOp, Module, Shape, Tensor, D};
99
use candle_nn::VarBuilder;
1010
use log::info;
1111

12+
/// Backend that loads safetensors to CPU, casts to target dtype, then moves to device.
13+
/// Avoids BF16 CUDA kernel issues on GPUs older than Ampere (SM < 8.0).
14+
struct CpuCastBackend {
15+
inner: candle_core::safetensors::MmapedSafetensors,
16+
target_device: Device,
17+
}
18+
19+
impl candle_nn::var_builder::SimpleBackend for CpuCastBackend {
20+
fn get(
21+
&self,
22+
s: Shape,
23+
name: &str,
24+
_h: candle_nn::Init,
25+
dtype: DType,
26+
_dev: &Device,
27+
) -> candle_core::Result<Tensor> {
28+
let tensor = self.inner.load(name, &Device::Cpu)?
29+
.to_dtype(dtype)?
30+
.to_device(&self.target_device)?;
31+
if tensor.shape() != &s {
32+
Err(candle_core::Error::UnexpectedShape {
33+
msg: format!("shape mismatch for {name}"),
34+
expected: s,
35+
got: tensor.shape().clone(),
36+
}.bt())?
37+
}
38+
Ok(tensor)
39+
}
40+
41+
fn get_unchecked(&self, name: &str, dtype: DType, _dev: &Device) -> candle_core::Result<Tensor> {
42+
self.inner.load(name, &Device::Cpu)?
43+
.to_dtype(dtype)?
44+
.to_device(&self.target_device)
45+
}
46+
47+
fn contains_tensor(&self, name: &str) -> bool {
48+
self.inner.get(name).is_ok()
49+
}
50+
}
51+
1252
use super::acoustic_connector::AcousticConnector;
1353
use super::config_1_5b::*;
1454
use super::ddpm::DpmSolverPP;
@@ -69,10 +109,15 @@ impl VibeVoice1_5B {
69109
let common_cfg = config.into_config();
70110

71111
info!("Loading VibeVoice-1.5B...");
72-
let dtype = DType::BF16;
112+
let dtype = DType::F16;
73113

74114
let vb = unsafe {
75-
VarBuilder::from_mmaped_safetensors(weight_paths, dtype, device)?
115+
// Load via CPU to handle BF16→F16 cast (BF16 CUDA kernels need SM 8.0+).
116+
// MmapedSafetensors loads BF16 on CPU, casts to F16, then moves to device.
117+
let tensors = candle_core::safetensors::MmapedSafetensors::multi(weight_paths)?;
118+
let backend: Box<dyn candle_nn::var_builder::SimpleBackend> =
119+
Box::new(CpuCastBackend { inner: tensors, target_device: device.clone() });
120+
VarBuilder::from_backend(backend, dtype, device.clone())
76121
};
77122

78123
// Single LM (28 layers)

cake-core/src/utils/wav.rs

Lines changed: 14 additions & 10 deletions
Original file line numberDiff line numberDiff line change
@@ -31,19 +31,20 @@ pub fn encode_wav_bytes(samples: &[f32], sample_rate: u32) -> Vec<u8> {
3131
// data chunk
3232
buf.extend_from_slice(b"data");
3333
buf.extend_from_slice(&data_size.to_le_bytes());
34-
for &s in samples {
35-
let clamped = s.clamp(-1.0, 1.0);
36-
let i = (clamped * 32767.0) as i16;
37-
buf.extend_from_slice(&i.to_le_bytes());
38-
}
34+
// Convert all samples to i16 bytes in one pass — single extend_from_slice
35+
let sample_bytes: Vec<u8> = samples
36+
.iter()
37+
.flat_map(|&s| ((s.clamp(-1.0, 1.0) * 32767.0) as i16).to_le_bytes())
38+
.collect();
39+
buf.extend_from_slice(&sample_bytes);
3940
buf
4041
}
4142

4243
/// Save PCM f32 samples as a WAV file (16-bit PCM, mono).
4344
pub fn save_wav(samples: &[f32], path: &Path, sample_rate: u32) -> Result<()> {
44-
use std::io::Write;
45+
use std::io::{BufWriter, Write};
4546
let data_size = (samples.len() * 2) as u32;
46-
let mut f = std::fs::File::create(path)?;
47+
let mut f = BufWriter::new(std::fs::File::create(path)?);
4748
f.write_all(b"RIFF")?;
4849
f.write_all(&(36 + data_size).to_le_bytes())?;
4950
f.write_all(b"WAVEfmt ")?;
@@ -56,9 +57,12 @@ pub fn save_wav(samples: &[f32], path: &Path, sample_rate: u32) -> Result<()> {
5657
f.write_all(&16u16.to_le_bytes())?;
5758
f.write_all(b"data")?;
5859
f.write_all(&data_size.to_le_bytes())?;
59-
for &s in samples {
60-
f.write_all(&((s.clamp(-1.0, 1.0) * 32767.0) as i16).to_le_bytes())?;
61-
}
60+
// Batch-convert all samples to i16 bytes, then write once
61+
let sample_bytes: Vec<u8> = samples
62+
.iter()
63+
.flat_map(|&s| ((s.clamp(-1.0, 1.0) * 32767.0) as i16).to_le_bytes())
64+
.collect();
65+
f.write_all(&sample_bytes)?;
6266
Ok(())
6367
}
6468

0 commit comments

Comments
 (0)