use linear_b

This commit is contained in:
jhqxxx
2026-01-15 23:17:34 +08:00
parent d9b803d27e
commit f4e95e3937
5 changed files with 91 additions and 125 deletions
+27 -19
View File
@@ -1,8 +1,8 @@
use anyhow::{Result, anyhow};
use candle_core::{D, IndexOp, Tensor};
use candle_nn::{
Conv2d, Embedding, Init, Linear, Module, RmsNorm, VarBuilder, embedding, linear,
linear_no_bias, rms_norm,
Conv2d, Embedding, Init, Linear, Module, RmsNorm, VarBuilder, embedding, linear, linear_b,
rms_norm,
};
use crate::{
@@ -284,23 +284,31 @@ impl HunYuanVLAttention {
) -> Result<Self> {
let num_kv_groups = num_attention_heads / num_key_value_heads;
let scaling = 1f64 / f64::sqrt(head_dim as f64);
let (q_proj, k_proj, v_proj, o_proj) = if attention_bias {
let q_proj = linear(hidden_size, num_attention_heads * head_dim, vb.pp("q_proj"))?;
let k_proj = linear(hidden_size, num_key_value_heads * head_dim, vb.pp("k_proj"))?;
let v_proj = linear(hidden_size, num_key_value_heads * head_dim, vb.pp("v_proj"))?;
let o_proj = linear(num_attention_heads * head_dim, hidden_size, vb.pp("o_proj"))?;
(q_proj, k_proj, v_proj, o_proj)
} else {
let q_proj =
linear_no_bias(hidden_size, num_attention_heads * head_dim, vb.pp("q_proj"))?;
let k_proj =
linear_no_bias(hidden_size, num_key_value_heads * head_dim, vb.pp("k_proj"))?;
let v_proj =
linear_no_bias(hidden_size, num_key_value_heads * head_dim, vb.pp("v_proj"))?;
let o_proj =
linear_no_bias(num_attention_heads * head_dim, hidden_size, vb.pp("o_proj"))?;
(q_proj, k_proj, v_proj, o_proj)
};
let q_proj = linear_b(
hidden_size,
num_attention_heads * head_dim,
attention_bias,
vb.pp("q_proj"),
)?;
let k_proj = linear_b(
hidden_size,
num_key_value_heads * head_dim,
attention_bias,
vb.pp("k_proj"),
)?;
let v_proj = linear_b(
hidden_size,
num_key_value_heads * head_dim,
attention_bias,
vb.pp("v_proj"),
)?;
let o_proj = linear_b(
num_attention_heads * head_dim,
hidden_size,
attention_bias,
vb.pp("o_proj"),
)?;
let query_layernorm = rms_norm(head_dim, rms_norm_eps, vb.pp("query_layernorm"))?;
let key_layernorm = rms_norm(head_dim, rms_norm_eps, vb.pp("key_layernorm"))?;
Ok(Self {