use std::path::Path;
use std::sync::atomic::{AtomicBool, AtomicUsize, Ordering};
use std::sync::{Arc, Mutex};
use anyhow::Result;
use metal::{Buffer, ComputeCommandEncoderRef, ComputePipelineState};
use crate::backend::metal::{
ElementwiseParams, FlashAttnParams, MetalContext, MetalParams, QkNormRopeParams, shaders,
};
use crate::gguf::GgufFile;
use crate::model::audio_decoder::{DetokenizerConfig, DetokenizerWeights};
fn sz1d(x: u64) -> metal::MTLSize {
metal::MTLSize::new(x, 1, 1)
}
struct MetalWeight {
buf: Buffer, m: u32,
params_buf: Buffer,
}
struct DetokLayerGpu {
operator_norm: Buffer,
ffn_norm: Buffer,
ffn_w1: MetalWeight,
ffn_w2: MetalWeight,
ffn_w3: MetalWeight,
conv_in_proj: Option<MetalWeight>,
conv_out_proj: Option<MetalWeight>,
conv_weight: Option<Buffer>,
wq: Option<MetalWeight>,
wk: Option<MetalWeight>,
wv: Option<MetalWeight>,
wo: Option<MetalWeight>,
q_norm: Option<Buffer>,
k_norm: Option<Buffer>,
}
struct Pipelines {
gemv_f32: ComputePipelineState,
memcpy_f32: ComputePipelineState,
cast_f32_to_f16: ComputePipelineState,
mul_out: ComputePipelineState,
add_inplace: ComputePipelineState,
silu_mul_inplace: ComputePipelineState,
rmsnorm: ComputePipelineState,
qk_norm_rope: ComputePipelineState,
flash_attention: ComputePipelineState,
conv1d: ComputePipelineState,
exp_polar: ComputePipelineState,
overlap_add: ComputePipelineState,
}
struct Params {
rmsnorm_hs: Buffer,
elementwise_is: Buffer,
conv1d: Buffer,
}
pub struct MetalAudioDecoder {
ctx: MetalContext,
cfg: DetokenizerConfig,
pipes: Pipelines,
params: Params,
depthformer: Option<MetalDepthformer>,
layers: Vec<DetokLayerGpu>,
output_norm: Buffer,
lin_w: MetalWeight,
lin_b: Buffer,
idft_basis: MetalWeight,
hann: Buffer,
hidden_buf: Buffer,
normed_buf: Buffer,
accum_scratch: Buffer, proj_buf: Buffer, bx_buf: Buffer, conv_out_buf: Buffer, gate_buf: Buffer,
up_buf: Buffer,
q_buf: Buffer, k_buf: Buffer, v_buf: Buffer, rope_freqs_dummy: Buffer,
attn_out_buf: Buffer,
spectrum_buf: Buffer, tokens_buf: Buffer,
conv_bufs: Vec<Option<Buffer>>,
kv_k: Vec<Option<Buffer>>,
kv_v: Vec<Option<Buffer>>,
n_past: AtomicUsize,
infer_lock: Mutex<()>,
session_active: AtomicBool,
}
impl MetalAudioDecoder {
pub fn from_gguf(gguf: &Arc<GgufFile>, _vocoder_path: &Path) -> Result<Self> {
let cpu_dec = crate::model::audio_decoder::AudioDecoderWeights::from_gguf(gguf)?;
let ctx = MetalContext::new()?;
let (_, _, conv_in_cols, _) = gguf.tensor_meta("lfm.layers.0.conv.in_proj.weight")?;
let n_embd = conv_in_cols;
let q_norm_t = gguf.get_tensor("lfm.layers.2.self_attn.q_layernorm.weight")?;
let head_dim = *q_norm_t.shape().first().unwrap_or(&0);
anyhow::ensure!(head_dim > 0, "invalid q_layernorm head_dim");
anyhow::ensure!(
head_dim <= 128,
"detokenizer head_dim {head_dim} > 128; metal flash_attention cannot size q_shared"
);
let (_, q_rows, _, _) = gguf.tensor_meta("lfm.layers.2.self_attn.q_proj.weight")?;
anyhow::ensure!(
q_rows.is_multiple_of(head_dim),
"q_proj rows ({q_rows}) must be a multiple of head_dim ({head_dim})"
);
let n_head = q_rows / head_dim;
let (_, k_rows, _, _) = gguf.tensor_meta("lfm.layers.2.self_attn.k_proj.weight")?;
anyhow::ensure!(
k_rows.is_multiple_of(head_dim),
"k_proj rows ({k_rows}) must be a multiple of head_dim ({head_dim})"
);
let n_kv = k_rows / head_dim;
anyhow::ensure!(
n_kv > 0 && n_head.is_multiple_of(n_kv),
"detokenizer GQA requires n_kv > 0 and n_head divisible by n_kv (n_head={n_head}, n_kv={n_kv})"
);
let (_, ffn_rows, _, _) = gguf.tensor_meta("lfm.layers.0.feed_forward.w1.weight")?;
let ffn_dim = ffn_rows;
let layer_is_conv = vec![true, true, false, true, false, true, false, true];
let n_layer = layer_is_conv.len();
let kv_dim = n_kv * head_dim;
let cfg = DetokenizerConfig {
n_layer,
n_embd,
n_head,
n_head_kv: n_kv,
n_embd_head: head_dim,
ffn_dim,
d_conv: 2,
rms_norm_eps: 1e-5,
rope_freq_base: 1_000_000.0,
swa_window_size: 30,
n_codes: 8,
n_fft: 1280,
hop_length: 320,
sample_rate: 24000,
layer_is_conv,
};
let pipes = Pipelines {
gemv_f32: ctx.create_pipeline(shaders::GEMV_F32, "gemv_f32")?,
memcpy_f32: ctx.create_pipeline(shaders::ELEMENTWISE, "memcpy_f32")?,
cast_f32_to_f16: ctx.create_pipeline(shaders::ELEMENTWISE, "cast_f32_to_f16")?,
mul_out: ctx.create_pipeline(shaders::ELEMENTWISE, "mul_out")?,
add_inplace: ctx.create_pipeline(shaders::ELEMENTWISE, "add_inplace")?,
silu_mul_inplace: ctx.create_pipeline(shaders::ELEMENTWISE, "silu_mul_inplace")?,
rmsnorm: ctx.create_pipeline(shaders::RMSNORM, "rmsnorm")?,
qk_norm_rope: ctx.create_pipeline(shaders::QK_NORM_ROPE, "qk_norm_rope")?,
flash_attention: ctx.create_pipeline(shaders::FLASH_ATTENTION, "flash_attention")?,
conv1d: ctx.create_pipeline(shaders::CONV1D, "conv1d_depthwise")?,
exp_polar: ctx.create_pipeline(shaders::EXP_POLAR, "exp_polar")?,
overlap_add: ctx.create_pipeline(shaders::OVERLAP_ADD, "overlap_add")?,
};
let eps_bits = cfg.rms_norm_eps.to_bits();
let params = Params {
rmsnorm_hs: ctx.upload_bytes(bytemuck::cast_slice(&[
n_embd as u32,
eps_bits,
0u32,
0u32,
])),
elementwise_is: ctx.upload_bytes(bytemuck::cast_slice(&[ffn_dim as u32, 0u32])),
conv1d: ctx.upload_bytes(bytemuck::cast_slice(&[n_embd as u32, 3u32, 2u32, 0u32])),
};
let make_weight = |name: &str| -> Result<MetalWeight> {
let t = gguf.get_tensor(name)?;
let f32_data = t.to_f32_vec();
let shape = t.shape();
let (rows, cols) = match shape.len() {
1 => (1, shape[0]),
2 => (shape[1], shape[0]),
_ => anyhow::bail!("unexpected rank for {name}"),
};
let buf = ctx.upload_f32(&f32_data);
let params_buf = ctx.upload_bytes(bytemuck::cast_slice(&[rows as u32, cols as u32]));
Ok(MetalWeight {
buf,
m: rows as u32,
params_buf,
})
};
let mut layers = Vec::with_capacity(n_layer);
for i in 0..n_layer {
let pfx = format!("lfm.layers.{i}");
let is_conv = cfg.layer_is_conv[i];
let (cin, cop, cw) = if is_conv {
(
Some(make_weight(&format!("{pfx}.conv.in_proj.weight"))?),
Some(make_weight(&format!("{pfx}.conv.out_proj.weight"))?),
Some(
ctx.upload_f32(
&gguf
.get_tensor(&format!("{pfx}.conv.conv.weight"))?
.to_f32_vec(),
),
),
)
} else {
(None, None, None)
};
let (wq, wk, wv, wo, qn, kn) = if !is_conv {
(
Some(make_weight(&format!("{pfx}.self_attn.q_proj.weight"))?),
Some(make_weight(&format!("{pfx}.self_attn.k_proj.weight"))?),
Some(make_weight(&format!("{pfx}.self_attn.v_proj.weight"))?),
Some(make_weight(&format!("{pfx}.self_attn.out_proj.weight"))?),
Some(
ctx.upload_f32(
&gguf
.get_tensor(&format!("{pfx}.self_attn.q_layernorm.weight"))?
.to_f32_vec(),
),
),
Some(
ctx.upload_f32(
&gguf
.get_tensor(&format!("{pfx}.self_attn.k_layernorm.weight"))?
.to_f32_vec(),
),
),
)
} else {
(None, None, None, None, None, None)
};
layers.push(DetokLayerGpu {
operator_norm: ctx.upload_f32(
&gguf
.get_tensor(&format!("{pfx}.operator_norm.weight"))?
.to_f32_vec(),
),
ffn_norm: ctx.upload_f32(
&gguf
.get_tensor(&format!("{pfx}.ffn_norm.weight"))?
.to_f32_vec(),
),
ffn_w1: make_weight(&format!("{pfx}.feed_forward.w1.weight"))?,
ffn_w2: make_weight(&format!("{pfx}.feed_forward.w2.weight"))?,
ffn_w3: make_weight(&format!("{pfx}.feed_forward.w3.weight"))?,
conv_in_proj: cin,
conv_out_proj: cop,
conv_weight: cw,
wq,
wk,
wv,
wo,
q_norm: qn,
k_norm: kn,
});
}
let output_norm =
ctx.upload_f32(&gguf.get_tensor("lfm.embedding_norm.weight")?.to_f32_vec());
let lin_w = make_weight("lin.weight")?;
let lin_b = ctx.upload_f32(&gguf.get_tensor("lin.bias")?.to_f32_vec());
let n_fft_bins = cfg.n_fft / 2 + 1;
let idft_basis = {
let basis = crate::model::audio_decoder::build_idft_basis(cfg.n_fft);
let buf = ctx.upload_f32(&basis);
let params_buf = ctx.upload_bytes(bytemuck::cast_slice(&[
cfg.n_fft as u32,
(n_fft_bins * 2) as u32,
]));
MetalWeight {
buf,
m: cfg.n_fft as u32,
params_buf,
}
};
let hann = ctx.upload_f32(&crate::model::audio_decoder::build_hann(cfg.n_fft));
let spectrum_size = 6 * (cfg.n_fft / 2 + 1) * 2;
let ab = |n: usize| ctx.create_buffer((n * 4) as u64);
let rope_freqs_dummy = ctx.freq_factors_dummy();
let mut conv_bufs = vec![None; n_layer];
let mut kv_k = vec![None; n_layer];
let mut kv_v = vec![None; n_layer];
for i in 0..n_layer {
if cfg.layer_is_conv[i] {
conv_bufs[i] = Some(ab(cfg.d_conv * n_embd));
} else {
let kv_bytes = (cfg.swa_window_size * kv_dim * 2) as u64;
kv_k[i] = Some(ctx.create_buffer(kv_bytes));
kv_v[i] = Some(ctx.create_buffer(kv_bytes));
}
}
let hidden_buf = ab(n_embd);
let normed_buf = ab(n_embd);
let accum_scratch = ab(n_embd.max(ffn_dim));
let proj_buf = ab(3 * n_embd);
let bx_buf = ab(n_embd);
let conv_out_buf = ab(n_embd);
let gate_buf = ab(ffn_dim);
let up_buf = ab(ffn_dim);
let q_buf = ab(n_head * head_dim);
let k_buf = ab(n_kv * head_dim);
let v_buf = ab(n_kv * head_dim);
let attn_out_buf = ab(n_embd);
let spectrum_buf = ab(spectrum_size);
let tokens_buf = ab(6 * n_embd);
let depthformer = match MetalDepthformer::from_gguf(
gguf,
_vocoder_path,
&cpu_dec.depthformer_config,
&cpu_dec.decoder_config,
) {
Ok(df) => {
tracing::debug!("Metal depthformer loaded");
Some(df)
}
Err(e) => {
tracing::debug!("Metal depthformer failed: {e}, using CPU");
None
}
};
Ok(Self {
ctx,
cfg,
pipes,
params,
depthformer,
layers,
rope_freqs_dummy,
output_norm,
lin_w,
lin_b,
idft_basis,
hann,
hidden_buf,
normed_buf,
accum_scratch,
proj_buf,
bx_buf,
conv_out_buf,
gate_buf,
up_buf,
q_buf,
k_buf,
v_buf,
attn_out_buf,
spectrum_buf,
tokens_buf,
conv_bufs,
kv_k,
kv_v,
n_past: AtomicUsize::new(0),
infer_lock: Mutex::new(()),
session_active: AtomicBool::new(false),
})
}
pub fn reset(&self) {
let _guard = self.infer_lock.lock().unwrap_or_else(|e| e.into_inner());
self.n_past.store(0, Ordering::Relaxed);
for b in self.conv_bufs.iter().flatten() {
unsafe {
std::ptr::write_bytes(b.contents() as *mut u8, 0, b.length() as usize);
}
}
}
pub fn config(&self) -> &DetokenizerConfig {
&self.cfg
}
fn encode_gemv(
&self,
enc: &ComputeCommandEncoderRef,
w: &MetalWeight,
input: &Buffer,
output: &Buffer,
) {
let groups = w.m as u64;
enc.set_compute_pipeline_state(&self.pipes.gemv_f32);
enc.set_buffer(0, Some(&w.buf), 0);
enc.set_buffer(1, Some(input), 0);
enc.set_buffer(2, Some(output), 0);
enc.set_buffer(3, Some(&w.params_buf), 0);
enc.dispatch_thread_groups(sz1d(groups), sz1d(32));
}
fn encode_gemv_accum(
&self,
enc: &ComputeCommandEncoderRef,
w: &MetalWeight,
input: &Buffer,
output: &Buffer,
scratch: &Buffer,
) {
self.encode_gemv(enc, w, input, scratch);
self.barrier(enc);
let n = w.m;
let params = ElementwiseParams::new(n);
enc.set_compute_pipeline_state(&self.pipes.add_inplace);
enc.set_buffer(0, Some(output), 0);
enc.set_buffer(1, Some(scratch), 0);
params.set(enc, 2);
enc.dispatch_thread_groups(sz1d((n as u64).div_ceil(256)), sz1d(256));
}
fn encode_rmsnorm(&self, enc: &ComputeCommandEncoderRef, buf: &Buffer, weight: &Buffer) {
enc.set_compute_pipeline_state(&self.pipes.rmsnorm);
enc.set_buffer(0, Some(buf), 0); enc.set_buffer(1, Some(buf), 0); enc.set_buffer(2, Some(weight), 0);
enc.set_buffer(3, Some(&self.params.rmsnorm_hs), 0);
enc.dispatch_thread_groups(sz1d(1), sz1d(256));
}
fn encode_rmsnorm_out(
&self,
enc: &ComputeCommandEncoderRef,
src: &Buffer,
dst: &Buffer,
weight: &Buffer,
) {
enc.set_compute_pipeline_state(&self.pipes.rmsnorm);
enc.set_buffer(0, Some(src), 0);
enc.set_buffer(1, Some(dst), 0);
enc.set_buffer(2, Some(weight), 0);
enc.set_buffer(3, Some(&self.params.rmsnorm_hs), 0);
enc.dispatch_thread_groups(sz1d(1), sz1d(256));
}
fn encode_memcpy(
&self,
enc: &ComputeCommandEncoderRef,
src: &Buffer,
src_off: u64,
dst: &Buffer,
dst_off: u64,
n: usize,
) {
let params = ElementwiseParams::new(n as u32);
enc.set_compute_pipeline_state(&self.pipes.memcpy_f32);
enc.set_buffer(0, Some(src), src_off);
enc.set_buffer(1, Some(dst), dst_off);
params.set(enc, 2);
enc.dispatch_thread_groups(sz1d((n as u64).div_ceil(256)), sz1d(256));
}
fn encode_cast_f32_to_f16(
&self,
enc: &ComputeCommandEncoderRef,
src: &Buffer,
dst: &Buffer,
dst_off_bytes: u64,
n: u32,
) {
let params = ElementwiseParams::new(n);
enc.set_compute_pipeline_state(&self.pipes.cast_f32_to_f16);
enc.set_buffer(0, Some(src), 0);
enc.set_buffer(1, Some(dst), dst_off_bytes);
params.set(enc, 2);
enc.dispatch_thread_groups(sz1d((n as u64).div_ceil(256)), sz1d(256));
}
#[allow(clippy::too_many_arguments)]
fn encode_mul_out(
&self,
enc: &ComputeCommandEncoderRef,
a: &Buffer,
a_off: u64,
b: &Buffer,
b_off: u64,
dst: &Buffer,
n: usize,
) {
let params = ElementwiseParams::new(n as u32);
enc.set_compute_pipeline_state(&self.pipes.mul_out);
enc.set_buffer(0, Some(a), a_off);
enc.set_buffer(1, Some(b), b_off);
enc.set_buffer(2, Some(dst), 0);
params.set(enc, 3);
enc.dispatch_thread_groups(sz1d((n as u64).div_ceil(256)), sz1d(256));
}
fn encode_silu_mul(&self, enc: &ComputeCommandEncoderRef) {
enc.set_compute_pipeline_state(&self.pipes.silu_mul_inplace);
enc.set_buffer(0, Some(&self.gate_buf), 0);
enc.set_buffer(1, Some(&self.up_buf), 0);
enc.set_buffer(2, Some(&self.params.elementwise_is), 0);
enc.dispatch_thread_groups(sz1d((self.cfg.ffn_dim as u64).div_ceil(256)), sz1d(256));
}
fn encode_conv1d(
&self,
enc: &ComputeCommandEncoderRef,
input: &Buffer,
rbuf: &Buffer,
weight: &Buffer,
output: &Buffer,
) {
let n = self.cfg.n_embd as u64;
enc.set_compute_pipeline_state(&self.pipes.conv1d);
enc.set_buffer(0, Some(input), 0);
enc.set_buffer(1, Some(rbuf), 0);
enc.set_buffer(2, Some(weight), 0);
enc.set_buffer(3, Some(output), 0);
enc.set_buffer(4, Some(&self.params.conv1d), 0);
enc.dispatch_thread_groups(sz1d(n.div_ceil(256)), sz1d(256));
}
#[inline(always)]
fn barrier(&self, _enc: &ComputeCommandEncoderRef) {}
fn encode_conv_layer(
&self,
enc: &ComputeCommandEncoderRef,
lw: &DetokLayerGpu,
il: usize,
n_tokens: usize,
) {
let hs = self.cfg.n_embd;
let rbuf = self.conv_bufs[il].as_ref().unwrap();
let cw = lw.conv_weight.as_ref().unwrap();
let cin = lw.conv_in_proj.as_ref().unwrap();
let cop = lw.conv_out_proj.as_ref().unwrap();
for t in 0..n_tokens {
let tok_off = (t * hs * 4) as u64;
self.encode_memcpy(enc, &self.tokens_buf, tok_off, &self.hidden_buf, 0, hs);
self.barrier(enc);
self.encode_rmsnorm_out(enc, &self.hidden_buf, &self.normed_buf, &lw.operator_norm);
self.barrier(enc);
self.encode_gemv(enc, cin, &self.normed_buf, &self.proj_buf);
self.barrier(enc);
self.encode_mul_out(
enc,
&self.proj_buf,
0, &self.proj_buf,
(2 * hs * 4) as u64, &self.bx_buf,
hs,
);
self.barrier(enc);
self.encode_conv1d(enc, &self.bx_buf, rbuf, cw, &self.conv_out_buf);
self.barrier(enc);
self.encode_mul_out(
enc,
&self.proj_buf,
(hs * 4) as u64, &self.conv_out_buf,
0,
&self.bx_buf,
hs,
); self.barrier(enc);
self.encode_gemv_accum(
enc,
cop,
&self.bx_buf,
&self.hidden_buf,
&self.accum_scratch,
);
self.barrier(enc);
self.encode_rmsnorm_out(enc, &self.hidden_buf, &self.normed_buf, &lw.ffn_norm);
self.barrier(enc);
self.encode_gemv(enc, &lw.ffn_w1, &self.normed_buf, &self.gate_buf);
self.encode_gemv(enc, &lw.ffn_w3, &self.normed_buf, &self.up_buf);
self.barrier(enc);
self.encode_silu_mul(enc);
self.barrier(enc);
self.encode_gemv_accum(
enc,
&lw.ffn_w2,
&self.gate_buf,
&self.hidden_buf,
&self.accum_scratch,
);
self.barrier(enc);
self.encode_memcpy(enc, &self.hidden_buf, 0, &self.tokens_buf, tok_off, hs);
self.barrier(enc);
}
}
fn encode_attn_layer(
&self,
enc: &ComputeCommandEncoderRef,
lw: &DetokLayerGpu,
il: usize,
n_tokens: usize,
n_past: usize,
) {
let hs = self.cfg.n_embd;
let hd = self.cfg.n_embd_head;
let n_heads = self.cfg.n_head as u32;
let n_kv = self.cfg.n_head_kv as u32;
let kv_dim = n_kv as usize * hd;
let k_cache = self.kv_k[il].as_ref().unwrap();
let v_cache = self.kv_v[il].as_ref().unwrap();
let wq = lw.wq.as_ref().unwrap();
let wk = lw.wk.as_ref().unwrap();
let wv = lw.wv.as_ref().unwrap();
let wo = lw.wo.as_ref().unwrap();
let qn = lw.q_norm.as_ref().unwrap();
let kn = lw.k_norm.as_ref().unwrap();
let freq_bits = self.cfg.rope_freq_base.to_bits();
let eps_bits = self.cfg.rms_norm_eps.to_bits();
for t in 0..n_tokens {
let tok_off = (t * hs * 4) as u64;
let pos = n_past + t;
self.encode_memcpy(enc, &self.tokens_buf, tok_off, &self.hidden_buf, 0, hs);
self.barrier(enc);
self.encode_rmsnorm_out(enc, &self.hidden_buf, &self.normed_buf, &lw.operator_norm);
self.barrier(enc);
self.encode_gemv(enc, wk, &self.normed_buf, &self.k_buf);
self.encode_gemv(enc, wv, &self.normed_buf, &self.v_buf);
self.barrier(enc);
let k_rope = QkNormRopeParams {
pos: pos as u32,
n_heads: 0,
n_kv_heads: n_kv,
head_dim: hd as u32,
eps_bits,
freq_base_bits: freq_bits,
rope_type: 0, has_freq_factors: 0,
has_qk_norm: 1,
};
enc.set_compute_pipeline_state(&self.pipes.qk_norm_rope);
enc.set_buffer(0, Some(&self.k_buf), 0);
enc.set_buffer(1, Some(&self.k_buf), 0);
enc.set_buffer(2, Some(kn), 0);
enc.set_buffer(3, Some(kn), 0);
k_rope.bind(enc, &self.rope_freqs_dummy);
enc.dispatch_thread_groups(sz1d(n_kv as u64), sz1d(256));
self.barrier(enc);
let cache_pos = pos % self.cfg.swa_window_size;
let cache_off = (cache_pos * kv_dim * 2) as u64;
self.encode_cast_f32_to_f16(enc, &self.k_buf, k_cache, cache_off, kv_dim as u32);
self.encode_cast_f32_to_f16(enc, &self.v_buf, v_cache, cache_off, kv_dim as u32);
self.barrier(enc);
}
let seq_len = (n_past + n_tokens).min(self.cfg.swa_window_size);
let scale = 1.0f32 / (hd as f32).sqrt();
for t in 0..n_tokens {
let tok_off = (t * hs * 4) as u64;
let pos = n_past + t;
self.encode_memcpy(enc, &self.tokens_buf, tok_off, &self.hidden_buf, 0, hs);
self.barrier(enc);
self.encode_rmsnorm_out(enc, &self.hidden_buf, &self.normed_buf, &lw.operator_norm);
self.barrier(enc);
self.encode_gemv(enc, wq, &self.normed_buf, &self.q_buf);
self.barrier(enc);
let q_rope = QkNormRopeParams {
pos: pos as u32,
n_heads,
n_kv_heads: 0,
head_dim: hd as u32,
eps_bits,
freq_base_bits: freq_bits,
rope_type: 0, has_freq_factors: 0,
has_qk_norm: 1,
};
enc.set_compute_pipeline_state(&self.pipes.qk_norm_rope);
enc.set_buffer(0, Some(&self.q_buf), 0);
enc.set_buffer(1, Some(&self.q_buf), 0);
enc.set_buffer(2, Some(qn), 0);
enc.set_buffer(3, Some(qn), 0);
q_rope.bind(enc, &self.rope_freqs_dummy);
enc.dispatch_thread_groups(sz1d(n_heads as u64), sz1d(256));
self.barrier(enc);
let attn_params = FlashAttnParams {
n_heads,
n_kv_heads: n_kv,
head_dim: hd as u32,
kv_dim: kv_dim as u32,
seq_len: seq_len as u32,
scale_bits: scale.to_bits(),
_pad0: 0,
_pad1: 0,
};
enc.set_compute_pipeline_state(&self.pipes.flash_attention);
enc.set_buffer(0, Some(&self.q_buf), 0);
enc.set_buffer(1, Some(k_cache), 0);
enc.set_buffer(2, Some(v_cache), 0);
enc.set_buffer(3, Some(&self.attn_out_buf), 0);
attn_params.set(enc, 4);
enc.dispatch_thread_groups(sz1d(n_heads as u64), sz1d(256));
self.barrier(enc);
self.encode_gemv_accum(
enc,
wo,
&self.attn_out_buf,
&self.hidden_buf,
&self.accum_scratch,
);
self.barrier(enc);
self.encode_rmsnorm_out(enc, &self.hidden_buf, &self.normed_buf, &lw.ffn_norm);
self.barrier(enc);
self.encode_gemv(enc, &lw.ffn_w1, &self.normed_buf, &self.gate_buf);
self.encode_gemv(enc, &lw.ffn_w3, &self.normed_buf, &self.up_buf);
self.barrier(enc);
self.encode_silu_mul(enc);
self.barrier(enc);
self.encode_gemv_accum(
enc,
&lw.ffn_w2,
&self.gate_buf,
&self.hidden_buf,
&self.accum_scratch,
);
self.barrier(enc);
self.encode_memcpy(enc, &self.hidden_buf, 0, &self.tokens_buf, tok_off, hs);
self.barrier(enc);
}
}
pub fn detokenize_to_spectrum(
&self,
cpu_weights: &DetokenizerWeights,
codes: &[i32],
) -> Vec<f32> {
let _guard = self.infer_lock.lock().unwrap_or_else(|e| e.into_inner());
let hs = self.cfg.n_embd;
let n_frames = 6usize;
let n_fft_bins = self.cfg.n_fft / 2 + 1; let spectrum_per_frame = n_fft_bins * 2;
let tokens = {
use crate::model::audio_decoder::{detok_embed_codes, upsample};
let emb = detok_embed_codes(cpu_weights, codes);
upsample(&emb, hs, n_frames)
};
unsafe {
let dst = self.tokens_buf.contents() as *mut f32;
std::ptr::copy_nonoverlapping(tokens.as_ptr(), dst, tokens.len());
}
let n_past = self.n_past.load(Ordering::Relaxed);
for (il, lw) in self.layers.iter().enumerate() {
let cb = self.ctx.queue.new_command_buffer();
let enc = cb.new_compute_command_encoder();
if self.cfg.layer_is_conv[il] {
self.encode_conv_layer(enc, lw, il, n_frames);
} else {
self.encode_attn_layer(enc, lw, il, n_frames, n_past);
}
enc.end_encoding();
cb.commit();
}
{
let cb = self.ctx.queue.new_command_buffer();
let enc = cb.new_compute_command_encoder();
for t in 0..n_frames {
let tok_off = (t * hs * 4) as u64;
let spec_off = (t * spectrum_per_frame * 4) as u64;
self.encode_memcpy(enc, &self.tokens_buf, tok_off, &self.hidden_buf, 0, hs);
self.barrier(enc);
self.encode_rmsnorm(enc, &self.hidden_buf, &self.output_norm);
self.barrier(enc);
enc.set_compute_pipeline_state(&self.pipes.gemv_f32);
enc.set_buffer(0, Some(&self.lin_w.buf), 0);
enc.set_buffer(1, Some(&self.hidden_buf), 0);
enc.set_buffer(2, Some(&self.spectrum_buf), spec_off);
enc.set_buffer(3, Some(&self.lin_w.params_buf), 0);
enc.dispatch_thread_groups(sz1d(self.lin_w.m as u64), sz1d(32));
self.barrier(enc);
let bias_n = spectrum_per_frame as u32;
let bias_params = ElementwiseParams::new(bias_n);
enc.set_compute_pipeline_state(&self.pipes.add_inplace);
enc.set_buffer(0, Some(&self.spectrum_buf), spec_off);
enc.set_buffer(1, Some(&self.lin_b), 0);
bias_params.set(enc, 2);
enc.dispatch_thread_groups(sz1d((bias_n as u64).div_ceil(256)), sz1d(256));
self.barrier(enc);
}
enc.end_encoding();
cb.commit();
cb.wait_until_completed();
}
self.n_past.store(n_past + n_frames, Ordering::Relaxed);
let total = n_frames * spectrum_per_frame;
let mut spectrum = vec![0.0f32; total];
unsafe {
let src = self.spectrum_buf.contents() as *const f32;
std::ptr::copy_nonoverlapping(src, spectrum.as_mut_ptr(), total);
}
spectrum
}
pub fn istft_to_pcm(&self, spectrum: &[f32], n_fft: usize, hop_length: usize) -> Vec<f32> {
let _guard = self.infer_lock.lock().unwrap_or_else(|e| e.into_inner());
if n_fft != self.cfg.n_fft || hop_length != self.cfg.hop_length {
return crate::model::audio_decoder::istft_to_pcm(spectrum, n_fft, hop_length);
}
let bins = n_fft / 2 + 1;
let frame_size = bins * 2;
let n_frames = spectrum.len() / frame_size;
if n_frames == 0 {
return vec![];
}
let spec_buf = self.ctx.upload_f32(spectrum);
let halfspec = self.ctx.create_buffer((n_frames * frame_size * 4) as u64);
let time_domain = self.ctx.create_buffer((n_frames * n_fft * 4) as u64);
let pcm_buf = self.ctx.create_buffer((n_frames * hop_length * 4) as u64);
let ep_params = self
.ctx
.upload_bytes(bytemuck::cast_slice(&[n_frames as u32, bins as u32]));
let oa_params = self.ctx.upload_bytes(bytemuck::cast_slice(&[
n_frames as u32,
n_fft as u32,
hop_length as u32,
0u32,
]));
let cb = self.ctx.queue.new_command_buffer();
let enc = cb.new_compute_command_encoder();
enc.set_compute_pipeline_state(&self.pipes.exp_polar);
enc.set_buffer(0, Some(&spec_buf), 0);
enc.set_buffer(1, Some(&halfspec), 0);
enc.set_buffer(2, Some(&ep_params), 0);
enc.dispatch_thread_groups(sz1d(((n_frames * bins) as u64).div_ceil(256)), sz1d(256));
self.barrier(enc);
enc.set_compute_pipeline_state(&self.pipes.gemv_f32);
enc.set_buffer(0, Some(&self.idft_basis.buf), 0);
enc.set_buffer(3, Some(&self.idft_basis.params_buf), 0);
for f in 0..n_frames {
enc.set_buffer(1, Some(&halfspec), (f * frame_size * 4) as u64);
enc.set_buffer(2, Some(&time_domain), (f * n_fft * 4) as u64);
enc.dispatch_thread_groups(sz1d(self.idft_basis.m as u64), sz1d(32));
}
self.barrier(enc);
enc.set_compute_pipeline_state(&self.pipes.overlap_add);
enc.set_buffer(0, Some(&time_domain), 0);
enc.set_buffer(1, Some(&self.hann), 0);
enc.set_buffer(2, Some(&pcm_buf), 0);
enc.set_buffer(3, Some(&oa_params), 0);
enc.dispatch_thread_groups(
sz1d(((n_frames * hop_length) as u64).div_ceil(256)),
sz1d(256),
);
enc.end_encoding();
cb.commit();
cb.wait_until_completed();
let total = n_frames * hop_length;
let mut pcm = vec![0.0f32; total];
unsafe {
let src = pcm_buf.contents() as *const f32;
std::ptr::copy_nonoverlapping(src, pcm.as_mut_ptr(), total);
}
let padding = (n_fft - hop_length) / 2;
if pcm.len() > padding {
pcm.drain(..padding);
}
pcm
}
}
struct DfLayerGpu {
operator_norm: Buffer,
wqkv: MetalWeight,
q_norm: Buffer,
k_norm: Buffer,
wo: MetalWeight,
ffn_norm: Buffer,
w1: MetalWeight,
w2: MetalWeight,
w3: MetalWeight,
}
struct DfPipelines {
gemv_f32: ComputePipelineState,
add_inplace: ComputePipelineState,
cast_f32_to_f16: ComputePipelineState,
rmsnorm: ComputePipelineState,
qk_norm_rope: ComputePipelineState,
flash_attention: ComputePipelineState,
silu_mul_inplace: ComputePipelineState,
}
pub struct MetalDepthformer {
ctx: MetalContext,
df_cfg: crate::model::audio_decoder::DepthformerConfig,
dec_cfg: crate::model::audio_decoder::DecoderConfig,
pipes: DfPipelines,
layers: Vec<DfLayerGpu>,
depth_linear_slices: Vec<MetalWeight>, depth_linear_biases: Vec<Buffer>, codebook_norms: Vec<Buffer>,
codebook_to_logits: Vec<MetalWeight>, codebook_emb_f32: Vec<Buffer>, src_emb_buf: Buffer,
dl_cols: usize,
hidden_buf: Buffer,
normed_buf: Buffer,
accum_buf: Buffer, qkv_buf: Buffer,
attn_out_buf: Buffer,
gate_buf: Buffer,
up_buf: Buffer,
logits_buf: Buffer,
rope_freqs_dummy: Buffer,
kv_k: Vec<Buffer>,
kv_v: Vec<Buffer>,
n_past: AtomicUsize,
rmsnorm_params: Buffer,
elementwise_is: Buffer,
}
impl MetalDepthformer {
pub fn from_gguf(
gguf: &GgufFile,
_vocoder_path: &Path,
df_cfg: &crate::model::audio_decoder::DepthformerConfig,
dec_cfg: &crate::model::audio_decoder::DecoderConfig,
) -> Result<Self> {
let ctx = MetalContext::new()?;
let n_embd = df_cfg.n_embd;
let hd = df_cfg.n_embd_head;
let n_kv = df_cfg.n_head_kv;
let kv_dim = n_kv * hd;
let pipes = DfPipelines {
gemv_f32: ctx.create_pipeline(shaders::GEMV_F32, "gemv_f32")?,
add_inplace: ctx.create_pipeline(shaders::ELEMENTWISE, "add_inplace")?,
cast_f32_to_f16: ctx.create_pipeline(shaders::ELEMENTWISE, "cast_f32_to_f16")?,
rmsnorm: ctx.create_pipeline(shaders::RMSNORM, "rmsnorm")?,
qk_norm_rope: ctx.create_pipeline(shaders::QK_NORM_ROPE, "qk_norm_rope")?,
flash_attention: ctx.create_pipeline(shaders::FLASH_ATTENTION, "flash_attention")?,
silu_mul_inplace: ctx.create_pipeline(shaders::ELEMENTWISE, "silu_mul_inplace")?,
};
let make_f32 = |name: &str| -> Result<MetalWeight> {
let t = gguf.get_tensor(name)?;
let f32_data = t.to_f32_vec();
let shape = t.shape();
let (rows, cols) = match shape.len() {
1 => (1, shape[0]),
2 => (shape[1], shape[0]),
_ => anyhow::bail!("unexpected rank for {name}"),
};
let buf = ctx.upload_f32(&f32_data);
let params_buf = ctx.upload_bytes(bytemuck::cast_slice(&[rows as u32, cols as u32]));
Ok(MetalWeight {
buf,
m: rows as u32,
params_buf,
})
};
let mut layers = Vec::with_capacity(df_cfg.n_layer);
for i in 0..df_cfg.n_layer {
let pfx = format!("depthformer.layers.{i}");
layers.push(DfLayerGpu {
operator_norm: ctx.upload_f32(
&gguf
.get_tensor(&format!("{pfx}.operator_norm.weight"))?
.to_f32_vec(),
),
wqkv: make_f32(&format!("{pfx}.operator.qkv_proj.weight"))?,
q_norm: ctx.upload_f32(
&gguf
.get_tensor(&format!("{pfx}.operator.attention.q_layernorm.weight"))?
.to_f32_vec(),
),
k_norm: ctx.upload_f32(
&gguf
.get_tensor(&format!("{pfx}.operator.attention.k_layernorm.weight"))?
.to_f32_vec(),
),
wo: make_f32(&format!("{pfx}.operator.out_proj.weight"))?,
ffn_norm: ctx.upload_f32(
&gguf
.get_tensor(&format!("{pfx}.ffn_norm.weight"))?
.to_f32_vec(),
),
w1: make_f32(&format!("{pfx}.feed_forward.w1.weight"))?,
w2: make_f32(&format!("{pfx}.feed_forward.w2.weight"))?,
w3: make_f32(&format!("{pfx}.feed_forward.w3.weight"))?,
});
}
let dl_tensor = gguf.get_tensor("depth_linear.weight")?;
let dl_f32 = dl_tensor.to_f32_vec();
let dl_shape = dl_tensor.shape();
let dl_cols = dl_shape[0]; let dl_rows = dl_shape[1]; let n_embd_d = dl_rows / dec_cfg.n_codebook;
let mut depth_linear_slices = Vec::with_capacity(dec_cfg.n_codebook);
for j in 0..dec_cfg.n_codebook {
let start = j * n_embd_d * dl_cols;
let end = start + n_embd_d * dl_cols;
let slice_data = &dl_f32[start..end];
let buf = ctx.upload_f32(slice_data);
let params_buf =
ctx.upload_bytes(bytemuck::cast_slice(&[n_embd_d as u32, dl_cols as u32]));
depth_linear_slices.push(MetalWeight {
buf,
m: n_embd_d as u32,
params_buf,
});
}
let dl_b = gguf
.get_tensor("depth_linear.bias")
.map(|t| t.to_f32_vec())
.unwrap_or_default();
let mut depth_linear_biases = Vec::with_capacity(dec_cfg.n_codebook);
for j in 0..dec_cfg.n_codebook {
let b_start = j * n_embd_d;
let b_end = (b_start + n_embd_d).min(dl_b.len());
let b_slice = if b_start < dl_b.len() {
&dl_b[b_start..b_end]
} else {
&[]
};
let mut b_vec = vec![0.0f32; n_embd_d];
b_vec[..b_slice.len()].copy_from_slice(b_slice);
depth_linear_biases.push(ctx.upload_f32(&b_vec));
}
let mut codebook_norms = Vec::with_capacity(dec_cfg.n_codebook);
let mut codebook_to_logits = Vec::with_capacity(dec_cfg.n_codebook);
let mut codebook_emb_f32 = Vec::with_capacity(dec_cfg.n_codebook);
for j in 0..dec_cfg.n_codebook {
let pfx = format!("depth_embeddings.{j}");
codebook_norms.push(
ctx.upload_f32(
&gguf
.get_tensor(&format!("{pfx}.embedding_norm.weight"))?
.to_f32_vec(),
),
);
codebook_to_logits.push(make_f32(&format!("{pfx}.to_logits.weight"))?);
codebook_emb_f32.push(
ctx.upload_f32(
&gguf
.get_tensor(&format!("{pfx}.embedding.weight"))?
.to_f32_vec(),
),
);
}
let eps_bits = df_cfg.rms_norm_eps.to_bits();
let rmsnorm_params =
ctx.upload_bytes(bytemuck::cast_slice(&[n_embd as u32, eps_bits, 0u32, 0u32]));
let elementwise_is = ctx.upload_bytes(bytemuck::cast_slice(&[df_cfg.ffn_dim as u32, 0u32]));
let q_dim = df_cfg.n_head * hd;
let qkv_dim = q_dim + 2 * kv_dim;
let buf = |n: usize| ctx.create_buffer((n * 4) as u64);
let rope_freqs_dummy = ctx.freq_factors_dummy();
let mut kv_k = Vec::with_capacity(df_cfg.n_layer);
let mut kv_v = Vec::with_capacity(df_cfg.n_layer);
for _ in 0..df_cfg.n_layer {
kv_k.push(ctx.create_buffer((df_cfg.max_seq_len * kv_dim * 2) as u64));
kv_v.push(ctx.create_buffer((df_cfg.max_seq_len * kv_dim * 2) as u64));
}
let src_emb_buf = buf(dl_cols);
let hidden_buf = buf(n_embd);
let normed_buf = buf(n_embd);
let accum_buf = buf(n_embd.max(df_cfg.ffn_dim));
let qkv_buf = buf(qkv_dim);
let attn_out_buf = buf(q_dim);
let gate_buf = buf(df_cfg.ffn_dim);
let up_buf = buf(df_cfg.ffn_dim);
let logits_buf = buf(dec_cfg.n_vocab);
Ok(Self {
ctx,
df_cfg: df_cfg.clone(),
dec_cfg: dec_cfg.clone(),
pipes,
layers,
rope_freqs_dummy,
depth_linear_slices,
depth_linear_biases,
codebook_norms,
codebook_to_logits,
codebook_emb_f32,
src_emb_buf,
dl_cols,
hidden_buf,
normed_buf,
accum_buf,
qkv_buf,
attn_out_buf,
gate_buf,
up_buf,
logits_buf,
kv_k,
kv_v,
n_past: std::sync::atomic::AtomicUsize::new(0),
rmsnorm_params,
elementwise_is,
})
}
pub fn reset(&self) {
self.n_past.store(0, std::sync::atomic::Ordering::Relaxed);
}
pub fn sample_frame(&self, embedding: &[f32], temperature: f32, top_k: usize) -> [i32; 8] {
let cfg = &self.df_cfg;
let dec = &self.dec_cfg;
let n_embd = cfg.n_embd;
let n_head = cfg.n_head as u32;
let n_kv = cfg.n_head_kv as u32;
let hd = cfg.n_embd_head;
let kv_dim = n_kv as usize * hd;
let q_dim = n_head as usize * hd;
let scale = 1.0f32 / (hd as f32).sqrt();
let eps_bits = cfg.rms_norm_eps.to_bits();
let freq_bits = cfg.rope_freq_base.to_bits();
self.reset();
unsafe {
let dst = self.src_emb_buf.contents() as *mut f32;
let copy_len = embedding.len().min(self.dl_cols);
if copy_len > 0 {
std::ptr::copy_nonoverlapping(embedding.as_ptr(), dst, copy_len);
}
if copy_len < self.dl_cols {
std::ptr::write_bytes(dst.add(copy_len), 0, self.dl_cols - copy_len);
}
}
let mut codes = [0i32; 8];
let mut prev_token: i32 = -1;
for (j, code) in codes.iter_mut().enumerate().take(dec.n_codebook) {
let pos = self.n_past.load(std::sync::atomic::Ordering::Relaxed);
let cb = self.ctx.queue.new_command_buffer();
let enc = cb.new_compute_command_encoder();
let dl = &self.depth_linear_slices[j];
self.encode_df_gemv(enc, dl, &self.src_emb_buf, &self.hidden_buf);
if let Some(b_buf) = self.depth_linear_biases.get(j) {
let params = ElementwiseParams::new(n_embd as u32);
enc.set_compute_pipeline_state(&self.pipes.add_inplace);
enc.set_buffer(0, Some(&self.hidden_buf), 0);
enc.set_buffer(1, Some(b_buf), 0);
params.set(enc, 2);
enc.dispatch_thread_groups(sz1d((n_embd as u64).div_ceil(256)), sz1d(256));
}
if j > 0 && prev_token >= 0 {
let emb_buf = &self.codebook_emb_f32[j - 1];
let tok = prev_token as usize;
let offset = (tok * n_embd * std::mem::size_of::<f32>()) as u64;
let params = ElementwiseParams::new(n_embd as u32);
enc.set_compute_pipeline_state(&self.pipes.add_inplace);
enc.set_buffer(0, Some(&self.hidden_buf), 0);
enc.set_buffer(1, Some(emb_buf), offset);
params.set(enc, 2);
enc.dispatch_thread_groups(sz1d((n_embd as u64).div_ceil(256)), sz1d(256));
}
for (il, lw) in self.layers.iter().enumerate() {
enc.set_compute_pipeline_state(&self.pipes.rmsnorm);
enc.set_buffer(0, Some(&self.hidden_buf), 0);
enc.set_buffer(1, Some(&self.normed_buf), 0);
enc.set_buffer(2, Some(&lw.operator_norm), 0);
enc.set_buffer(3, Some(&self.rmsnorm_params), 0);
enc.dispatch_thread_groups(sz1d(1), sz1d(256));
self.encode_df_gemv(enc, &lw.wqkv, &self.normed_buf, &self.qkv_buf);
let rope_params = QkNormRopeParams {
pos: pos as u32,
n_heads: n_head,
n_kv_heads: n_kv,
head_dim: hd as u32,
eps_bits,
freq_base_bits: freq_bits,
rope_type: 1, has_freq_factors: 0,
has_qk_norm: 1,
};
enc.set_compute_pipeline_state(&self.pipes.qk_norm_rope);
enc.set_buffer(0, Some(&self.qkv_buf), 0); enc.set_buffer(1, Some(&self.qkv_buf), (q_dim * 4) as u64); enc.set_buffer(2, Some(&lw.q_norm), 0);
enc.set_buffer(3, Some(&lw.k_norm), 0);
rope_params.bind(enc, &self.rope_freqs_dummy);
enc.dispatch_thread_groups(sz1d((n_head + n_kv) as u64), sz1d(256));
let k_src_off = (q_dim * 4) as u64;
let v_src_off = ((q_dim + kv_dim) * 4) as u64;
let cache_off = (pos * kv_dim * 2) as u64;
let copy_params = ElementwiseParams::new(kv_dim as u32);
enc.set_compute_pipeline_state(&self.pipes.cast_f32_to_f16);
enc.set_buffer(0, Some(&self.qkv_buf), k_src_off);
enc.set_buffer(1, Some(&self.kv_k[il]), cache_off);
copy_params.set(enc, 2);
enc.dispatch_thread_groups(sz1d((kv_dim as u64).div_ceil(256)), sz1d(256));
enc.set_compute_pipeline_state(&self.pipes.cast_f32_to_f16);
enc.set_buffer(0, Some(&self.qkv_buf), v_src_off);
enc.set_buffer(1, Some(&self.kv_v[il]), cache_off);
copy_params.set(enc, 2);
enc.dispatch_thread_groups(sz1d((kv_dim as u64).div_ceil(256)), sz1d(256));
let seq_len = pos + 1;
let attn_params = FlashAttnParams {
n_heads: n_head,
n_kv_heads: n_kv,
head_dim: hd as u32,
kv_dim: kv_dim as u32,
seq_len: seq_len as u32,
scale_bits: scale.to_bits(),
_pad0: 0,
_pad1: 0,
};
enc.set_compute_pipeline_state(&self.pipes.flash_attention);
enc.set_buffer(0, Some(&self.qkv_buf), 0); enc.set_buffer(1, Some(&self.kv_k[il]), 0);
enc.set_buffer(2, Some(&self.kv_v[il]), 0);
enc.set_buffer(3, Some(&self.attn_out_buf), 0);
attn_params.set(enc, 4);
enc.dispatch_thread_groups(sz1d(n_head as u64), sz1d(256));
self.encode_df_gemv_accum(enc, &lw.wo, &self.attn_out_buf, &self.hidden_buf);
enc.set_compute_pipeline_state(&self.pipes.rmsnorm);
enc.set_buffer(0, Some(&self.hidden_buf), 0);
enc.set_buffer(1, Some(&self.normed_buf), 0);
enc.set_buffer(2, Some(&lw.ffn_norm), 0);
enc.set_buffer(3, Some(&self.rmsnorm_params), 0);
enc.dispatch_thread_groups(sz1d(1), sz1d(256));
self.encode_df_gemv(enc, &lw.w1, &self.normed_buf, &self.gate_buf);
self.encode_df_gemv(enc, &lw.w3, &self.normed_buf, &self.up_buf);
enc.set_compute_pipeline_state(&self.pipes.silu_mul_inplace);
enc.set_buffer(0, Some(&self.gate_buf), 0);
enc.set_buffer(1, Some(&self.up_buf), 0);
enc.set_buffer(2, Some(&self.elementwise_is), 0);
enc.dispatch_thread_groups(sz1d((cfg.ffn_dim as u64).div_ceil(256)), sz1d(256));
self.encode_df_gemv_accum(enc, &lw.w2, &self.gate_buf, &self.hidden_buf);
}
self.n_past.store(pos + 1, Ordering::Relaxed);
enc.set_compute_pipeline_state(&self.pipes.rmsnorm);
enc.set_buffer(0, Some(&self.hidden_buf), 0);
enc.set_buffer(1, Some(&self.normed_buf), 0);
enc.set_buffer(2, Some(&self.codebook_norms[j]), 0);
enc.set_buffer(3, Some(&self.rmsnorm_params), 0);
enc.dispatch_thread_groups(sz1d(1), sz1d(256));
self.encode_df_gemv(
enc,
&self.codebook_to_logits[j],
&self.normed_buf,
&self.logits_buf,
);
enc.end_encoding();
cb.commit();
cb.wait_until_completed();
let logits = unsafe {
std::slice::from_raw_parts(self.logits_buf.contents() as *const f32, dec.n_vocab)
};
let sampled = if logits.is_empty() {
0
} else if !temperature.is_finite() || temperature <= 0.0 || top_k <= 1 {
crate::sampler::argmax(logits) as i32
} else {
let mut logits_vec = logits.to_vec();
let inv_temp = 1.0 / temperature;
for l in &mut logits_vec {
*l *= inv_temp;
}
crate::backend::cpu::softmax_inplace(&mut logits_vec);
let mut indices: Vec<usize> = (0..logits_vec.len()).collect();
indices.sort_unstable_by(|&a, &b| logits_vec[b].total_cmp(&logits_vec[a]));
let k = top_k.clamp(1, logits_vec.len());
indices.truncate(k);
let sum: f32 = indices.iter().map(|&i| logits_vec[i]).sum();
let mut r = rand::random::<f32>() * sum;
let mut picked = indices.first().copied().unwrap_or(0);
for &i in &indices {
r -= logits_vec[i];
if r <= 0.0 {
picked = i;
break;
}
}
picked as i32
};
*code = sampled;
prev_token = sampled;
}
codes
}
fn encode_df_gemv(
&self,
enc: &ComputeCommandEncoderRef,
w: &MetalWeight,
input: &Buffer,
output: &Buffer,
) {
enc.set_compute_pipeline_state(&self.pipes.gemv_f32);
enc.set_buffer(0, Some(&w.buf), 0);
enc.set_buffer(1, Some(input), 0);
enc.set_buffer(2, Some(output), 0);
enc.set_buffer(3, Some(&w.params_buf), 0);
enc.dispatch_thread_groups(sz1d(w.m as u64), sz1d(32));
}
fn encode_df_gemv_accum(
&self,
enc: &ComputeCommandEncoderRef,
w: &MetalWeight,
input: &Buffer,
output: &Buffer,
) {
self.encode_df_gemv(enc, w, input, &self.accum_buf);
let n = w.m;
let params = ElementwiseParams::new(n);
enc.set_compute_pipeline_state(&self.pipes.add_inplace);
enc.set_buffer(0, Some(output), 0);
enc.set_buffer(1, Some(&self.accum_buf), 0);
params.set(enc, 2);
enc.dispatch_thread_groups(sz1d((n as u64).div_ceil(256)), sz1d(256));
}
}
impl crate::model::audio_decoder::AudioGpu for MetalAudioDecoder {
fn supports_depthformer(&self) -> bool {
self.depthformer.is_some()
}
fn sample_audio_frame(&self, embedding: &[f32], temperature: f32, top_k: usize) -> [i32; 8] {
let _guard = self.infer_lock.lock().unwrap_or_else(|e| e.into_inner());
match &self.depthformer {
Some(df) => df.sample_frame(embedding, temperature, top_k),
None => {
tracing::warn!("Metal depthformer not loaded, returning empty audio frame");
[0; 8]
}
}
}
fn detokenize_to_spectrum(
&self,
cpu_weights: &crate::model::audio_decoder::DetokenizerWeights,
codes: &[i32],
) -> Vec<f32> {
self.detokenize_to_spectrum(cpu_weights, codes)
}
fn istft_to_pcm(&self, spectrum: &[f32], n_fft: usize, hop_length: usize) -> Vec<f32> {
self.istft_to_pcm(spectrum, n_fft, hop_length)
}
fn reset_depthformer(&self) {
let _guard = self.infer_lock.lock().unwrap_or_else(|e| e.into_inner());
if let Some(df) = &self.depthformer {
df.reset();
}
}
fn reset_detokenizer(&self) {
self.reset();
}
fn try_acquire_session(&self) -> bool {
self.session_active
.compare_exchange(false, true, Ordering::Acquire, Ordering::Relaxed)
.is_ok()
}
fn release_session(&self) {
self.session_active.store(false, Ordering::Release);
}
}