update fmt

This commit is contained in:
jhqxxx
2026-01-15 23:20:16 +08:00
parent f4e95e3937
commit 45b53e41cd
3 changed files with 8 additions and 6 deletions
+2 -2
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_b,
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};
+5 -2
View File
@@ -1,7 +1,10 @@
use anyhow::Result;
use candle_core::{D, IndexOp, Tensor};
use candle_nn::{
Activation, Conv2d, Embedding, Init, LayerNorm, Linear, Module, RmsNorm, VarBuilder, embedding, linear, linear_b, 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;
@@ -77,7 +80,7 @@ impl Attention {
let head_dim = dim / num_heads;
let scaling = 1.0 / (head_dim as f64).sqrt();
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;
+1 -2
View File
@@ -1,8 +1,7 @@
use anyhow::Result;
use candle_core::Tensor;
use candle_nn::{
Embedding, Linear, Module, RmsNorm, VarBuilder, embedding, linear_b, linear_no_bias,
rms_norm,
Embedding, Linear, Module, RmsNorm, VarBuilder, embedding, linear_b, linear_no_bias, rms_norm,
};
use crate::{