add qwen2.5vl model
This commit is contained in:
Generated
+3907
File diff suppressed because it is too large
Load Diff
+18
@@ -7,3 +7,21 @@ license = "Apache-2.0"
|
|||||||
description = "aha"
|
description = "aha"
|
||||||
|
|
||||||
[dependencies]
|
[dependencies]
|
||||||
|
candle-core = { git = "https://github.com/huggingface/candle.git", version = "0.9.1"}
|
||||||
|
candle-nn = { git = "https://github.com/huggingface/candle.git", version = "0.9.1"}
|
||||||
|
candle-transformers = { git = "https://github.com/huggingface/candle.git", version = "0.9.1"}
|
||||||
|
candle-flash-attn = { git = "https://github.com/huggingface/candle.git", optional = true }
|
||||||
|
serde = "1.0.226"
|
||||||
|
serde_json = "1.0.145"
|
||||||
|
anyhow = "1.0.100"
|
||||||
|
ffmpeg-next = "8.0.0"
|
||||||
|
image = "0.25.8"
|
||||||
|
reqwest = { version = "0.12.23", features = ["blocking"] }
|
||||||
|
base64 = "0.22.1"
|
||||||
|
num = "0.4.3"
|
||||||
|
minijinja = "2.12.0"
|
||||||
|
tokenizers = "0.22.1"
|
||||||
|
openai_dive = "1.3.0"
|
||||||
|
[features]
|
||||||
|
flash-attn=["candle-flash-attn"]
|
||||||
|
cuda=["candle-nn/cuda", "candle-core/cuda", "candle-transformers/cuda"]
|
||||||
|
|||||||
Binary file not shown.
|
After Width: | Height: | Size: 62 KiB |
Binary file not shown.
|
After Width: | Height: | Size: 115 KiB |
@@ -0,0 +1,104 @@
|
|||||||
|
use crate::utils::utils::string_to_static_str;
|
||||||
|
use anyhow::{Result, anyhow};
|
||||||
|
use minijinja::{Environment, Value as MiniJinjaValue, context};
|
||||||
|
use openai_dive::v1::resources::chat::ChatCompletionParameters;
|
||||||
|
|
||||||
|
pub fn get_template(path: String) -> Result<String> {
|
||||||
|
let tokenizer_config_file = path.clone() + "/tokenizer_config.json";
|
||||||
|
assert!(
|
||||||
|
std::path::Path::new(&tokenizer_config_file).exists(),
|
||||||
|
"tokenizer_config.json not exists in model path"
|
||||||
|
);
|
||||||
|
let tokenizer_config: serde_json::Value =
|
||||||
|
serde_json::from_slice(&std::fs::read(tokenizer_config_file)?)
|
||||||
|
.map_err(|e| anyhow!(format!("load tokenizer_config file error:{}", e)))?;
|
||||||
|
let chat_template = tokenizer_config["chat_template"]
|
||||||
|
.as_str()
|
||||||
|
.ok_or(anyhow!(format!("chat_template to str error")))?;
|
||||||
|
// 修复模板中的问题行
|
||||||
|
let fixed_template = chat_template
|
||||||
|
.replace(
|
||||||
|
"message.content.startswith('<tool_response>')",
|
||||||
|
"message.content is startingwith('<tool_response>')", // 使用minijinja中的 is startingwith 替换
|
||||||
|
)
|
||||||
|
.replace(
|
||||||
|
"message.content.endswith('</tool_response>')",
|
||||||
|
"message.content is endingwith('</tool_response>')", // 使用minijinja中的 is endingwith 替换
|
||||||
|
)
|
||||||
|
.replace(
|
||||||
|
"content.split('</think>')[0].rstrip('\\n').split('<think>')[-1].lstrip('\\n')",
|
||||||
|
"((content | split('</think>'))[0] | rstrip('\\n') | split('<think>'))[-1] | lstrip('\\n')", // 使用自定义的split, rstrip, lstrip过滤器替换
|
||||||
|
)
|
||||||
|
.replace(
|
||||||
|
"content.split('</think>')[-1].lstrip('\\n')",
|
||||||
|
"(content | split('</think>'))[-1] | lstrip('\\n')", // 使用自定义的过滤器替换
|
||||||
|
)
|
||||||
|
.replace(
|
||||||
|
"reasoning_content.strip('\\n')",
|
||||||
|
"reasoning_content | strip('\\n')", // 使用自定义的过滤器替换
|
||||||
|
)
|
||||||
|
.replace(
|
||||||
|
"content.lstrip('\\n')",
|
||||||
|
"content | lstrip('\\n')", // 使用自定义的过滤器替换
|
||||||
|
);
|
||||||
|
Ok(fixed_template)
|
||||||
|
}
|
||||||
|
|
||||||
|
pub struct ChatTemplate<'a> {
|
||||||
|
env: Environment<'a>,
|
||||||
|
}
|
||||||
|
|
||||||
|
impl<'a> ChatTemplate<'a> {
|
||||||
|
pub fn init(path: &str) -> Result<Self> {
|
||||||
|
let path = path.to_string();
|
||||||
|
assert!(
|
||||||
|
std::path::Path::new(&path).exists(),
|
||||||
|
"model path file not exists"
|
||||||
|
);
|
||||||
|
let template = get_template(path)?;
|
||||||
|
let template = string_to_static_str(template);
|
||||||
|
// 加载jinjaenv处理chat_template
|
||||||
|
let mut env = Environment::new();
|
||||||
|
// 添加自定义过滤器
|
||||||
|
env.add_filter("tojson", |v: MiniJinjaValue| {
|
||||||
|
serde_json::to_string(&v).unwrap()
|
||||||
|
});
|
||||||
|
|
||||||
|
env.add_filter("split", |s: String, delimiter: String| {
|
||||||
|
s.split(&delimiter)
|
||||||
|
.map(|s| s.to_string())
|
||||||
|
.collect::<Vec<String>>()
|
||||||
|
});
|
||||||
|
|
||||||
|
// 添加 lstrip 过滤器
|
||||||
|
env.add_filter("lstrip", |s: String, chars: Option<String>| match chars {
|
||||||
|
Some(chars_str) => s.trim_start_matches(chars_str.as_str()).to_string(),
|
||||||
|
None => s.trim_start().to_string(),
|
||||||
|
});
|
||||||
|
|
||||||
|
// 添加 rstrip 过滤器
|
||||||
|
env.add_filter("rstrip", |s: String, chars: Option<String>| match chars {
|
||||||
|
Some(chars_str) => s.trim_end_matches(chars_str.as_str()).to_string(),
|
||||||
|
None => s.trim_end().to_string(),
|
||||||
|
});
|
||||||
|
// let template = get_template(path.to_string())?;
|
||||||
|
let _ = env.add_template("chat", template);
|
||||||
|
|
||||||
|
Ok(Self { env })
|
||||||
|
}
|
||||||
|
|
||||||
|
pub fn apply_chat_template(&self, messages: &ChatCompletionParameters) -> Result<String> {
|
||||||
|
let context = context! {
|
||||||
|
messages => &messages.messages,
|
||||||
|
add_generation_prompt => true,
|
||||||
|
};
|
||||||
|
let template = self
|
||||||
|
.env
|
||||||
|
.get_template("chat")
|
||||||
|
.map_err(|e| anyhow!(format!("render template error {}", e)))?;
|
||||||
|
let message_str = template
|
||||||
|
.render(context)
|
||||||
|
.map_err(|e| anyhow!(format!("render template error {}", e)))?;
|
||||||
|
Ok(message_str)
|
||||||
|
}
|
||||||
|
}
|
||||||
@@ -0,0 +1 @@
|
|||||||
|
pub mod chat_template;
|
||||||
+28
@@ -0,0 +1,28 @@
|
|||||||
|
use crate::models::{GenerateModel, qwen2_5vl::generate::Qwen2_5VLGenerateModel};
|
||||||
|
use anyhow::{Ok, Result};
|
||||||
|
use candle_core::{DType, Device};
|
||||||
|
pub mod chat_template;
|
||||||
|
pub mod models;
|
||||||
|
pub mod position_embed;
|
||||||
|
pub mod tokenizer;
|
||||||
|
pub mod utils;
|
||||||
|
|
||||||
|
pub enum ModelType {
|
||||||
|
Qwen2_5VL,
|
||||||
|
}
|
||||||
|
|
||||||
|
impl ModelType {
|
||||||
|
pub fn init(
|
||||||
|
model_type: ModelType,
|
||||||
|
model_path: &str,
|
||||||
|
device: Option<&Device>,
|
||||||
|
dtype: Option<DType>,
|
||||||
|
) -> Result<Box<dyn GenerateModel>> {
|
||||||
|
match model_type {
|
||||||
|
ModelType::Qwen2_5VL => {
|
||||||
|
let model = Qwen2_5VLGenerateModel::init(model_path, device, dtype)?;
|
||||||
|
Ok(Box::new(model))
|
||||||
|
}
|
||||||
|
}
|
||||||
|
}
|
||||||
|
}
|
||||||
|
|||||||
@@ -0,0 +1,11 @@
|
|||||||
|
pub mod qwen2_5vl;
|
||||||
|
use anyhow::Result;
|
||||||
|
use candle_core::{DType, Device};
|
||||||
|
use openai_dive::v1::resources::chat::ChatCompletionParameters;
|
||||||
|
|
||||||
|
pub trait GenerateModel {
|
||||||
|
fn init(path: &str, device: Option<&Device>, dtype: Option<DType>) -> Result<Self>
|
||||||
|
where
|
||||||
|
Self: Sized;
|
||||||
|
fn generate(&mut self, mes: ChatCompletionParameters) -> Result<String>;
|
||||||
|
}
|
||||||
@@ -0,0 +1,97 @@
|
|||||||
|
use candle_nn::Activation;
|
||||||
|
|
||||||
|
#[derive(Debug, Clone, PartialEq, serde::Deserialize)]
|
||||||
|
pub struct VisionConfig {
|
||||||
|
pub depth: usize,
|
||||||
|
pub hidden_act: Activation,
|
||||||
|
pub hidden_size: usize,
|
||||||
|
pub intermediate_size: usize,
|
||||||
|
pub num_heads: usize,
|
||||||
|
pub in_chans: usize,
|
||||||
|
pub out_hidden_size: usize,
|
||||||
|
pub patch_size: usize,
|
||||||
|
pub spatial_merge_size: usize,
|
||||||
|
pub spatial_patch_size: usize,
|
||||||
|
pub window_size: usize,
|
||||||
|
pub fullatt_block_indexes: Vec<usize>,
|
||||||
|
pub tokens_per_second: usize,
|
||||||
|
pub temporal_patch_size: usize,
|
||||||
|
}
|
||||||
|
#[derive(Debug, Clone, PartialEq, serde::Deserialize)]
|
||||||
|
pub struct RopeScaling {
|
||||||
|
pub r#type: String,
|
||||||
|
pub mrope_section: Vec<usize>,
|
||||||
|
}
|
||||||
|
|
||||||
|
#[derive(Debug, Clone, PartialEq, serde::Deserialize)]
|
||||||
|
pub struct Config {
|
||||||
|
pub attention_dropout: f32,
|
||||||
|
pub bos_token_id: usize,
|
||||||
|
pub eos_token_id: usize,
|
||||||
|
pub vision_start_token_id: usize,
|
||||||
|
pub vision_end_token_id: usize,
|
||||||
|
pub vision_token_id: usize,
|
||||||
|
pub image_token_id: usize,
|
||||||
|
pub video_token_id: usize,
|
||||||
|
pub hidden_act: Activation,
|
||||||
|
pub hidden_size: usize,
|
||||||
|
pub initializer_range: f32,
|
||||||
|
pub intermediate_size: usize,
|
||||||
|
pub max_position_embeddings: usize,
|
||||||
|
pub max_window_layers: usize,
|
||||||
|
pub num_attention_heads: usize,
|
||||||
|
pub num_hidden_layers: usize,
|
||||||
|
pub num_key_value_heads: usize,
|
||||||
|
pub rms_norm_eps: f64,
|
||||||
|
pub rope_theta: f32,
|
||||||
|
pub sliding_window: usize,
|
||||||
|
pub tie_word_embeddings: bool,
|
||||||
|
pub torch_dtype: String,
|
||||||
|
pub use_sliding_window: bool,
|
||||||
|
pub vision_config: VisionConfig,
|
||||||
|
pub rope_scaling: RopeScaling,
|
||||||
|
pub vocab_size: usize,
|
||||||
|
}
|
||||||
|
|
||||||
|
pub struct VisionSetting {
|
||||||
|
pub image_factor: u32,
|
||||||
|
pub min_pixels: u32,
|
||||||
|
pub max_pixels: u32,
|
||||||
|
pub max_ratio: u32,
|
||||||
|
pub temporal_patch_size: usize,
|
||||||
|
pub patch_size: usize,
|
||||||
|
pub merge_size: usize,
|
||||||
|
pub video_min_pixels: u32,
|
||||||
|
pub video_max_pixels: u32,
|
||||||
|
pub video_total_pixels: u32,
|
||||||
|
pub frame_factor: u32,
|
||||||
|
pub fps: f32,
|
||||||
|
pub fps_min_frames: u32,
|
||||||
|
pub fps_max_frames: u32,
|
||||||
|
pub image_mean: Vec<f32>,
|
||||||
|
pub image_std: Vec<f32>,
|
||||||
|
}
|
||||||
|
|
||||||
|
impl VisionSetting {
|
||||||
|
pub fn default() -> Self {
|
||||||
|
Self {
|
||||||
|
image_factor: 28,
|
||||||
|
min_pixels: 4 * 28 * 28,
|
||||||
|
max_pixels: 16384 * 28 * 28,
|
||||||
|
// max_pixels: 1000 * 28 * 28,
|
||||||
|
max_ratio: 200,
|
||||||
|
temporal_patch_size: 2,
|
||||||
|
patch_size: 14,
|
||||||
|
merge_size: 2,
|
||||||
|
video_min_pixels: 128 * 28 * 28,
|
||||||
|
video_max_pixels: 768 * 28 * 28,
|
||||||
|
video_total_pixels: 24576 * 28 * 28,
|
||||||
|
frame_factor: 2,
|
||||||
|
fps: 2.0,
|
||||||
|
fps_min_frames: 4,
|
||||||
|
fps_max_frames: 768,
|
||||||
|
image_mean: vec![0.48145466_f32, 0.4578275, 0.40821073],
|
||||||
|
image_std: vec![0.26862954, 0.26130258, 0.27577711],
|
||||||
|
}
|
||||||
|
}
|
||||||
|
}
|
||||||
@@ -0,0 +1,137 @@
|
|||||||
|
use crate::models::qwen2_5vl::config::Config;
|
||||||
|
use crate::utils::utils::{find_safetensors_files, get_device, get_dtype};
|
||||||
|
use crate::{
|
||||||
|
chat_template::chat_template::ChatTemplate,
|
||||||
|
models::{
|
||||||
|
GenerateModel,
|
||||||
|
qwen2_5vl::{model::Qwen2_5VLModel, processor::Qwen2_5VLProcessor},
|
||||||
|
},
|
||||||
|
tokenizer::tokenizer::TokenizerModel,
|
||||||
|
};
|
||||||
|
use anyhow::{Result, anyhow};
|
||||||
|
use candle_core::{D, DType, Device, IndexOp, Tensor};
|
||||||
|
use candle_nn::VarBuilder;
|
||||||
|
use candle_transformers::generation::LogitsProcessor;
|
||||||
|
use openai_dive::v1::resources::chat::ChatCompletionParameters;
|
||||||
|
|
||||||
|
pub struct Qwen2_5VLGenerateModel<'a> {
|
||||||
|
chat_template: ChatTemplate<'a>,
|
||||||
|
tokenizer: TokenizerModel,
|
||||||
|
pre_processor: Qwen2_5VLProcessor,
|
||||||
|
qwen2_5_vl: Qwen2_5VLModel,
|
||||||
|
device: Device,
|
||||||
|
dtype: DType,
|
||||||
|
}
|
||||||
|
|
||||||
|
impl<'a> GenerateModel for Qwen2_5VLGenerateModel<'a> {
|
||||||
|
fn init(path: &str, device: Option<&Device>, dtype: Option<DType>) -> Result<Self> {
|
||||||
|
let chat_template = ChatTemplate::init(path)?;
|
||||||
|
let tokenizer = TokenizerModel::init(path)?;
|
||||||
|
let config_path = path.to_string() + "/config.json";
|
||||||
|
let cfg: Config = serde_json::from_slice(&std::fs::read(config_path)?)?;
|
||||||
|
let device = &get_device(device);
|
||||||
|
let cfg_dtype = cfg.torch_dtype.as_str();
|
||||||
|
let dtype = get_dtype(dtype, cfg_dtype);
|
||||||
|
let pre_processor = Qwen2_5VLProcessor::new(device, dtype)?;
|
||||||
|
|
||||||
|
let model_list = find_safetensors_files(&path)?;
|
||||||
|
let vb = unsafe { VarBuilder::from_mmaped_safetensors(&model_list, dtype, device)? };
|
||||||
|
let qwen2_5_vl = Qwen2_5VLModel::new(cfg, vb)?;
|
||||||
|
Ok(Qwen2_5VLGenerateModel {
|
||||||
|
chat_template,
|
||||||
|
tokenizer,
|
||||||
|
pre_processor,
|
||||||
|
qwen2_5_vl,
|
||||||
|
device: device.clone(),
|
||||||
|
dtype: dtype,
|
||||||
|
})
|
||||||
|
}
|
||||||
|
fn generate(&mut self, mes: ChatCompletionParameters) -> Result<String> {
|
||||||
|
let temperature = match mes.temperature {
|
||||||
|
Some(temp) => Some(temp as f64),
|
||||||
|
None => None,
|
||||||
|
};
|
||||||
|
let top_p = match mes.top_p {
|
||||||
|
Some(tp) => Some(tp as f64),
|
||||||
|
None => None,
|
||||||
|
};
|
||||||
|
let mut logit_processor = LogitsProcessor::new(34562, temperature, top_p);
|
||||||
|
let mes_render = self.chat_template.apply_chat_template(&mes)?;
|
||||||
|
let input = self.pre_processor.process_info(&mes, &mes_render)?;
|
||||||
|
let mut input_ids = self
|
||||||
|
.tokenizer
|
||||||
|
.text_encode(input.replace_text.clone(), &self.device)?;
|
||||||
|
let mut seq_len = input_ids.dim(1)?;
|
||||||
|
let mut seqlen_offset = 0;
|
||||||
|
let end_of_text_id = self.qwen2_5_vl.cfg.bos_token_id as u32;
|
||||||
|
let im_end_id = self.qwen2_5_vl.cfg.eos_token_id as u32;
|
||||||
|
let mut pixel_values = if input.pixel_values.is_some() {
|
||||||
|
Some(&input.pixel_values.unwrap().clone())
|
||||||
|
} else {
|
||||||
|
None
|
||||||
|
};
|
||||||
|
let image_grid_thw = if input.image_grid_thw.is_some() {
|
||||||
|
Some(&input.image_grid_thw.unwrap().clone())
|
||||||
|
} else {
|
||||||
|
None
|
||||||
|
};
|
||||||
|
let mut pixel_values_video = if input.pixel_values_video.is_some() {
|
||||||
|
Some(&input.pixel_values_video.unwrap().clone())
|
||||||
|
} else {
|
||||||
|
None
|
||||||
|
};
|
||||||
|
let video_grid_thw = if input.video_grid_thw.is_some() {
|
||||||
|
Some(&input.video_grid_thw.unwrap().clone())
|
||||||
|
} else {
|
||||||
|
None
|
||||||
|
};
|
||||||
|
let second_per_grid_ts = if input.second_per_grid_ts.is_some() {
|
||||||
|
Some(input.second_per_grid_ts.unwrap().clone())
|
||||||
|
} else {
|
||||||
|
None
|
||||||
|
};
|
||||||
|
|
||||||
|
let mut mask = Tensor::ones_like(&input_ids)?;
|
||||||
|
let mut cache_position = Tensor::ones_like(&input_ids.i(0)?)?
|
||||||
|
.to_dtype(candle_core::DType::F64)?
|
||||||
|
.cumsum(D::Minus1)?
|
||||||
|
.to_dtype(candle_core::DType::U32)?
|
||||||
|
.broadcast_sub(&Tensor::new(vec![1_u32], input_ids.device())?)?;
|
||||||
|
|
||||||
|
let mut generate = Vec::new();
|
||||||
|
let sample_len = match mes.max_tokens {
|
||||||
|
Some(max) => max,
|
||||||
|
None => 512,
|
||||||
|
};
|
||||||
|
for _ in 0..sample_len {
|
||||||
|
let logits = self.qwen2_5_vl.forward(
|
||||||
|
&input_ids,
|
||||||
|
pixel_values,
|
||||||
|
image_grid_thw,
|
||||||
|
pixel_values_video,
|
||||||
|
video_grid_thw,
|
||||||
|
&mask,
|
||||||
|
Some(&cache_position),
|
||||||
|
seqlen_offset,
|
||||||
|
second_per_grid_ts.clone(),
|
||||||
|
)?;
|
||||||
|
let logits = logits.squeeze(0)?.squeeze(0)?.to_dtype(DType::F32)?;
|
||||||
|
let next_token = logit_processor.sample(&logits)?;
|
||||||
|
generate.push(next_token);
|
||||||
|
if next_token == end_of_text_id || next_token == im_end_id {
|
||||||
|
break;
|
||||||
|
}
|
||||||
|
seqlen_offset += seq_len;
|
||||||
|
seq_len = 1;
|
||||||
|
input_ids = Tensor::from_vec(vec![next_token], (1, 1), &self.device)?;
|
||||||
|
let appendd_mask = Tensor::ones((1, 1), mask.dtype(), &self.device)?;
|
||||||
|
mask = Tensor::cat(&[mask, appendd_mask], 1)?;
|
||||||
|
cache_position = Tensor::from_vec(vec![seqlen_offset as u32], 1, &self.device)?;
|
||||||
|
pixel_values = None;
|
||||||
|
pixel_values_video = None;
|
||||||
|
}
|
||||||
|
let res = self.tokenizer.token_decode(generate)?;
|
||||||
|
self.qwen2_5_vl.clear_kv_cache();
|
||||||
|
Ok(res)
|
||||||
|
}
|
||||||
|
}
|
||||||
@@ -0,0 +1,4 @@
|
|||||||
|
pub mod config;
|
||||||
|
pub mod model;
|
||||||
|
pub mod processor;
|
||||||
|
pub mod generate;
|
||||||
File diff suppressed because it is too large
Load Diff
@@ -0,0 +1,462 @@
|
|||||||
|
use std::collections::HashMap;
|
||||||
|
|
||||||
|
use anyhow::{Result, anyhow};
|
||||||
|
use candle_core::{DType, Device, IndexOp, Shape, Tensor};
|
||||||
|
use ffmpeg_next as ffmpeg;
|
||||||
|
use image::DynamicImage;
|
||||||
|
use num::integer::lcm;
|
||||||
|
use openai_dive::v1::resources::chat::{
|
||||||
|
ChatCompletionParameters, ChatMessage, ChatMessageContent, ChatMessageContentPart,
|
||||||
|
};
|
||||||
|
|
||||||
|
use crate::{
|
||||||
|
models::qwen2_5vl::config::VisionSetting,
|
||||||
|
utils::{
|
||||||
|
img_utils::get_image,
|
||||||
|
utils::{ceil_by_factor, floor_by_factor, round_by_factor},
|
||||||
|
},
|
||||||
|
};
|
||||||
|
|
||||||
|
#[derive(Clone)]
|
||||||
|
pub struct VisionInput {
|
||||||
|
pub data: Tensor,
|
||||||
|
pub grid_thw: Tensor,
|
||||||
|
}
|
||||||
|
|
||||||
|
#[derive(Clone)]
|
||||||
|
pub struct GeneralInput {
|
||||||
|
pub replace_text: String,
|
||||||
|
pub pixel_values: Option<Tensor>,
|
||||||
|
pub image_grid_thw: Option<Tensor>,
|
||||||
|
pub pixel_values_video: Option<Tensor>,
|
||||||
|
pub video_grid_thw: Option<Tensor>,
|
||||||
|
pub second_per_grid_ts: Option<Vec<f32>>,
|
||||||
|
}
|
||||||
|
|
||||||
|
pub struct Qwen2_5VLProcessor {
|
||||||
|
vision_setting: VisionSetting,
|
||||||
|
device: Device,
|
||||||
|
dtype: DType,
|
||||||
|
image_token: String,
|
||||||
|
video_token: String,
|
||||||
|
}
|
||||||
|
|
||||||
|
impl Qwen2_5VLProcessor {
|
||||||
|
pub fn new(device: &Device, dtype: DType) -> Result<Self> {
|
||||||
|
let vision_setting = VisionSetting::default();
|
||||||
|
let image_token = "<|image_pad|>".to_string();
|
||||||
|
let video_token = "<|video_pad|>".to_string();
|
||||||
|
Ok(Self {
|
||||||
|
vision_setting,
|
||||||
|
device: device.clone(),
|
||||||
|
dtype,
|
||||||
|
image_token,
|
||||||
|
video_token,
|
||||||
|
})
|
||||||
|
}
|
||||||
|
|
||||||
|
pub fn extract_vision_info(
|
||||||
|
&self,
|
||||||
|
mes: &ChatCompletionParameters,
|
||||||
|
) -> Result<HashMap<String, Vec<String>>> {
|
||||||
|
let mut vision_map = HashMap::new();
|
||||||
|
vision_map.insert("image".to_string(), Vec::new());
|
||||||
|
vision_map.insert("video".to_string(), Vec::new());
|
||||||
|
for chat_mes in mes.messages.clone() {
|
||||||
|
match chat_mes {
|
||||||
|
ChatMessage::User { content, name } => match content {
|
||||||
|
ChatMessageContent::ContentPart(part_vec) => {
|
||||||
|
for part in part_vec {
|
||||||
|
match part {
|
||||||
|
ChatMessageContentPart::Image(img_part) => {
|
||||||
|
let img_url = img_part.image_url;
|
||||||
|
vision_map.get_mut("image").unwrap().push(img_url.url);
|
||||||
|
}
|
||||||
|
_ => {}
|
||||||
|
}
|
||||||
|
}
|
||||||
|
}
|
||||||
|
_ => {}
|
||||||
|
},
|
||||||
|
_ => {}
|
||||||
|
}
|
||||||
|
}
|
||||||
|
Ok(vision_map)
|
||||||
|
}
|
||||||
|
|
||||||
|
pub fn process_img(
|
||||||
|
&self,
|
||||||
|
img: &DynamicImage,
|
||||||
|
img_mean: &Tensor,
|
||||||
|
img_std: &Tensor,
|
||||||
|
) -> Result<Tensor> {
|
||||||
|
let img_h = img.height();
|
||||||
|
let img_w = img.width();
|
||||||
|
// h,w resize成 28的倍数
|
||||||
|
let (resize_h, resize_w) = smart_resize(img_h, img_w, &self.vision_setting, true, None)?;
|
||||||
|
let img = img.resize_exact(resize_w, resize_h, image::imageops::FilterType::CatmullRom);
|
||||||
|
let img_vec = img.to_rgb8().into_raw();
|
||||||
|
// (h, w, c) => (c, h, w)
|
||||||
|
let img_tensor = Tensor::from_slice(
|
||||||
|
&img_vec,
|
||||||
|
(resize_h as usize, resize_w as usize, 3),
|
||||||
|
&self.device,
|
||||||
|
)?
|
||||||
|
.permute((2, 0, 1))?
|
||||||
|
.to_dtype(self.dtype)?;
|
||||||
|
// 0-255 rescale to 0-1
|
||||||
|
let img_tensor = img_tensor.affine(1.0 / 255.0, 0.)?;
|
||||||
|
// normalize
|
||||||
|
let img_tensor = img_tensor
|
||||||
|
.broadcast_sub(&img_mean)?
|
||||||
|
.broadcast_div(&img_std)?;
|
||||||
|
// (c, h, w) => (1, c, h, w)
|
||||||
|
let img_tensor = img_tensor.unsqueeze(0)?;
|
||||||
|
Ok(img_tensor)
|
||||||
|
}
|
||||||
|
|
||||||
|
pub fn process_vision_tensor(&self, img_tensor: &Tensor) -> Result<(Tensor, Tensor)> {
|
||||||
|
let channel = img_tensor.dim(1)?;
|
||||||
|
let grid_t = img_tensor.dim(0)? / self.vision_setting.temporal_patch_size;
|
||||||
|
let grid_h = img_tensor.dim(2)? / self.vision_setting.patch_size;
|
||||||
|
let grid_w = img_tensor.dim(3)? / self.vision_setting.patch_size;
|
||||||
|
let shape = Shape::from(vec![
|
||||||
|
grid_t,
|
||||||
|
self.vision_setting.temporal_patch_size,
|
||||||
|
channel,
|
||||||
|
grid_h / self.vision_setting.merge_size,
|
||||||
|
self.vision_setting.merge_size,
|
||||||
|
self.vision_setting.patch_size,
|
||||||
|
grid_w / self.vision_setting.merge_size,
|
||||||
|
self.vision_setting.merge_size,
|
||||||
|
self.vision_setting.patch_size,
|
||||||
|
]);
|
||||||
|
let img_tensor = img_tensor.reshape(shape)?;
|
||||||
|
// shape to // grid_t,
|
||||||
|
// grid_h / merge_size,
|
||||||
|
// grid_w / merge_size,
|
||||||
|
// merge_size,
|
||||||
|
// merge_size,
|
||||||
|
// channel,
|
||||||
|
// temporal_patch_size,
|
||||||
|
// patch_size,
|
||||||
|
// patch_size,
|
||||||
|
let img_tensor = img_tensor.permute(vec![0, 3, 6, 4, 7, 2, 1, 5, 8])?;
|
||||||
|
let img_tensor = img_tensor
|
||||||
|
.reshape((
|
||||||
|
grid_t * grid_h * grid_w,
|
||||||
|
channel
|
||||||
|
* self.vision_setting.temporal_patch_size
|
||||||
|
* self.vision_setting.patch_size
|
||||||
|
* self.vision_setting.patch_size,
|
||||||
|
))?
|
||||||
|
.contiguous()?;
|
||||||
|
let grid_thw = Tensor::from_vec(
|
||||||
|
vec![grid_t as u32, grid_h as u32, grid_w as u32],
|
||||||
|
(1, 3),
|
||||||
|
&self.device,
|
||||||
|
)?;
|
||||||
|
Ok((img_tensor, grid_thw))
|
||||||
|
}
|
||||||
|
|
||||||
|
pub fn process_images(
|
||||||
|
&self,
|
||||||
|
imgs: Vec<DynamicImage>,
|
||||||
|
img_mean: &Tensor,
|
||||||
|
img_std: &Tensor,
|
||||||
|
) -> Result<VisionInput> {
|
||||||
|
let mut pixel_values_vec = Vec::new();
|
||||||
|
let mut vision_grid_thws_vec = Vec::new();
|
||||||
|
|
||||||
|
for img in imgs {
|
||||||
|
let img_tensor = self.process_img(&img, &img_mean, &img_std)?;
|
||||||
|
let img_tensor = Tensor::cat(&[&img_tensor, &img_tensor], 0)?.contiguous()?;
|
||||||
|
let (img_tensor, grid_thw) = self.process_vision_tensor(&img_tensor)?;
|
||||||
|
pixel_values_vec.push(img_tensor);
|
||||||
|
vision_grid_thws_vec.push(grid_thw);
|
||||||
|
}
|
||||||
|
let pixel_values = Tensor::cat(&pixel_values_vec, 0)?;
|
||||||
|
let vision_grid_thws = Tensor::cat(&vision_grid_thws_vec, 0)?;
|
||||||
|
Ok(VisionInput {
|
||||||
|
data: pixel_values,
|
||||||
|
grid_thw: vision_grid_thws,
|
||||||
|
})
|
||||||
|
}
|
||||||
|
|
||||||
|
pub fn process_videos(
|
||||||
|
&self,
|
||||||
|
data: Vec<Tensor>,
|
||||||
|
img_mean: &Tensor,
|
||||||
|
img_std: &Tensor,
|
||||||
|
) -> Result<VisionInput> {
|
||||||
|
let mut pixel_values_vec = Vec::new();
|
||||||
|
let mut vision_grid_thws_vec = Vec::new();
|
||||||
|
for single_video in data {
|
||||||
|
// 0-255 rescale to 0-1
|
||||||
|
let video_tensor = single_video.to_dtype(self.dtype)?.affine(1.0 / 255.0, 0.)?;
|
||||||
|
// normalize
|
||||||
|
let video_tensor = video_tensor
|
||||||
|
.broadcast_sub(&img_mean)?
|
||||||
|
.broadcast_div(&img_std)?
|
||||||
|
.contiguous()?;
|
||||||
|
let (video_tensor, video_grid_thw) = self.process_vision_tensor(&video_tensor)?;
|
||||||
|
pixel_values_vec.push(video_tensor);
|
||||||
|
vision_grid_thws_vec.push(video_grid_thw);
|
||||||
|
}
|
||||||
|
let pixel_values = Tensor::cat(&pixel_values_vec, 0)?.contiguous()?;
|
||||||
|
let vision_grid_thws = Tensor::cat(&vision_grid_thws_vec, 0)?.contiguous()?;
|
||||||
|
Ok(VisionInput {
|
||||||
|
data: pixel_values,
|
||||||
|
grid_thw: vision_grid_thws,
|
||||||
|
})
|
||||||
|
}
|
||||||
|
|
||||||
|
pub fn process_info(
|
||||||
|
&self,
|
||||||
|
messages: &ChatCompletionParameters,
|
||||||
|
text: &str,
|
||||||
|
) -> Result<GeneralInput> {
|
||||||
|
let mut pixel_values = None;
|
||||||
|
let mut image_grid_thw = None;
|
||||||
|
let mut pixel_values_video = None;
|
||||||
|
let mut video_grid_thw = None;
|
||||||
|
let mut second_per_grid_ts = None;
|
||||||
|
let vision_map = self.extract_vision_info(messages)?;
|
||||||
|
let img_mean =
|
||||||
|
Tensor::from_slice(&self.vision_setting.image_mean, (3, 1, 1), &self.device)?
|
||||||
|
.to_dtype(self.dtype)?;
|
||||||
|
let img_std = Tensor::from_slice(&self.vision_setting.image_std, (3, 1, 1), &self.device)?
|
||||||
|
.to_dtype(self.dtype)?;
|
||||||
|
for (key, vec) in vision_map {
|
||||||
|
// println!("key: {}, \nvalue: {:?}", key, vec);
|
||||||
|
if key.eq("image") {
|
||||||
|
let mut file_vec = Vec::new();
|
||||||
|
for file in &vec {
|
||||||
|
let image = get_image(file);
|
||||||
|
match image {
|
||||||
|
Ok(img) => file_vec.push(img),
|
||||||
|
Err(e) => println!("get_image err: {:?}", e),
|
||||||
|
};
|
||||||
|
}
|
||||||
|
if file_vec.len() > 0 {
|
||||||
|
let vision_input = self.process_images(file_vec, &img_mean, &img_std);
|
||||||
|
match vision_input {
|
||||||
|
Ok(img_input) => {
|
||||||
|
pixel_values = Some(img_input.data);
|
||||||
|
image_grid_thw = Some(img_input.grid_thw);
|
||||||
|
}
|
||||||
|
Err(e) => println!("img process_images err: {:?}", e),
|
||||||
|
};
|
||||||
|
}
|
||||||
|
}
|
||||||
|
if key.eq("video") {
|
||||||
|
let mut file_vec = Vec::new();
|
||||||
|
for file in &vec {
|
||||||
|
let video_data = get_video_data(file, &self.vision_setting, &self.device);
|
||||||
|
match video_data {
|
||||||
|
Ok(tensor) => file_vec.push(tensor),
|
||||||
|
Err(e) => println!("get_video_data err: {:?}", e),
|
||||||
|
};
|
||||||
|
}
|
||||||
|
if file_vec.len() > 0 {
|
||||||
|
let vision_input = self.process_videos(file_vec, &img_mean, &img_std);
|
||||||
|
match vision_input {
|
||||||
|
Ok(video_input) => {
|
||||||
|
let video_num = video_input.grid_thw.dim(0)?;
|
||||||
|
pixel_values_video = Some(video_input.data);
|
||||||
|
video_grid_thw = Some(video_input.grid_thw);
|
||||||
|
let second_per_grid = vec![
|
||||||
|
self.vision_setting.temporal_patch_size
|
||||||
|
as f32
|
||||||
|
/ self.vision_setting.fps;
|
||||||
|
video_num
|
||||||
|
];
|
||||||
|
second_per_grid_ts = Some(second_per_grid);
|
||||||
|
}
|
||||||
|
Err(e) => println!("video process_videos err: {:?}", e),
|
||||||
|
};
|
||||||
|
}
|
||||||
|
}
|
||||||
|
}
|
||||||
|
let merge_length = self.vision_setting.merge_size.pow(2);
|
||||||
|
let mut text = text.to_string();
|
||||||
|
if image_grid_thw.is_some() {
|
||||||
|
let mut index = 0;
|
||||||
|
while text.contains(&self.image_token) {
|
||||||
|
let grid_i = image_grid_thw.as_ref().unwrap().i(index)?;
|
||||||
|
let repeat_num =
|
||||||
|
grid_i.to_vec1::<u32>()?.iter().product::<u32>() as usize / merge_length;
|
||||||
|
let replace = "<|placeholder|>".repeat(repeat_num);
|
||||||
|
text = text.replacen(&self.image_token, &replace, 1);
|
||||||
|
index += 1;
|
||||||
|
}
|
||||||
|
text = text.replace("<|placeholder|>", &self.image_token);
|
||||||
|
}
|
||||||
|
if video_grid_thw.is_some() {
|
||||||
|
let mut index = 0;
|
||||||
|
while text.contains(&self.video_token) {
|
||||||
|
let grid_i = video_grid_thw.as_ref().unwrap().i(index)?;
|
||||||
|
let repeat_num =
|
||||||
|
grid_i.to_vec1::<u32>()?.iter().product::<u32>() as usize / merge_length;
|
||||||
|
let replace = "<|placeholder|>".repeat(repeat_num);
|
||||||
|
text = text.replacen(&self.video_token, &replace, 1);
|
||||||
|
index += 1;
|
||||||
|
}
|
||||||
|
text = text.replace("<|placeholder|>", &self.video_token);
|
||||||
|
}
|
||||||
|
let input = GeneralInput {
|
||||||
|
replace_text: text,
|
||||||
|
pixel_values,
|
||||||
|
image_grid_thw,
|
||||||
|
pixel_values_video,
|
||||||
|
video_grid_thw,
|
||||||
|
second_per_grid_ts,
|
||||||
|
};
|
||||||
|
Ok(input)
|
||||||
|
}
|
||||||
|
}
|
||||||
|
|
||||||
|
pub fn smart_resize(
|
||||||
|
img_h: u32,
|
||||||
|
img_w: u32,
|
||||||
|
vision_setting: &VisionSetting,
|
||||||
|
is_img: bool,
|
||||||
|
video_ratio: Option<u32>,
|
||||||
|
) -> Result<(u32, u32)> {
|
||||||
|
if std::cmp::max(img_h, img_w) / std::cmp::min(img_h, img_w) > vision_setting.max_ratio {
|
||||||
|
return Err(anyhow!(format!(
|
||||||
|
"absolute aspect ratio mush be smaller than {}, got {}",
|
||||||
|
vision_setting.max_ratio,
|
||||||
|
std::cmp::max(img_h, img_w) / std::cmp::min(img_h, img_w)
|
||||||
|
)));
|
||||||
|
}
|
||||||
|
let mut image_factor = vision_setting.image_factor;
|
||||||
|
if let Some(ratio) = video_ratio {
|
||||||
|
image_factor = lcm(image_factor, ratio);
|
||||||
|
}
|
||||||
|
let mut h_bar = std::cmp::max(image_factor, round_by_factor(img_h, image_factor));
|
||||||
|
let mut w_bar = std::cmp::max(image_factor, round_by_factor(img_w, image_factor));
|
||||||
|
let mut max_pixels = 0u32;
|
||||||
|
let mut min_pixels = 0u32;
|
||||||
|
if is_img {
|
||||||
|
min_pixels = vision_setting.min_pixels;
|
||||||
|
max_pixels = vision_setting.max_pixels;
|
||||||
|
} else {
|
||||||
|
min_pixels = vision_setting.video_min_pixels;
|
||||||
|
max_pixels = vision_setting.video_max_pixels;
|
||||||
|
}
|
||||||
|
if h_bar * w_bar > max_pixels {
|
||||||
|
let beta = ((img_h * img_w) as f32 / max_pixels as f32).sqrt();
|
||||||
|
h_bar = floor_by_factor(img_h as f32 / beta, image_factor);
|
||||||
|
w_bar = floor_by_factor(img_w as f32 / beta, image_factor);
|
||||||
|
} else if h_bar * w_bar < min_pixels {
|
||||||
|
let beta = (min_pixels as f32 / (img_h * img_w) as f32).sqrt();
|
||||||
|
h_bar = ceil_by_factor(img_h as f32 * beta, image_factor);
|
||||||
|
w_bar = ceil_by_factor(img_w as f32 * beta, image_factor);
|
||||||
|
}
|
||||||
|
Ok((h_bar, w_bar))
|
||||||
|
}
|
||||||
|
|
||||||
|
pub fn get_video_data(
|
||||||
|
file: &String,
|
||||||
|
vision_setting: &VisionSetting,
|
||||||
|
device: &Device,
|
||||||
|
) -> Result<Tensor> {
|
||||||
|
ffmpeg::init().map_err(|e| anyhow!(format!("Failed to initialize ffmpeg: {}", e)))?;
|
||||||
|
|
||||||
|
let mut ictx = ffmpeg::format::input(&file)
|
||||||
|
.map_err(|e| anyhow!(format!("Failed to open video file: {}", e)))?;
|
||||||
|
let input = ictx
|
||||||
|
.streams()
|
||||||
|
.best(ffmpeg::media::Type::Video)
|
||||||
|
.ok_or_else(|| anyhow!(format!("No video stream found")))?;
|
||||||
|
let video_stream_index = input.index();
|
||||||
|
let context_decoder = ffmpeg::codec::context::Context::from_parameters(input.parameters())
|
||||||
|
.map_err(|e| anyhow!(format!("Failed to crate decoder context: {}", e)))?;
|
||||||
|
let mut decoder = context_decoder
|
||||||
|
.decoder()
|
||||||
|
.video()
|
||||||
|
.map_err(|e| anyhow!(format!("Failed to decoder video: {}", e)))?;
|
||||||
|
|
||||||
|
let video_h = decoder.height();
|
||||||
|
let video_w = decoder.width();
|
||||||
|
let format = decoder.format();
|
||||||
|
|
||||||
|
let frames = input.frames();
|
||||||
|
let rate = (input.rate().0 as f32 / input.rate().1 as f32).round() as u32;
|
||||||
|
// 1s取两帧
|
||||||
|
let min_frames = ceil_by_factor(
|
||||||
|
vision_setting.fps_min_frames as f32,
|
||||||
|
vision_setting.frame_factor,
|
||||||
|
);
|
||||||
|
let max_frames = floor_by_factor(
|
||||||
|
vision_setting.fps_max_frames as f32,
|
||||||
|
vision_setting.frame_factor,
|
||||||
|
);
|
||||||
|
let nframes = (frames as f32 / rate as f32 * vision_setting.fps) as u32;
|
||||||
|
let nframes = std::cmp::min(std::cmp::max(nframes, min_frames), max_frames);
|
||||||
|
let nframes = round_by_factor(nframes, vision_setting.frame_factor);
|
||||||
|
let sample_interval = (frames as f32 / nframes as f32).round() as u32;
|
||||||
|
let mut frame_id = 0_u32;
|
||||||
|
|
||||||
|
// 图片帧使用scaler reshape的时候需要保证宽高是16的倍数,不然reshape出来的是损坏的图片
|
||||||
|
// 所以计算resize的目标宽高时,需要用16和image_factor的最小公倍数
|
||||||
|
let (resize_h, resize_w) = smart_resize(video_h, video_w, vision_setting, false, Some(16))?;
|
||||||
|
let mut scaler = ffmpeg::software::scaling::context::Context::get(
|
||||||
|
format,
|
||||||
|
video_w,
|
||||||
|
video_h,
|
||||||
|
ffmpeg::format::Pixel::RGB24,
|
||||||
|
resize_w,
|
||||||
|
resize_h,
|
||||||
|
ffmpeg::software::scaling::flag::Flags::BILINEAR
|
||||||
|
| ffmpeg::software::scaling::flag::Flags::ACCURATE_RND,
|
||||||
|
)
|
||||||
|
.map_err(|e| anyhow!(format!("Failed to crate scaler: {}", e)))?;
|
||||||
|
|
||||||
|
let mut frames_vec = Vec::new();
|
||||||
|
let mut receive_and_process_decoded_frames =
|
||||||
|
|decoder: &mut ffmpeg::decoder::Video| -> Result<()> {
|
||||||
|
let mut decoded = ffmpeg::frame::Video::empty();
|
||||||
|
while decoder.receive_frame(&mut decoded).is_ok() {
|
||||||
|
if frame_id % sample_interval == 0 {
|
||||||
|
let mut rgb_frame = ffmpeg::frame::Video::empty();
|
||||||
|
scaler
|
||||||
|
.run(&decoded, &mut rgb_frame)
|
||||||
|
.map_err(|e| anyhow!(format!("Failed to scaler run decoded: {}", e)))?;
|
||||||
|
|
||||||
|
// save_file(&rgb_frame, frame_id as usize);
|
||||||
|
let frame_data = rgb_frame.data(0);
|
||||||
|
let frame_tensor = Tensor::from_slice(
|
||||||
|
frame_data,
|
||||||
|
(resize_h as usize, resize_w as usize, 3),
|
||||||
|
device,
|
||||||
|
)?
|
||||||
|
.permute((2, 0, 1))?;
|
||||||
|
frames_vec.push(frame_tensor);
|
||||||
|
}
|
||||||
|
frame_id += 1;
|
||||||
|
}
|
||||||
|
Ok(())
|
||||||
|
};
|
||||||
|
|
||||||
|
for (stream, packet) in ictx.packets() {
|
||||||
|
if stream.index() == video_stream_index {
|
||||||
|
decoder
|
||||||
|
.send_packet(&packet)
|
||||||
|
.map_err(|e| anyhow!(format!("Failed to send packet: {}", e)))?;
|
||||||
|
receive_and_process_decoded_frames(&mut decoder)?;
|
||||||
|
}
|
||||||
|
}
|
||||||
|
decoder
|
||||||
|
.send_eof()
|
||||||
|
.map_err(|e| anyhow!(format!("Failed to decoder.send_eof(): {}", e)))?;
|
||||||
|
receive_and_process_decoded_frames(&mut decoder)?;
|
||||||
|
|
||||||
|
if frames_vec.is_empty() {
|
||||||
|
return Err(anyhow!("No frames extracted from video".to_string()));
|
||||||
|
}
|
||||||
|
// (t, c, h, w)
|
||||||
|
let frames_tensor = Tensor::stack(&frames_vec, 0)?.contiguous()?;
|
||||||
|
Ok(frames_tensor)
|
||||||
|
}
|
||||||
@@ -0,0 +1 @@
|
|||||||
|
pub mod rope;
|
||||||
@@ -0,0 +1,182 @@
|
|||||||
|
use anyhow::Result;
|
||||||
|
use candle_core::{D, DType, Device, IndexOp, Tensor};
|
||||||
|
use candle_transformers::models::deepseek2::SplitOp;
|
||||||
|
|
||||||
|
use crate::models::qwen2_5vl::config::RopeScaling;
|
||||||
|
|
||||||
|
pub fn compute_default_rope_parameters(dim: usize, base: f32) -> Vec<f32> {
|
||||||
|
let inv_freq: Vec<f32> = (0..dim)
|
||||||
|
.step_by(2)
|
||||||
|
.map(|i| 1.0_f32 / base.powf(i as f32 / dim as f32))
|
||||||
|
.collect();
|
||||||
|
inv_freq
|
||||||
|
}
|
||||||
|
|
||||||
|
pub fn rotate_half(x: &Tensor) -> Result<Tensor> {
|
||||||
|
let half_dim = x.dim(D::Minus1)? / 2;
|
||||||
|
let x1 = x.narrow(D::Minus1, 0, half_dim)?;
|
||||||
|
let x2 = x.narrow(D::Minus1, half_dim, half_dim)?;
|
||||||
|
let x2 = x2.affine(-1.0, 0.0)?;
|
||||||
|
let rotate_x = Tensor::cat(&[&x2, &x1], D::Minus1)?.contiguous()?;
|
||||||
|
Ok(rotate_x)
|
||||||
|
}
|
||||||
|
|
||||||
|
pub fn apply_multimodel_rotary_pos_emb(
|
||||||
|
q: &Tensor,
|
||||||
|
k: &Tensor,
|
||||||
|
cos: &Tensor,
|
||||||
|
sin: &Tensor,
|
||||||
|
mrope_section: Vec<usize>,
|
||||||
|
) -> Result<(Tensor, Tensor)> {
|
||||||
|
let mrope_section = mrope_section.repeat(2);
|
||||||
|
let cos_select: Vec<Tensor> = cos
|
||||||
|
.split(&mrope_section, D::Minus1)?
|
||||||
|
.iter()
|
||||||
|
.enumerate()
|
||||||
|
.map(|(i, m)| m.i(i % 3).unwrap())
|
||||||
|
.collect();
|
||||||
|
let cos = Tensor::cat(&cos_select, D::Minus1)?
|
||||||
|
.unsqueeze(1)?
|
||||||
|
.contiguous()?;
|
||||||
|
let sin_select: Vec<Tensor> = sin
|
||||||
|
.split(&mrope_section, D::Minus1)?
|
||||||
|
.iter()
|
||||||
|
.enumerate()
|
||||||
|
.map(|(i, m)| m.i(i % 3).unwrap())
|
||||||
|
.collect();
|
||||||
|
let sin = Tensor::cat(&sin_select, D::Minus1)?
|
||||||
|
.unsqueeze(1)?
|
||||||
|
.contiguous()?;
|
||||||
|
let q_embed = q
|
||||||
|
.broadcast_mul(&cos)?
|
||||||
|
.add(&rotate_half(&q)?.broadcast_mul(&sin)?)?;
|
||||||
|
let k_embed = k
|
||||||
|
.broadcast_mul(&cos)?
|
||||||
|
.add(&rotate_half(&k)?.broadcast_mul(&sin)?)?;
|
||||||
|
Ok((q_embed, k_embed))
|
||||||
|
}
|
||||||
|
|
||||||
|
pub fn apply_rotary_pos_emb_vision(
|
||||||
|
q: &Tensor,
|
||||||
|
k: &Tensor,
|
||||||
|
cos: &Tensor,
|
||||||
|
sin: &Tensor,
|
||||||
|
) -> Result<(Tensor, Tensor)> {
|
||||||
|
// q, k -> (seq_len, num_heads, head_dim)
|
||||||
|
// cos, sin -> (seq_len, head_dim) -> (seq_len, 1, head_dim)
|
||||||
|
let cos = cos.unsqueeze(D::Minus2)?;
|
||||||
|
let sin = sin.unsqueeze(D::Minus2)?;
|
||||||
|
let q_embed = q
|
||||||
|
.broadcast_mul(&cos)?
|
||||||
|
.add(&rotate_half(q)?.broadcast_mul(&sin)?)?;
|
||||||
|
let k_embed = k
|
||||||
|
.broadcast_mul(&cos)?
|
||||||
|
.add(&rotate_half(k)?.broadcast_mul(&sin)?)?;
|
||||||
|
Ok((q_embed, k_embed))
|
||||||
|
}
|
||||||
|
|
||||||
|
pub fn apply_rotary_pos_emb(
|
||||||
|
q: &Tensor,
|
||||||
|
k: &Tensor,
|
||||||
|
cos: &Tensor,
|
||||||
|
sin: &Tensor,
|
||||||
|
) -> Result<(Tensor, Tensor)> {
|
||||||
|
// sin/cos: (bs, 1, seq_len, head_dim)
|
||||||
|
// q/k: (bs, n_head, seq_len, head_dim)
|
||||||
|
let q_embed = q
|
||||||
|
.broadcast_mul(&cos)?
|
||||||
|
.add(&rotate_half(q)?.broadcast_mul(&sin)?)?;
|
||||||
|
let k_embed = k
|
||||||
|
.broadcast_mul(&cos)?
|
||||||
|
.add(&rotate_half(k)?.broadcast_mul(&sin)?)?;
|
||||||
|
Ok((q_embed, k_embed))
|
||||||
|
}
|
||||||
|
|
||||||
|
#[derive(Debug, Clone)]
|
||||||
|
pub struct Qwen2_5VLTextRotaryEmbedding {
|
||||||
|
inv_freq: Vec<f32>,
|
||||||
|
}
|
||||||
|
|
||||||
|
impl Qwen2_5VLTextRotaryEmbedding {
|
||||||
|
pub fn new(dim: usize, theta_base: f32) -> Self {
|
||||||
|
let inv_freq = compute_default_rope_parameters(dim, theta_base);
|
||||||
|
Self { inv_freq }
|
||||||
|
}
|
||||||
|
pub fn forward(
|
||||||
|
&self,
|
||||||
|
position_ids: &Tensor,
|
||||||
|
dtype: DType,
|
||||||
|
mrope_section: Vec<usize>,
|
||||||
|
) -> Result<(Tensor, Tensor)> {
|
||||||
|
// position_ids shape: (3, bs, position) -> (3, bs, 1, position)
|
||||||
|
let position_ids_expanded = position_ids
|
||||||
|
.unsqueeze(D::Minus2)?
|
||||||
|
.to_dtype(DType::F32)?
|
||||||
|
.contiguous()?;
|
||||||
|
// inv_freq Vec<f32> -> Tensor(1, 1, head_dim / 2, 1) -> (3, bs, head_dim / 2, 1)
|
||||||
|
let inv_freq_expanded = Tensor::from_vec(
|
||||||
|
self.inv_freq.clone(),
|
||||||
|
(1, 1, self.inv_freq.len(), 1),
|
||||||
|
position_ids.device(),
|
||||||
|
)?
|
||||||
|
.broadcast_as((3, position_ids.dim(1)?, self.inv_freq.len(), 1))?
|
||||||
|
.to_dtype(DType::F32)?
|
||||||
|
.contiguous()?;
|
||||||
|
|
||||||
|
// (3, bs, head_dim / 2, 1) matmul (3, bs, 1, position)
|
||||||
|
// -> (3, bs, head_dim / 2, seq_len) -> (3, bs, seq_len, head_dim / 2)
|
||||||
|
let freqs = inv_freq_expanded
|
||||||
|
.matmul(&position_ids_expanded)?
|
||||||
|
.transpose(2, 3)?;
|
||||||
|
// let freqs = position_ids_expanded.matmul(&inv_freq_expanded)?;
|
||||||
|
// (3, bs, seq_len, head_dim / 2) -> (3, bs, seq_len, head_dim)
|
||||||
|
let emb = Tensor::cat(&[&freqs, &freqs], D::Minus1)?.contiguous()?;
|
||||||
|
let cos = emb.cos()?;
|
||||||
|
let sin = emb.sin()?;
|
||||||
|
let mrope_section = mrope_section.repeat(2);
|
||||||
|
let cos_select: Vec<Tensor> = cos
|
||||||
|
.split(&mrope_section, D::Minus1)?
|
||||||
|
.iter()
|
||||||
|
.enumerate()
|
||||||
|
.map(|(i, m)| m.i(i % 3).unwrap())
|
||||||
|
.collect();
|
||||||
|
// (bs, seq_len, head_dim) -> (bs, 1, seq_len, head_dim)
|
||||||
|
let cos = Tensor::cat(&cos_select, D::Minus1)?
|
||||||
|
.unsqueeze(1)?
|
||||||
|
.contiguous()?;
|
||||||
|
let sin_select: Vec<Tensor> = sin
|
||||||
|
.split(&mrope_section, D::Minus1)?
|
||||||
|
.iter()
|
||||||
|
.enumerate()
|
||||||
|
.map(|(i, m)| m.i(i % 3).unwrap())
|
||||||
|
.collect();
|
||||||
|
// (bs, seq_len, head_dim) -> (bs, 1, seq_len, head_dim)
|
||||||
|
let sin = Tensor::cat(&sin_select, D::Minus1)?
|
||||||
|
.unsqueeze(1)?
|
||||||
|
.contiguous()?;
|
||||||
|
Ok((cos.to_dtype(dtype)?, sin.to_dtype(dtype)?))
|
||||||
|
}
|
||||||
|
}
|
||||||
|
|
||||||
|
#[derive(Debug, Clone)]
|
||||||
|
pub struct Qwen2_5VisionRotaryEmbedding {
|
||||||
|
inv_freq: Vec<f32>,
|
||||||
|
}
|
||||||
|
|
||||||
|
impl Qwen2_5VisionRotaryEmbedding {
|
||||||
|
pub fn new(dim: usize, theta_base: Option<f32>) -> Self {
|
||||||
|
let theta_base = match theta_base {
|
||||||
|
Some(theta) => theta,
|
||||||
|
None => 10000.0_f32,
|
||||||
|
};
|
||||||
|
let inv_freq = compute_default_rope_parameters(dim, theta_base);
|
||||||
|
Self { inv_freq }
|
||||||
|
}
|
||||||
|
|
||||||
|
pub fn forward(&self, seqlen: usize, device: &Device) -> Result<Tensor> {
|
||||||
|
let seq = Tensor::arange(0.0_f32, seqlen as f32, device)?.reshape((seqlen, 1))?;
|
||||||
|
let inv_freq = Tensor::from_vec(self.inv_freq.clone(), (1, self.inv_freq.len()), device)?;
|
||||||
|
let freqs = seq.matmul(&inv_freq)?;
|
||||||
|
Ok(freqs)
|
||||||
|
}
|
||||||
|
}
|
||||||
@@ -0,0 +1 @@
|
|||||||
|
pub mod tokenizer;
|
||||||
@@ -0,0 +1,44 @@
|
|||||||
|
use anyhow::{Result, anyhow};
|
||||||
|
use candle_core::{Device, Tensor};
|
||||||
|
use tokenizers::Tokenizer;
|
||||||
|
|
||||||
|
pub struct TokenizerModel {
|
||||||
|
tokenizer: Tokenizer,
|
||||||
|
}
|
||||||
|
|
||||||
|
impl TokenizerModel {
|
||||||
|
pub fn init(path: &str) -> Result<Self> {
|
||||||
|
let path = path.to_string();
|
||||||
|
assert!(
|
||||||
|
std::path::Path::new(&path).exists(),
|
||||||
|
"model path file not exists"
|
||||||
|
);
|
||||||
|
let tokenizer_file = path.clone() + "/tokenizer.json";
|
||||||
|
assert!(
|
||||||
|
std::path::Path::new(&tokenizer_file).exists(),
|
||||||
|
"tokenizer.json not exists in model path"
|
||||||
|
);
|
||||||
|
let tokenizer = Tokenizer::from_file(tokenizer_file)
|
||||||
|
.map_err(|e| anyhow!(format!("tokenizer from file error{}", e)))?;
|
||||||
|
Ok(Self { tokenizer })
|
||||||
|
}
|
||||||
|
|
||||||
|
pub fn text_encode(&self, text: String, device: &Device) -> Result<Tensor> {
|
||||||
|
let token_id = self
|
||||||
|
.tokenizer
|
||||||
|
.encode(text, true)
|
||||||
|
.map_err(|e| anyhow!(format!("tokenizer encode error: {}", e)))?
|
||||||
|
.get_ids()
|
||||||
|
.to_vec();
|
||||||
|
let token_tensor = Tensor::from_slice(&token_id, (1, token_id.len()), device)?;
|
||||||
|
Ok(token_tensor)
|
||||||
|
}
|
||||||
|
|
||||||
|
pub fn token_decode(&self, tokens: Vec<u32>) -> Result<String> {
|
||||||
|
let decode = self
|
||||||
|
.tokenizer
|
||||||
|
.decode(&tokens, true)
|
||||||
|
.map_err(|e| anyhow!(format!("tokenizer encode error{}", e)))?;
|
||||||
|
Ok(decode)
|
||||||
|
}
|
||||||
|
}
|
||||||
@@ -0,0 +1,62 @@
|
|||||||
|
use std::io::Cursor;
|
||||||
|
|
||||||
|
use anyhow::{Result, anyhow};
|
||||||
|
use base64::{Engine, engine::general_purpose};
|
||||||
|
use image::{DynamicImage, ImageReader};
|
||||||
|
|
||||||
|
pub fn load_image_from_url(url: &str) -> Result<DynamicImage> {
|
||||||
|
let response = reqwest::blocking::get(url)
|
||||||
|
.map_err(|e| anyhow!(format!("Failed to fetch image from url: {}", e)))?;
|
||||||
|
let bytes = response
|
||||||
|
.bytes()
|
||||||
|
.map_err(|e| anyhow!(format!("Failed to get image bytes: {}", e)))?;
|
||||||
|
|
||||||
|
let cursor = Cursor::new(bytes);
|
||||||
|
let img = ImageReader::new(cursor)
|
||||||
|
.with_guessed_format()
|
||||||
|
.map_err(|e| anyhow!(format!("Failed to read image format: {}", e)))?
|
||||||
|
.decode()
|
||||||
|
.map_err(|e| anyhow!(format!("Failed to decode image: {}", e)))?;
|
||||||
|
Ok(img)
|
||||||
|
}
|
||||||
|
|
||||||
|
pub fn load_image_from_base64(base64_data: &str) -> Result<DynamicImage> {
|
||||||
|
let image_data = general_purpose::STANDARD
|
||||||
|
.decode(base64_data)
|
||||||
|
.map_err(|e| anyhow!(format!("Failed to decode image: {}", e)))?;
|
||||||
|
let cursor = Cursor::new(image_data);
|
||||||
|
let img = ImageReader::new(cursor)
|
||||||
|
.with_guessed_format()
|
||||||
|
.map_err(|e| anyhow!(format!("Failed to read image format: {}", e)))?
|
||||||
|
.decode()
|
||||||
|
.map_err(|e| anyhow!(format!("Failed to decode image: {}", e)))?;
|
||||||
|
Ok(img)
|
||||||
|
}
|
||||||
|
|
||||||
|
pub fn get_image(file: &String) -> Result<DynamicImage> {
|
||||||
|
let mut img = None;
|
||||||
|
if file.starts_with("http://") || file.starts_with("https://") {
|
||||||
|
img = Some(load_image_from_url(&file)?);
|
||||||
|
}
|
||||||
|
if file.starts_with("file://") {
|
||||||
|
let mut path = file.clone();
|
||||||
|
path = path.split_off(7);
|
||||||
|
img = Some(
|
||||||
|
ImageReader::open(path)
|
||||||
|
.map_err(|e| anyhow!(format!("Failed to open file: {}", e)))?
|
||||||
|
.decode()
|
||||||
|
.map_err(|e| anyhow!(format!("Failed to decode image: {}", e)))?,
|
||||||
|
);
|
||||||
|
}
|
||||||
|
if file.starts_with("data:image") {
|
||||||
|
if file.contains("base64,") {
|
||||||
|
let data: Vec<&str> = file.split("base64,").collect();
|
||||||
|
let data = data[1];
|
||||||
|
img = Some(load_image_from_base64(data)?);
|
||||||
|
}
|
||||||
|
}
|
||||||
|
if img.is_some() {
|
||||||
|
return Ok(img.unwrap());
|
||||||
|
}
|
||||||
|
Err(anyhow!("get image from message failed".to_string()))
|
||||||
|
}
|
||||||
@@ -0,0 +1,4 @@
|
|||||||
|
pub mod tensor_utils;
|
||||||
|
pub mod utils;
|
||||||
|
pub mod img_utils;
|
||||||
|
pub mod video_utils;
|
||||||
@@ -0,0 +1,202 @@
|
|||||||
|
use anyhow::{Result, anyhow};
|
||||||
|
use candle_core::{D, DType, Device, IndexOp, Tensor, shape::Dim};
|
||||||
|
|
||||||
|
pub fn split(t: &Tensor, splits: &[usize], dim: D) -> Result<Vec<Tensor>> {
|
||||||
|
let dim = dim.to_index(t.shape(), "split")?;
|
||||||
|
let mut split_res = Vec::new();
|
||||||
|
let mut index = 0;
|
||||||
|
for split in splits {
|
||||||
|
split_res.push(t.narrow(dim, index, *split)?);
|
||||||
|
index += *split;
|
||||||
|
}
|
||||||
|
Ok(split_res)
|
||||||
|
}
|
||||||
|
|
||||||
|
pub fn safe_arg_sort_last_dim(t: &Tensor, ascending: bool) -> Result<Tensor> {
|
||||||
|
// tensor在GPU上时,维度超过1024, arg_sort_last_dim方法会报错
|
||||||
|
// 所以维度大于1024时,放到CPU上处理
|
||||||
|
let last_dim = t.dims()[t.rank() - 1];
|
||||||
|
if last_dim <= 1024 {
|
||||||
|
let t = t.arg_sort_last_dim(ascending)?;
|
||||||
|
Ok(t)
|
||||||
|
} else {
|
||||||
|
let cpu_tensor = t.to_device(&Device::Cpu)?;
|
||||||
|
let sorted_indices = cpu_tensor.arg_sort_last_dim(ascending)?;
|
||||||
|
let t = sorted_indices.to_device(t.device())?;
|
||||||
|
Ok(t)
|
||||||
|
}
|
||||||
|
}
|
||||||
|
|
||||||
|
pub fn nonzero_index_vec(mask: &Tensor) -> Result<Vec<u32>> {
|
||||||
|
// 根据mask矩阵选出其中不为0的元素所在索引, 返回vec
|
||||||
|
// 只能处理1维数据
|
||||||
|
let mut mask = mask.clone();
|
||||||
|
if mask.dtype() != DType::U32 {
|
||||||
|
mask = mask.to_dtype(DType::U32)?;
|
||||||
|
}
|
||||||
|
match mask.rank() {
|
||||||
|
0 => {
|
||||||
|
return Err(anyhow!(format!(
|
||||||
|
"input rank must > 0, the input tensor rank: {}",
|
||||||
|
mask.rank()
|
||||||
|
)));
|
||||||
|
}
|
||||||
|
1 => {
|
||||||
|
let mask_vector = mask.to_vec1::<u32>()?;
|
||||||
|
let indices: Vec<u32> = mask_vector
|
||||||
|
.iter()
|
||||||
|
.enumerate()
|
||||||
|
.filter_map(|(idx, &val)| if val != 0 { Some(idx as u32) } else { None })
|
||||||
|
.collect();
|
||||||
|
Ok(indices)
|
||||||
|
}
|
||||||
|
_ => {
|
||||||
|
return Err(anyhow!(format!(
|
||||||
|
"input rank not support, the input tensor rank: {}",
|
||||||
|
mask.rank()
|
||||||
|
)));
|
||||||
|
}
|
||||||
|
}
|
||||||
|
}
|
||||||
|
|
||||||
|
pub fn nonzero_index(mask: &Tensor) -> Result<Tensor> {
|
||||||
|
// 根据mask矩阵选出其中不为1的元素所在索引, 返回Tensor
|
||||||
|
let indices_tensor = match mask.rank() {
|
||||||
|
0 => {
|
||||||
|
return Err(anyhow!(format!(
|
||||||
|
"input rank must > 0, the input tensor rank: {}",
|
||||||
|
mask.rank()
|
||||||
|
)));
|
||||||
|
}
|
||||||
|
1 => {
|
||||||
|
let index_vec = nonzero_index_vec(mask)?;
|
||||||
|
let indices_tensor = Tensor::from_slice(&index_vec, index_vec.len(), mask.device())?;
|
||||||
|
indices_tensor
|
||||||
|
}
|
||||||
|
_ => {
|
||||||
|
return Err(anyhow!(format!(
|
||||||
|
"input rank must == 1, the input tensor rank: {}",
|
||||||
|
mask.rank()
|
||||||
|
)));
|
||||||
|
}
|
||||||
|
};
|
||||||
|
Ok(indices_tensor)
|
||||||
|
}
|
||||||
|
|
||||||
|
pub fn zero_index_vec(mask: &Tensor) -> Result<Vec<u32>> {
|
||||||
|
// 根据mask矩阵选出其中为0的元素所在索引, 返回vec
|
||||||
|
// 只能处理1维数据
|
||||||
|
let mut mask = mask.clone();
|
||||||
|
if mask.dtype() != DType::U32 {
|
||||||
|
mask = mask.to_dtype(DType::U32)?;
|
||||||
|
}
|
||||||
|
match mask.rank() {
|
||||||
|
0 => {
|
||||||
|
return Err(anyhow!(format!(
|
||||||
|
"input rank must > 0, the input tensor rank: {}",
|
||||||
|
mask.rank()
|
||||||
|
)));
|
||||||
|
}
|
||||||
|
1 => {
|
||||||
|
let mask_vector = mask.to_vec1::<u32>()?;
|
||||||
|
let indices: Vec<u32> = mask_vector
|
||||||
|
.iter()
|
||||||
|
.enumerate()
|
||||||
|
.filter_map(|(idx, &val)| if val == 0 { Some(idx as u32) } else { None })
|
||||||
|
.collect();
|
||||||
|
Ok(indices)
|
||||||
|
}
|
||||||
|
_ => {
|
||||||
|
return Err(anyhow!(format!(
|
||||||
|
"input rank not support, the input tensor rank: {}",
|
||||||
|
mask.rank()
|
||||||
|
)));
|
||||||
|
}
|
||||||
|
}
|
||||||
|
}
|
||||||
|
|
||||||
|
pub fn zero_index(mask: &Tensor) -> Result<Tensor> {
|
||||||
|
let index_vec = zero_index_vec(mask)?;
|
||||||
|
let indices_tensor = Tensor::from_slice(&index_vec, index_vec.len(), mask.device())?;
|
||||||
|
Ok(indices_tensor)
|
||||||
|
}
|
||||||
|
|
||||||
|
pub fn nonzero_slice(mask: &Tensor) -> Result<Vec<(usize, usize)>> {
|
||||||
|
// 根据mask矩阵选出其中非0的元素所在索引
|
||||||
|
// 根据索引获取连续索引间隔
|
||||||
|
// 如不为零索引元素为[0, 3, 4, 5, 8, 9]
|
||||||
|
// 间隔为: [(0, 1), (3, 6), (8, 10)]
|
||||||
|
// 索引前闭后开
|
||||||
|
let mut index_vec = nonzero_index_vec(mask)?;
|
||||||
|
match index_vec.len() {
|
||||||
|
0 => {
|
||||||
|
return Ok(vec![]);
|
||||||
|
}
|
||||||
|
1 => {
|
||||||
|
return Ok(vec![(index_vec[0] as usize, (index_vec[0] + 1) as usize)]);
|
||||||
|
}
|
||||||
|
_ => {
|
||||||
|
let mut vec_slice = vec![];
|
||||||
|
let mut start = index_vec.remove(0);
|
||||||
|
let mut last = start;
|
||||||
|
|
||||||
|
for i in index_vec {
|
||||||
|
if i == (last + 1) {
|
||||||
|
last = i;
|
||||||
|
continue;
|
||||||
|
} else {
|
||||||
|
vec_slice.push((start as usize, (last + 1) as usize));
|
||||||
|
start = i;
|
||||||
|
last = i;
|
||||||
|
}
|
||||||
|
}
|
||||||
|
vec_slice.push((start as usize, (last + 1) as usize));
|
||||||
|
Ok(vec_slice)
|
||||||
|
}
|
||||||
|
}
|
||||||
|
}
|
||||||
|
|
||||||
|
pub fn masked_scatter_dim0(original: &Tensor, replace: &Tensor, mask: &Tensor) -> Result<Tensor> {
|
||||||
|
// 根据mask中非0元素所在索引,使用replace中的数据替换掉original中的数据
|
||||||
|
// original: rank = 3: (bs, seq_len, hidden_dim)
|
||||||
|
// replace: rank = 2: (seq_len, hidden_dim)
|
||||||
|
// mask: rank = 2: (bs, seq_len)
|
||||||
|
// 推理时bs=1,为了方便替换,将bs squeeze,替换后再unsqueeze
|
||||||
|
// 按行替换
|
||||||
|
if original.dim(0)? != 1 || mask.dim(0)? != 1 {
|
||||||
|
return Err(anyhow!(format!(
|
||||||
|
"masked_scatter_dim0 original bs: {} or mask bs :{} not equal to 1 ",
|
||||||
|
original.dim(0)?,
|
||||||
|
mask.dim(0)? != 1
|
||||||
|
)));
|
||||||
|
}
|
||||||
|
let mut original = original.squeeze(0)?;
|
||||||
|
let mask = mask.squeeze(0)?;
|
||||||
|
let slices = nonzero_slice(&mask)?;
|
||||||
|
let mut sub_start = 0usize;
|
||||||
|
let mut sub_end = 0usize;
|
||||||
|
for (start, end) in slices {
|
||||||
|
sub_end = sub_start + (end - start);
|
||||||
|
let sub_replace = replace.i((sub_start..sub_end, ..))?;
|
||||||
|
original = original.slice_assign(&[(start..end), (0..original.dim(1)?)], &sub_replace)?;
|
||||||
|
sub_start = sub_end;
|
||||||
|
}
|
||||||
|
original = original.unsqueeze(0)?;
|
||||||
|
Ok(original)
|
||||||
|
}
|
||||||
|
|
||||||
|
pub fn get_equal_mask(input_ids: &Tensor, token_ids: u32) -> Result<Tensor> {
|
||||||
|
let image_token_id_tensor = Tensor::new(vec![token_ids], input_ids.device())?;
|
||||||
|
let mask = input_ids
|
||||||
|
.broadcast_eq(&image_token_id_tensor)?
|
||||||
|
.to_dtype(candle_core::DType::U32)?;
|
||||||
|
Ok(mask)
|
||||||
|
}
|
||||||
|
|
||||||
|
pub fn get_vision_next_indices(input_ids: &Tensor, token_id: u32) -> Result<Tensor> {
|
||||||
|
// input_ids -> shape: (seq_len)
|
||||||
|
let mask = get_equal_mask(&input_ids, token_id)?;
|
||||||
|
let indices = nonzero_index(&mask)?;
|
||||||
|
let indices = indices.broadcast_add(&Tensor::new(vec![1u32], input_ids.device())?)?;
|
||||||
|
Ok(indices)
|
||||||
|
}
|
||||||
@@ -0,0 +1,87 @@
|
|||||||
|
use anyhow::Result;
|
||||||
|
use candle_core::{DType, Device};
|
||||||
|
|
||||||
|
pub fn get_device(device: Option<&Device>) -> Device {
|
||||||
|
match device {
|
||||||
|
Some(d) => d.clone(),
|
||||||
|
None => {
|
||||||
|
#[cfg(feature = "cuda")]
|
||||||
|
{
|
||||||
|
Device::new_cuda(0).unwrap_or(Device::Cpu)
|
||||||
|
}
|
||||||
|
#[cfg(not(feature = "cuda"))]
|
||||||
|
{
|
||||||
|
Device::Cpu
|
||||||
|
}
|
||||||
|
}
|
||||||
|
}
|
||||||
|
}
|
||||||
|
|
||||||
|
pub fn get_dtype(dtype: Option<DType>, cfg_dtype: &str) -> DType {
|
||||||
|
match dtype {
|
||||||
|
Some(d) => d,
|
||||||
|
None => {
|
||||||
|
#[cfg(feature = "cuda")]
|
||||||
|
{
|
||||||
|
match cfg_dtype {
|
||||||
|
"float32" | "float" => DType::F32,
|
||||||
|
"float64" | "double" => DType::F64,
|
||||||
|
"float16" => DType::F16,
|
||||||
|
"bfloat16" => DType::BF16,
|
||||||
|
"uint8" => DType::U8,
|
||||||
|
"int8" | "int16" | "int32" | "int64" => DType::I64,
|
||||||
|
_ => DType::F32,
|
||||||
|
}
|
||||||
|
}
|
||||||
|
#[cfg(not(feature = "cuda"))]
|
||||||
|
{
|
||||||
|
match cfg_dtype {
|
||||||
|
"float32" | "float" => DType::F32,
|
||||||
|
"float64" | "double" => DType::F64,
|
||||||
|
"float16" | "bfloat16" => DType::F16, // cpu上bfloat16有问题
|
||||||
|
"uint8" => DType::U8,
|
||||||
|
"int8" | "int16" | "int32" | "int64" => DType::I64,
|
||||||
|
_ => DType::F32,
|
||||||
|
}
|
||||||
|
}
|
||||||
|
}
|
||||||
|
}
|
||||||
|
}
|
||||||
|
|
||||||
|
pub fn string_to_static_str(s: String) -> &'static str {
|
||||||
|
Box::leak(s.into_boxed_str())
|
||||||
|
}
|
||||||
|
|
||||||
|
pub fn find_safetensors_files(path: &str) -> Result<Vec<String>> {
|
||||||
|
let mut files = Vec::new();
|
||||||
|
|
||||||
|
for entry in std::fs::read_dir(path)? {
|
||||||
|
let entry = entry?;
|
||||||
|
let file_path = entry.path();
|
||||||
|
|
||||||
|
if file_path.is_file() {
|
||||||
|
if let Some(extension) = file_path.extension() {
|
||||||
|
if extension == "safetensors" {
|
||||||
|
files.push(file_path.to_string_lossy().to_string());
|
||||||
|
}
|
||||||
|
}
|
||||||
|
}
|
||||||
|
}
|
||||||
|
|
||||||
|
Ok(files)
|
||||||
|
}
|
||||||
|
|
||||||
|
pub fn round_by_factor(num: u32, factor: u32) -> u32 {
|
||||||
|
let round = (num as f32 / factor as f32).round() as u32;
|
||||||
|
round * factor
|
||||||
|
}
|
||||||
|
|
||||||
|
pub fn floor_by_factor(num: f32, factor: u32) -> u32 {
|
||||||
|
let floor = (num / factor as f32).floor() as u32;
|
||||||
|
floor * factor
|
||||||
|
}
|
||||||
|
|
||||||
|
pub fn ceil_by_factor(num: f32, factor: u32) -> u32 {
|
||||||
|
let ceil = (num / factor as f32).ceil() as u32;
|
||||||
|
ceil * factor
|
||||||
|
}
|
||||||
@@ -0,0 +1,13 @@
|
|||||||
|
use ffmpeg_next as ffmpeg;
|
||||||
|
use std::{fs::File, io::Write};
|
||||||
|
|
||||||
|
#[allow(unused)]
|
||||||
|
fn save_file(
|
||||||
|
frame: &ffmpeg::frame::Video,
|
||||||
|
index: usize,
|
||||||
|
) -> std::result::Result<(), std::io::Error> {
|
||||||
|
let mut file = File::create(format!("frame{}.ppm", index))?;
|
||||||
|
file.write_all(format!("P6\n{} {}\n255\n", frame.width(), frame.height()).as_bytes())?;
|
||||||
|
file.write_all(frame.data(0))?;
|
||||||
|
Ok(())
|
||||||
|
}
|
||||||
@@ -0,0 +1,12 @@
|
|||||||
|
use anyhow::Result;
|
||||||
|
use aha::models::qwen2_5vl::config::Config;
|
||||||
|
|
||||||
|
#[test]
|
||||||
|
fn qwen2_5vl_config() -> Result<()> {
|
||||||
|
// cargo test qwen2_5vl_config -- --nocapture
|
||||||
|
let model_path = "/home/jhq/huggingface_model/Qwen/Qwen2.5-VL-3B-Instruct/";
|
||||||
|
let config_path = model_path.to_string() + "/config.json";
|
||||||
|
let config: Config = serde_json::from_slice(&std::fs::read(config_path)?)?;
|
||||||
|
println!("{:?}", config);
|
||||||
|
Ok(())
|
||||||
|
}
|
||||||
@@ -0,0 +1,56 @@
|
|||||||
|
use std::time::Instant;
|
||||||
|
|
||||||
|
use aha::{models::{qwen2_5vl::generate::Qwen2_5VLGenerateModel, GenerateModel}, ModelType};
|
||||||
|
use anyhow::{Result};
|
||||||
|
use candle_core::{DType, Device};
|
||||||
|
use openai_dive::v1::resources::chat::ChatCompletionParameters;
|
||||||
|
|
||||||
|
|
||||||
|
#[test]
|
||||||
|
fn qwen2_5vl_generate() -> Result<()> {
|
||||||
|
// test with cpu :(太慢了, : RUST_BACKTRACE=1 cargo test qwen2_5vl_generate -- --nocapture
|
||||||
|
// test with cuda: RUST_BACKTRACE=1 cargo test -F cuda qwen2_5vl_generate -- --nocapture
|
||||||
|
// test with cuda+flash-attn: RUST_BACKTRACE=1 cargo test -F cuda,flash-attn qwen2_5vl_generate -- --nocapture
|
||||||
|
let device = Device::cuda_if_available(0)?;
|
||||||
|
let dtype = DType::BF16;
|
||||||
|
|
||||||
|
let model_path = "/home/jhq/huggingface_model/Qwen/Qwen2.5-VL-3B-Instruct/";
|
||||||
|
|
||||||
|
let message = r#"
|
||||||
|
{
|
||||||
|
"model": "qwen2.5vl",
|
||||||
|
"messages": [
|
||||||
|
{
|
||||||
|
"role": "user",
|
||||||
|
"content": [
|
||||||
|
{
|
||||||
|
"type": "image",
|
||||||
|
"image_url":
|
||||||
|
{
|
||||||
|
"url": "file://./assets/img/ocr_test.png"
|
||||||
|
}
|
||||||
|
},
|
||||||
|
{
|
||||||
|
"type": "text",
|
||||||
|
"text": "请分析图片并提取所有可见文本内容,按从左到右、从上到下的布局,返回纯文本"
|
||||||
|
}
|
||||||
|
]
|
||||||
|
}
|
||||||
|
]
|
||||||
|
}
|
||||||
|
"#;
|
||||||
|
let mes:ChatCompletionParameters = serde_json::from_str(message)?;
|
||||||
|
let i_start = Instant::now();
|
||||||
|
// let mut model = Qwen2_5VLGenerateModel::init(model_path, &device, dtype)?;
|
||||||
|
let mut model = ModelType::init(ModelType::Qwen2_5VL, model_path, None, None)?;
|
||||||
|
let i_duration = i_start.elapsed();
|
||||||
|
println!("Time elapsed in load model is: {:?}", i_duration);
|
||||||
|
|
||||||
|
let i_start = Instant::now();
|
||||||
|
let result = model.generate(mes)?;
|
||||||
|
println!("generate: \n{}", result);
|
||||||
|
let i_duration = i_start.elapsed();
|
||||||
|
println!("Time elapsed in generate is: {:?}", i_duration);
|
||||||
|
|
||||||
|
Ok(())
|
||||||
|
}
|
||||||
Reference in New Issue
Block a user