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
+32 -73
View File
@@ -3,8 +3,8 @@ use candle_core::{D, Tensor};
use candle_nn::{
Activation, BatchNorm, BatchNormConfig, Conv1d, Conv1dConfig, Conv2d, Conv2dConfig, Embedding,
LayerNorm, LayerNormConfig, Linear, Module, RmsNorm, VarBuilder, batch_norm, conv1d,
conv1d_no_bias, conv2d, conv2d_no_bias, embedding, layer_norm, linear, linear_no_bias,
rms_norm,
conv1d_no_bias, conv2d, conv2d_no_bias, embedding, layer_norm, linear_b,
linear_no_bias, rms_norm,
};
use rayon::iter::{IndexedParallelIterator, IntoParallelRefIterator, ParallelIterator};
@@ -29,19 +29,9 @@ impl GateUpDownMLP {
act_fn: Activation,
bias: bool,
) -> Result<Self> {
let (gate_proj, up_proj, down_proj) = if bias {
(
linear(hidden_size, intermediate_size, vb.pp("gate_proj"))?,
linear(hidden_size, intermediate_size, vb.pp("up_proj"))?,
linear(intermediate_size, hidden_size, vb.pp("down_proj"))?,
)
} else {
(
linear_no_bias(hidden_size, intermediate_size, vb.pp("gate_proj"))?,
linear_no_bias(hidden_size, intermediate_size, vb.pp("up_proj"))?,
linear_no_bias(intermediate_size, hidden_size, vb.pp("down_proj"))?,
)
};
let gate_proj = linear_b(hidden_size, intermediate_size, bias, vb.pp("gate_proj"))?;
let up_proj = linear_b(hidden_size, intermediate_size, bias, vb.pp("up_proj"))?;
let down_proj = linear_b(intermediate_size, hidden_size, bias, vb.pp("down_proj"))?;
Ok(Self {
gate_proj,
up_proj,
@@ -78,17 +68,9 @@ impl TwoLinearMLP {
linear1_pp_name: &str,
linear2_pp_name: &str,
) -> Result<Self> {
let (linear1, linear2) = if bias {
(
linear(in_dim, middle_dim, vb.pp(linear1_pp_name))?,
linear(middle_dim, out_dim, vb.pp(linear2_pp_name))?,
)
} else {
(
linear_no_bias(in_dim, middle_dim, vb.pp(linear1_pp_name))?,
linear_no_bias(middle_dim, out_dim, vb.pp(linear2_pp_name))?,
)
};
let linear1 = linear_b(in_dim, middle_dim, bias, vb.pp(linear1_pp_name))?;
let linear2 = linear_b(middle_dim, out_dim, bias, vb.pp(linear2_pp_name))?;
Ok(Self {
linear1,
linear2,
@@ -141,53 +123,30 @@ impl NaiveAttention {
let k_proj_pp_name = k_proj_pp_name.unwrap_or("k_proj");
let v_proj_pp_name = v_proj_pp_name.unwrap_or("v_proj");
let o_proj_pp_name = o_proj_pp_name.unwrap_or("o_proj");
let (q_proj, k_proj, v_proj, o_proj) = if bias {
(
linear(
hidden_size,
num_attention_heads * head_dim,
vb.pp(q_proj_pp_name),
)?,
linear(
hidden_size,
num_key_value_heads * head_dim,
vb.pp(k_proj_pp_name),
)?,
linear(
hidden_size,
num_key_value_heads * head_dim,
vb.pp(v_proj_pp_name),
)?,
linear(
num_attention_heads * head_dim,
hidden_size,
vb.pp(o_proj_pp_name),
)?,
)
} else {
(
linear_no_bias(
hidden_size,
num_attention_heads * head_dim,
vb.pp(q_proj_pp_name),
)?,
linear_no_bias(
hidden_size,
num_key_value_heads * head_dim,
vb.pp(k_proj_pp_name),
)?,
linear_no_bias(
hidden_size,
num_key_value_heads * head_dim,
vb.pp(v_proj_pp_name),
)?,
linear_no_bias(
num_attention_heads * head_dim,
hidden_size,
vb.pp(o_proj_pp_name),
)?,
)
};
let q_proj = linear_b(
hidden_size,
num_attention_heads * head_dim,
bias,
vb.pp(q_proj_pp_name),
)?;
let k_proj = linear_b(
hidden_size,
num_key_value_heads * head_dim,
bias,
vb.pp(k_proj_pp_name),
)?;
let v_proj = linear_b(
hidden_size,
num_key_value_heads * head_dim,
bias,
vb.pp(v_proj_pp_name),
)?;
let o_proj = linear_b(
num_attention_heads * head_dim,
hidden_size,
bias,
vb.pp(o_proj_pp_name),
)?;
Ok(Self {
q_proj,
+3 -9
View File
@@ -1,10 +1,7 @@
use anyhow::Result;
use candle_core::{D, IndexOp, Tensor};
use candle_nn::{
Activation, Conv2d, Embedding, Init, LayerNorm, Linear, Module, RmsNorm, VarBuilder, embedding,
linear, linear_no_bias,
ops::{sigmoid, softmax},
rms_norm,
Activation, Conv2d, Embedding, Init, LayerNorm, Linear, Module, RmsNorm, VarBuilder, embedding, linear, linear_b, linear_no_bias, ops::{sigmoid, softmax}, rms_norm
};
use candle_transformers::models::segment_anything::LayerNorm2d;
@@ -79,11 +76,8 @@ impl Attention {
) -> Result<Self> {
let head_dim = dim / num_heads;
let scaling = 1.0 / (head_dim as f64).sqrt();
let qkv = if qkv_bias {
linear(dim, dim * 3, vb.pp("qkv"))?
} else {
linear_no_bias(dim, dim * 3, vb.pp("qkv"))?
};
let qkv = linear_b(dim, dim * 3, qkv_bias, vb.pp("qkv"))?;
let proj = linear(dim, dim, vb.pp("proj"))?;
let mut rel_pos_h = None;
let mut rel_pos_w = None;
+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 {
+26 -18
View File
@@ -1,7 +1,8 @@
use anyhow::Result;
use candle_core::Tensor;
use candle_nn::{
Embedding, Linear, Module, RmsNorm, VarBuilder, embedding, linear, linear_no_bias, rms_norm,
Embedding, Linear, Module, RmsNorm, VarBuilder, embedding, linear_b, linear_no_bias,
rms_norm,
};
use crate::{
@@ -36,23 +37,30 @@ impl Qwen3Attention {
let num_key_value_heads = config.num_key_value_heads;
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 config.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,
config.attention_bias,
vb.pp("q_proj"),
)?;
let k_proj = linear_b(
hidden_size,
num_key_value_heads * head_dim,
config.attention_bias,
vb.pp("k_proj"),
)?;
let v_proj = linear_b(
hidden_size,
num_key_value_heads * head_dim,
config.attention_bias,
vb.pp("v_proj"),
)?;
let o_proj = linear_b(
num_attention_heads * head_dim,
hidden_size,
config.attention_bias,
vb.pp("o_proj"),
)?;
let q_norm = rms_norm(head_dim, config.rms_norm_eps, vb.pp("q_norm"))?;
let k_norm = rms_norm(head_dim, config.rms_norm_eps, vb.pp("k_norm"))?;
Ok(Self {
+3 -6
View File
@@ -2,7 +2,7 @@ use anyhow::{Result, anyhow};
use candle_core::{D, DType, Device, IndexOp, Shape, Tensor};
use candle_nn::{
Activation, BatchNorm, Conv2d, Init, LayerNorm, Linear, Module, ModuleT, VarBuilder, linear,
linear_no_bias, ops::sigmoid,
linear_b, linear_no_bias, ops::sigmoid,
};
use crate::{
@@ -91,11 +91,8 @@ impl WindowAttention {
) -> Result<Self> {
let head_dim = dim / num_heads;
let scaling = 1.0 / (head_dim as f64).sqrt();
let qkv = if qkv_bias {
linear(dim, dim * 3, vb.pp("qkv"))?
} else {
linear_no_bias(dim, dim * 3, vb.pp("qkv"))?
};
let qkv = linear_b(dim, dim * 3, qkv_bias, vb.pp("qkv"))?;
let proj = linear(dim, dim, vb.pp("proj"))?;
let relative_position_bias_table = vb.get_with_hints(
((2 * window_size.0 - 1) * (2 * window_size.1 - 1), num_heads),