style: cargo fmt --all (18 files)
Auto-merged by ci-doctor.
This commit is contained in:
@@ -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)?;
|
||||
|
||||
Reference in New Issue
Block a user