From f4e95e3937b688fd04f08cf5c18099b2cdda9211 Mon Sep 17 00:00:00 2001 From: jhqxxx <18280426169@163.com> Date: Thu, 15 Jan 2026 23:17:34 +0800 Subject: [PATCH] use linear_b --- src/models/common/mod.rs | 105 ++++++++++--------------------- src/models/deepseek_ocr/model.rs | 12 +--- src/models/hunyuan_ocr/model.rs | 46 ++++++++------ src/models/qwen3/model.rs | 44 +++++++------ src/models/rmbg2_0/model.rs | 9 +-- 5 files changed, 91 insertions(+), 125 deletions(-) diff --git a/src/models/common/mod.rs b/src/models/common/mod.rs index 737ce7e..2b76938 100644 --- a/src/models/common/mod.rs +++ b/src/models/common/mod.rs @@ -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 { - 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 { - 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, diff --git a/src/models/deepseek_ocr/model.rs b/src/models/deepseek_ocr/model.rs index 4352a95..d9c6f0a 100644 --- a/src/models/deepseek_ocr/model.rs +++ b/src/models/deepseek_ocr/model.rs @@ -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 { 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; diff --git a/src/models/hunyuan_ocr/model.rs b/src/models/hunyuan_ocr/model.rs index 67492b9..4df4ffb 100644 --- a/src/models/hunyuan_ocr/model.rs +++ b/src/models/hunyuan_ocr/model.rs @@ -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 { 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 { diff --git a/src/models/qwen3/model.rs b/src/models/qwen3/model.rs index 8292c2a..c182af4 100644 --- a/src/models/qwen3/model.rs +++ b/src/models/qwen3/model.rs @@ -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 { diff --git a/src/models/rmbg2_0/model.rs b/src/models/rmbg2_0/model.rs index 9417750..edf2d1e 100644 --- a/src/models/rmbg2_0/model.rs +++ b/src/models/rmbg2_0/model.rs @@ -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 { 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),