update fmt

This commit is contained in:
jhqxxx
2026-03-14 20:06:20 +08:00
parent 0c6d4d01d9
commit 21caee0522
13 changed files with 761 additions and 697 deletions
+4 -2
View File
@@ -492,9 +492,11 @@ impl ImageEncoderViT {
}
pub fn forward(&self, xs: &Tensor) -> Result<Tensor> {
let mut x = self.patch_embed.forward(xs)?;
if self.pos_embed.is_some() {
// if self.pos_embed.is_some() {
if let Some(pos_emb) = &self.pos_embed {
let dim1 = x.dim(1)?;
let pos = self.get_abs_pos_sam(self.pos_embed.as_ref().unwrap(), dim1)?;
// let pos = self.get_abs_pos_sam(self.pos_embed.as_ref().unwrap(), dim1)?;
let pos = self.get_abs_pos_sam(pos_emb, dim1)?;
x = x.broadcast_add(&pos)?;
}
for blk in &self.blocks {
+8 -15
View File
@@ -268,19 +268,15 @@ impl EncoderLayerSANM {
self.self_attn
.forward(&xs, mask, mask_shfit_chunk, mask_att_chunk_encoder)?;
let x_concat = Tensor::cat(&[&xs, &attn], D::Minus1)?;
if self.in_dim == self.hidden_dim {
let x_concat = self
.concat_linear
.as_ref()
.unwrap()
if self.in_dim == self.hidden_dim
&& let Some(concat_linear) = &self.concat_linear
{
let x_concat = concat_linear
.forward(&x_concat)?
.affine(stoch_layer_coeff, 0.0)?;
xs = residual.add(&x_concat)?;
} else {
xs = self
.concat_linear
.as_ref()
.unwrap()
} else if let Some(concat_linear) = &self.concat_linear {
xs = concat_linear
.forward(&x_concat)?
.affine(stoch_layer_coeff, 0.0)?;
}
@@ -496,13 +492,10 @@ impl AdaptorEncoderLayer {
} else {
xs.clone()
};
if self.concat_linear.is_some() {
if let Some(concat_linear) = &self.concat_linear {
let attn = self.self_attn.forward(&xs, None, None, mask, false)?;
let x_concat = Tensor::cat(&[&xs, &attn], D::Minus1)?;
let x_concat = self
.concat_linear
.as_ref()
.unwrap()
let x_concat = concat_linear
.forward(&x_concat)?
.affine(stoch_layer_coeff, 0.0)?;
xs = residual.add(&x_concat)?;
+5 -3
View File
@@ -707,11 +707,13 @@ impl PaddleOCRVLModel {
self.rope_deltas = Some(rope_deltas);
} else {
let (bs, seq_len, _) = inputs_embeds.dims3()?;
let delta = if let Some(cache_position) = cache_position {
let delta = if let Some(cache_position) = cache_position
&& let Some(rope_deltas) = &self.rope_deltas
{
cache_position
.i(0)?
.to_dtype(self.rope_deltas.as_ref().unwrap().dtype())?
.broadcast_add(self.rope_deltas.as_ref().unwrap())?
.to_dtype(rope_deltas.dtype())?
.broadcast_add(rope_deltas)?
.contiguous()?
.to_dtype(candle_core::DType::U32)?
} else {
+5 -3
View File
@@ -1057,11 +1057,13 @@ impl Qwen2_5VLModel {
self.rope_deltas = Some(rope_deltas);
} else {
let (bs, seq_len, _) = inputs_embeds.dims3()?;
let delta = if let Some(cache_position) = cache_position {
let delta = if let Some(cache_position) = cache_position
&& let Some(rope_deltas) = &self.rope_deltas
{
cache_position
.i(0)?
.to_dtype(self.rope_deltas.as_ref().unwrap().dtype())?
.broadcast_add(self.rope_deltas.as_ref().unwrap())?
.to_dtype(rope_deltas.dtype())?
.broadcast_add(rope_deltas)?
.contiguous()?
.to_dtype(candle_core::DType::U32)?
} else {
+8 -8
View File
@@ -1316,26 +1316,26 @@ impl Qwen3_5Model {
video_grid_thw: Option<&Tensor>,
seqlen_offset: usize,
) -> Result<Tensor> {
let position_ids = if self.rope_deltas.is_none() {
let (position_ids, rope_deltas) =
self.get_rope_index(input_ids, image_grid_thw, video_grid_thw, None)?;
self.rope_deltas = Some(rope_deltas);
position_ids
} else {
let position_ids = if let Some(rope_deltas) = &self.rope_deltas {
let (bs, seq_len, _) = inputs_embeds.dims3()?;
Tensor::arange(
seqlen_offset as i64,
(seqlen_offset + seq_len) as i64,
input_ids.device(),
)?
.to_dtype(self.rope_deltas.as_ref().unwrap().dtype())?
.to_dtype(rope_deltas.dtype())?
.unsqueeze(0)?
.broadcast_as((bs, seq_len))?
.broadcast_add(self.rope_deltas.as_ref().unwrap())?
.broadcast_add(rope_deltas)?
.unsqueeze(0)?
.broadcast_as((3, bs, seq_len))?
.contiguous()?
.to_dtype(candle_core::DType::U32)?
} else {
let (position_ids, rope_deltas) =
self.get_rope_index(input_ids, image_grid_thw, video_grid_thw, None)?;
self.rope_deltas = Some(rope_deltas);
position_ids
};
Ok(position_ids)
}
+5 -3
View File
@@ -1013,11 +1013,13 @@ impl Qwen3VLModel {
self.rope_deltas = Some(rope_deltas);
} else {
let (bs, seq_len, _) = inputs_embeds.dims3()?;
let delta = if let Some(cache_position) = cache_position {
let delta = if let Some(cache_position) = cache_position
&& let Some(rope_deltas) = &self.rope_deltas
{
cache_position
.i(0)?
.to_dtype(self.rope_deltas.as_ref().unwrap().dtype())?
.broadcast_add(self.rope_deltas.as_ref().unwrap())?
.to_dtype(rope_deltas.dtype())?
.broadcast_add(rope_deltas)?
.contiguous()?
.to_dtype(candle_core::DType::U32)?
} else {
+3
View File
@@ -9,6 +9,9 @@ use candle_core::{DType, Device, IndexOp, Shape, Tensor};
use ffmpeg_next as ffmpeg;
use image::DynamicImage;
#[cfg(feature = "ffmpeg")]
use anyhow::anyhow;
#[cfg(feature = "ffmpeg")]
use crate::utils::video_utils::video_smart_resize;
use crate::{
+8 -5
View File
@@ -63,10 +63,10 @@ impl PatchEmbed {
xs = xs.pad_with_zeros(2, 0, self.patch_size - h % self.patch_size)?;
}
xs = self.proj.forward(&xs)?;
if self.norm.is_some() {
if let Some(norm) = &self.norm {
let (_, _, ph, pw) = xs.dims4()?;
xs = xs.flatten_from(2)?.transpose(1, 2)?;
xs = self.norm.as_ref().unwrap().forward(&xs)?;
xs = norm.forward(&xs)?;
xs = xs.transpose(1, 2)?.reshape(((), self.embed_dim, ph, pw))?;
}
Ok(xs)
@@ -1018,9 +1018,12 @@ impl BasicDecBlk {
pub fn new(vb: VarBuilder, in_c: usize, out_c: usize) -> Result<Self> {
let inter_channels = 64;
let conv_in = get_conv2d(vb.pp("conv_in"), in_c, inter_channels, 3, 1, 1, 1, 1, true)?;
let dec_att = ASPPDeformable::new(vb.pp("dec_att"), inter_channels, inter_channels, vec![
1, 3, 7,
])?;
let dec_att = ASPPDeformable::new(
vb.pp("dec_att"),
inter_channels,
inter_channels,
vec![1, 3, 7],
)?;
let conv_out = get_conv2d(
vb.pp("conv_out"),
inter_channels,
+3 -3
View File
@@ -33,11 +33,11 @@ impl SinusoidalPositionEncoderCat {
device,
)?
.reshape((seq_len, 1))?; // (seq_len, 1)
let inv_freq = if self.inv_freq.is_none() {
let inv_freq = if let Some(inv_freq) = &self.inv_freq {
inv_freq.clone()
} else {
let inv_freq = compute_default_rope_parameters(head_dim, 10000.0);
Tensor::from_slice(&inv_freq, (1, inv_freq.len()), device)?
} else {
self.inv_freq.as_ref().unwrap().clone()
};
let freqs = positions.matmul(&inv_freq)?; // (seq_len, dim / 2)
let sin = freqs.sin()?;