style: cargo fmt --all (18 files)

Auto-merged by ci-doctor.
This commit is contained in:
redclawsystems
2026-05-07 16:30:04 +00:00
parent 789448f196
commit ae53983c03
18 changed files with 250 additions and 168 deletions
+73 -23
View File
@@ -142,7 +142,11 @@ impl ConvStem {
},
vb.pp("conv3"),
)?;
Ok(Self { conv1, conv2, conv3 })
Ok(Self {
conv1,
conv2,
conv3,
})
}
/// Forward: `(B, 1, T_audio)` raw waveform → `(B, T_seq, 288)` where
@@ -167,9 +171,7 @@ impl ConvStem {
/// Loader: open the HF safetensors and construct a `ConvStem`. Useful
/// for the standalone Phase 8.5 smoke test.
pub fn load_conv_stem(weights_path: &std::path::Path, device: &Device) -> Result<ConvStem> {
let vb = unsafe {
VarBuilder::from_mmaped_safetensors(&[weights_path], DType::F32, device)
}?;
let vb = unsafe { VarBuilder::from_mmaped_safetensors(&[weights_path], DType::F32, device) }?;
ConvStem::new(vb.pp("model").pp("encoder"))
}
@@ -190,7 +192,13 @@ struct RotaryCache {
}
impl RotaryCache {
fn new(rotary_dim: usize, max_seq: usize, theta: f64, dtype: DType, dev: &Device) -> Result<Self> {
fn new(
rotary_dim: usize,
max_seq: usize,
theta: f64,
dtype: DType,
dev: &Device,
) -> Result<Self> {
assert!(rotary_dim.is_multiple_of(2), "rotary_dim must be even");
let inv_freq: Vec<f32> = (0..rotary_dim)
.step_by(2)
@@ -203,7 +211,11 @@ impl RotaryCache {
let freqs = positions.matmul(&inv_freq.reshape((1, rotary_dim / 2))?)?;
let cos = freqs.cos()?.to_dtype(dtype)?;
let sin = freqs.sin()?.to_dtype(dtype)?;
Ok(Self { cos, sin, rotary_dim })
Ok(Self {
cos,
sin,
rotary_dim,
})
}
/// Apply partial RoPE to `q` of shape `(B, H, T, head_dim)`.
@@ -251,7 +263,14 @@ impl EncoderAttention {
let k_proj = candle_nn::linear_no_bias(h, h, vb.pp("k_proj"))?;
let v_proj = candle_nn::linear_no_bias(h, h, vb.pp("v_proj"))?;
let o_proj = candle_nn::linear_no_bias(h, h, vb.pp("o_proj"))?;
Ok(Self { q_proj, k_proj, v_proj, o_proj, n_heads, head_dim })
Ok(Self {
q_proj,
k_proj,
v_proj,
o_proj,
n_heads,
head_dim,
})
}
fn forward(&self, xs: &Tensor, rope: &RotaryCache) -> Result<Tensor> {
@@ -260,9 +279,18 @@ impl EncoderAttention {
let k = self.k_proj.forward(xs)?;
let v = self.v_proj.forward(xs)?;
// (B, T, H) -> (B, H_heads, T, head_dim)
let q = q.reshape((b, t, self.n_heads, self.head_dim))?.transpose(1, 2)?.contiguous()?;
let k = k.reshape((b, t, self.n_heads, self.head_dim))?.transpose(1, 2)?.contiguous()?;
let v = v.reshape((b, t, self.n_heads, self.head_dim))?.transpose(1, 2)?.contiguous()?;
let q = q
.reshape((b, t, self.n_heads, self.head_dim))?
.transpose(1, 2)?
.contiguous()?;
let k = k
.reshape((b, t, self.n_heads, self.head_dim))?
.transpose(1, 2)?
.contiguous()?;
let v = v
.reshape((b, t, self.n_heads, self.head_dim))?
.transpose(1, 2)?
.contiguous()?;
let q = rope.apply(&q)?;
let k = rope.apply(&k)?;
// Collapse (B, H, T, D) -> (B*H, T, D) for the matmul. candle's
@@ -324,7 +352,12 @@ impl EncoderLayer {
let self_attn = EncoderAttention::new(cfg, vb.pp("self_attn"))?;
let post_attn_ln = layer_norm_weight_only(h, 1e-5, vb.pp("post_attention_layernorm"))?;
let mlp = EncoderMlp::new(cfg, vb.pp("mlp"))?;
Ok(Self { input_ln, self_attn, post_attn_ln, mlp })
Ok(Self {
input_ln,
self_attn,
post_attn_ln,
mlp,
})
}
fn forward(&self, xs: &Tensor, rope: &RotaryCache) -> Result<Tensor> {
@@ -369,8 +402,14 @@ impl Encoder {
for i in 0..cfg.encoder_num_hidden_layers {
layers.push(EncoderLayer::new(cfg, layer_vb.pp(i))?);
}
let final_ln = layer_norm_weight_only(cfg.hidden_size, 1e-5, vb.pp("encoder").pp("layer_norm"))?;
Ok(Self { stem, layers, final_ln, rope })
let final_ln =
layer_norm_weight_only(cfg.hidden_size, 1e-5, vb.pp("encoder").pp("layer_norm"))?;
Ok(Self {
stem,
layers,
final_ln,
rope,
})
}
pub fn forward(&self, pcm: &Tensor) -> Result<Tensor> {
@@ -389,9 +428,7 @@ pub fn load_encoder(
device: &Device,
cfg: &MoonshineConfig,
) -> Result<Encoder> {
let vb = unsafe {
VarBuilder::from_mmaped_safetensors(&[weights_path], DType::F32, device)
}?;
let vb = unsafe { VarBuilder::from_mmaped_safetensors(&[weights_path], DType::F32, device) }?;
Encoder::new(cfg, vb.pp("model"))
}
@@ -642,7 +679,11 @@ pub struct Decoder {
impl Decoder {
pub fn new(cfg: &MoonshineConfig, vb: VarBuilder) -> Result<Self> {
let embed = candle_nn::embedding(cfg.vocab_size, cfg.hidden_size, vb.pp("decoder").pp("embed_tokens"))?;
let embed = candle_nn::embedding(
cfg.vocab_size,
cfg.hidden_size,
vb.pp("decoder").pp("embed_tokens"),
)?;
let head_dim = cfg.hidden_size / cfg.decoder_num_attention_heads;
let rotary_dim = ((head_dim as f64 * cfg.partial_rotary_factor) as usize / 2) * 2;
let rope = RotaryCache::new(
@@ -661,7 +702,13 @@ impl Decoder {
// `encoder.layer_norm.weight`). Inspector dump confirmed.
let final_ln = layer_norm_weight_only(cfg.hidden_size, 1e-5, vb.pp("decoder").pp("norm"))?;
let embed_weight = embed.embeddings().clone();
Ok(Self { embed, layers, final_ln, rope, embed_weight })
Ok(Self {
embed,
layers,
final_ln,
rope,
embed_weight,
})
}
/// Forward over `tokens` shape `(B, T)` with encoder output `enc`
@@ -688,9 +735,7 @@ pub fn load_full(
device: &Device,
cfg: &MoonshineConfig,
) -> Result<(Encoder, Decoder)> {
let vb = unsafe {
VarBuilder::from_mmaped_safetensors(&[weights_path], DType::F32, device)
}?;
let vb = unsafe { VarBuilder::from_mmaped_safetensors(&[weights_path], DType::F32, device) }?;
let encoder = Encoder::new(cfg, vb.pp("model"))?;
let decoder = Decoder::new(cfg, vb.pp("model"))?;
Ok((encoder, decoder))
@@ -741,7 +786,9 @@ impl Decoder {
/// Convenience: load the HF tokenizer.json for Moonshine. Caller passes
/// the path returned by hf_hub.
pub fn load_tokenizer(path: &std::path::Path) -> std::result::Result<tokenizers::Tokenizer, Box<dyn std::error::Error + Send + Sync>> {
pub fn load_tokenizer(
path: &std::path::Path,
) -> std::result::Result<tokenizers::Tokenizer, Box<dyn std::error::Error + Send + Sync>> {
tokenizers::Tokenizer::from_file(path)
}
@@ -914,7 +961,10 @@ impl Decoder {
)?;
h = (h + attn_out)?;
let normed = layer.post_attn_ln.forward(&h)?;
let cross_out = layer.cross_attn.forward_step(&normed, &cache.cross_k[i], &cache.cross_v[i])?;
let cross_out =
layer
.cross_attn
.forward_step(&normed, &cache.cross_k[i], &cache.cross_v[i])?;
h = (h + cross_out)?;
let normed = layer.final_ln.forward(&h)?;
let mlp_out = layer.mlp.forward(&normed)?;