update note

This commit is contained in:
jhqxxx
2025-11-25 17:45:09 +08:00
parent 08a73bfcc5
commit de3861b731
3 changed files with 8 additions and 16 deletions
+4 -14
View File
@@ -151,8 +151,8 @@ impl Attention {
) -> Result<(Tensor, Tensor)> { ) -> Result<(Tensor, Tensor)> {
let (q_h, q_w) = q_size; let (q_h, q_w) = q_size;
let (k_h, k_w) = k_size; let (k_h, k_w) = k_size;
let rh = self.get_rel_pos(q_h, k_h, rel_pos_h)?; // (h, k, dim) let rh = self.get_rel_pos(q_h, k_h, rel_pos_h)?; // (q_h, k_h, dim)
let rw = self.get_rel_pos(q_w, k_w, rel_pos_w)?; // (w, k, dim) let rw = self.get_rel_pos(q_w, k_w, rel_pos_w)?; // (q_w, k_w, dim)
let (b, _, dim) = q.dims3()?; let (b, _, dim) = q.dims3()?;
let r_q = q.reshape((b, q_h, q_w, dim))?.contiguous()?; let r_q = q.reshape((b, q_h, q_w, dim))?.contiguous()?;
let r_q_ = r_q.unsqueeze(D::Minus2)?; // (b, q_h, q_w, 1, dim) let r_q_ = r_q.unsqueeze(D::Minus2)?; // (b, q_h, q_w, 1, dim)
@@ -610,7 +610,8 @@ impl CLIPVisionEmbeddings {
} }
fn get_abs_pos(&self, tgt_size: usize) -> Result<Tensor> { fn get_abs_pos(&self, tgt_size: usize) -> Result<Tensor> {
let abs_pos_new = self.pos_embeds.squeeze(0)?; println!("self.pos_embeds: {:?}", self.pos_embeds);
let abs_pos_new = self.pos_embeds.clone();
let (len, dim) = abs_pos_new.dims2()?; let (len, dim) = abs_pos_new.dims2()?;
let src_size = ((len - 1) as f32).sqrt() as usize; let src_size = ((len - 1) as f32).sqrt() as usize;
let tgt_size = (tgt_size as f32).sqrt() as usize; let tgt_size = (tgt_size as f32).sqrt() as usize;
@@ -991,17 +992,6 @@ impl DeepseekV2MoE {
Ok(final_xs) Ok(final_xs)
} }
// pub fn farward(&self, xs: &Tensor) -> Result<Tensor> {
// let identity = xs.clone();
// let (bs, seq_len, embedding_dim) = xs.dims3()?;
// let (topk_idx, topk_weight) = self.gate.forward(xs)?;
// let xs = xs.reshape((bs * seq_len, embedding_dim))?;
// let xs = self.moe_infer(&xs, &topk_idx, &topk_weight)?;
// let xs = xs.reshape((bs, seq_len, embedding_dim))?;
// let xs_shared_experts = self.shared_experts.forward(&identity)?;
// let xs = xs.add(&xs_shared_experts)?;
// Ok(xs)
// }
} }
impl Module for DeepseekV2MoE { impl Module for DeepseekV2MoE {
+3 -1
View File
@@ -130,6 +130,7 @@ pub fn find_closest_aspect_ratio(
best_ratio_diff = ratio_diff; best_ratio_diff = ratio_diff;
best_ratio = ratio; best_ratio = ratio;
} else if (ratio_diff - best_ratio_diff).abs() < 1e-10 { } else if (ratio_diff - best_ratio_diff).abs() < 1e-10 {
// 当多个候选比例具有相同的宽高比差异时,根据图像的实际面积来选择最优比例。
let target_area = 0.5 * (image_size as f64).powi(2) * (ratio.0 * ratio.1) as f64; let target_area = 0.5 * (image_size as f64).powi(2) * (ratio.0 * ratio.1) as f64;
if area as f64 > target_area { if area as f64 > target_area {
best_ratio = ratio; best_ratio = ratio;
@@ -148,6 +149,7 @@ pub fn dynamic_preprocess(
let orig_width = image.width(); let orig_width = image.width();
let orig_height = image.height(); let orig_height = image.height();
let aspect_ratio = orig_width as f64 / orig_height as f64; let aspect_ratio = orig_width as f64 / orig_height as f64;
// 控制分块数量在2-9之间
let target_ratios = generate_target_ratios_sorted(2, 9); let target_ratios = generate_target_ratios_sorted(2, 9);
let target_aspect_ratio = find_closest_aspect_ratio( let target_aspect_ratio = find_closest_aspect_ratio(
aspect_ratio, aspect_ratio,
@@ -196,7 +198,7 @@ pub fn resize_with_edge_padding(
) -> DynamicImage { ) -> DynamicImage {
// 按图像原比例resize,可能不是输入的宽高 // 按图像原比例resize,可能不是输入的宽高
let mut img = img.resize(width, height, image::imageops::FilterType::CatmullRom); let mut img = img.resize(width, height, image::imageops::FilterType::CatmullRom);
// 使用全0像素填充为输入宽高 // 使用输入像素颜色填充为输入宽高
if img.height() != height || img.width() != width { if img.height() != height || img.width() != width {
let (img_h, img_w) = (img.height(), img.width()); let (img_h, img_w) = (img.height(), img.width());
let img_buffer = img.to_rgb8(); let img_buffer = img.to_rgb8();
+1 -1
View File
@@ -71,7 +71,7 @@ fn deepseekocr_weight() -> Result<()> {
for m in &model_list { for m in &model_list {
let weights = safetensors::load(m, &device)?; let weights = safetensors::load(m, &device)?;
for (key, tensor) in weights.iter() { for (key, tensor) in weights.iter() {
if key.contains("lm_head") { if key.contains("rel_pos_h") {
println!("=== {} === {:?}", key, tensor.shape()); println!("=== {} === {:?}", key, tensor.shape());
} }
// println!("=== {} === {:?}", key, tensor.shape()); // println!("=== {} === {:?}", key, tensor.shape());