add hunyuan_ocr
This commit is contained in:
@@ -0,0 +1,105 @@
|
||||
use candle_nn::Activation;
|
||||
use serde::{Deserialize, Serialize};
|
||||
|
||||
#[derive(Debug, Clone, PartialEq, Serialize, Deserialize)]
|
||||
pub struct HunYuanVLConfig {
|
||||
pub attention_bias: bool,
|
||||
pub attention_dropout: f64,
|
||||
pub attention_head_dim: usize,
|
||||
pub bos_token_id: u32,
|
||||
pub eod_token_id: u32,
|
||||
pub eos_token_id: u32,
|
||||
pub head_dim: usize,
|
||||
pub hidden_act: Activation,
|
||||
pub hidden_size: usize,
|
||||
pub image_start_token_id: u32,
|
||||
pub image_end_token_id: u32,
|
||||
pub image_token_id: u32,
|
||||
pub image_newline_token_id: u32,
|
||||
pub initializer_range: f64,
|
||||
pub intermediate_size: usize,
|
||||
pub max_position_embeddings: usize,
|
||||
pub mlp_bias: bool,
|
||||
pub norm_type: String,
|
||||
pub num_attention_heads: usize,
|
||||
pub num_experts: usize,
|
||||
pub num_hidden_layers: usize,
|
||||
pub num_key_value_heads: usize,
|
||||
pub org_vocab_size: usize,
|
||||
pub pad_id: i32,
|
||||
pub pad_token_id: i32,
|
||||
pub pretraining_tp: i32,
|
||||
pub rms_norm_eps: f64,
|
||||
pub rope_scaling: HunYuanVLRopeScaling,
|
||||
pub rope_theta: f64,
|
||||
pub routed_scaling_factor: f64,
|
||||
pub sep_token_id: u32,
|
||||
pub text_end_id: u32,
|
||||
pub text_start_id: u32,
|
||||
pub tie_word_embeddings: bool,
|
||||
pub dtype: String,
|
||||
pub use_cache: bool,
|
||||
pub use_qk_norm: bool,
|
||||
pub use_cla: bool,
|
||||
pub vision_config: HunYuanVLVisionConfig,
|
||||
pub vocab_size: usize,
|
||||
}
|
||||
|
||||
#[derive(Debug, Clone, PartialEq, Serialize, Deserialize)]
|
||||
pub struct HunYuanVLRopeScaling {
|
||||
pub alpha: f64,
|
||||
pub beta_fast: i32,
|
||||
pub beta_slow: i32,
|
||||
pub factor: f64,
|
||||
pub mscale: f64,
|
||||
pub mscale_all_dim: f64,
|
||||
#[serde(rename = "type")]
|
||||
pub type_field: String,
|
||||
pub xdrope_section: Vec<usize>,
|
||||
}
|
||||
|
||||
#[derive(Debug, Clone, PartialEq, Serialize, Deserialize)]
|
||||
pub struct HunYuanVLVisionConfig {
|
||||
pub add_patchemb_bias: bool,
|
||||
pub attention_dropout: f64,
|
||||
pub cat_extra_token: i32,
|
||||
pub hidden_act: Activation,
|
||||
pub hidden_dropout: f64,
|
||||
pub hidden_size: usize,
|
||||
pub img_max_token_num: usize,
|
||||
pub intermediate_size: usize,
|
||||
pub interpolate_mode: String,
|
||||
pub max_image_size: usize,
|
||||
pub max_vit_seq_len: usize,
|
||||
pub num_attention_heads: usize,
|
||||
pub num_channels: usize,
|
||||
pub num_hidden_layers: usize,
|
||||
pub out_hidden_size: usize,
|
||||
pub patch_size: usize,
|
||||
pub rms_norm_eps: f64,
|
||||
pub spatial_merge_size: usize,
|
||||
}
|
||||
|
||||
#[derive(Debug, Clone, PartialEq, serde::Deserialize)]
|
||||
pub struct HunyuanOCRGenerationConfig {
|
||||
pub bos_token_id: usize,
|
||||
pub pad_token_id: usize,
|
||||
pub do_sample: bool,
|
||||
pub eos_token_id: Vec<usize>,
|
||||
pub top_p: f32,
|
||||
pub top_k: usize,
|
||||
pub temperature: f32,
|
||||
pub repetition_penalty: f32,
|
||||
}
|
||||
|
||||
#[derive(Debug, Clone, PartialEq, serde::Deserialize)]
|
||||
pub struct HunyuanOCRPreprocessorConfig {
|
||||
pub min_pixels: usize,
|
||||
pub max_pixels: usize,
|
||||
pub patch_size: usize,
|
||||
pub resample: usize,
|
||||
pub temporal_patch_size: usize,
|
||||
pub merge_size: usize,
|
||||
pub image_mean: Vec<f32>,
|
||||
pub image_std: Vec<f32>,
|
||||
}
|
||||
@@ -0,0 +1,218 @@
|
||||
use aha_openai_dive::v1::resources::chat::{
|
||||
ChatCompletionChunkResponse, ChatCompletionParameters, ChatCompletionResponse,
|
||||
};
|
||||
use anyhow::{Result, anyhow};
|
||||
use candle_core::{DType, Device, Tensor};
|
||||
use candle_nn::VarBuilder;
|
||||
use rocket::async_stream::stream;
|
||||
use rocket::futures::Stream;
|
||||
|
||||
use crate::{
|
||||
chat_template::ChatTemplate,
|
||||
models::{
|
||||
GenerateModel,
|
||||
hunyuan_ocr::{
|
||||
config::{HunYuanVLConfig, HunyuanOCRGenerationConfig},
|
||||
model::HunyuanVLModel,
|
||||
processor::HunyuanVLProcessor,
|
||||
},
|
||||
},
|
||||
tokenizer::TokenizerModel,
|
||||
utils::{
|
||||
build_completion_chunk_response, build_completion_response, find_type_files, get_device,
|
||||
get_dtype, get_logit_processor,
|
||||
},
|
||||
};
|
||||
|
||||
pub struct HunyuanOCRGenerateModel<'a> {
|
||||
chat_template: ChatTemplate<'a>,
|
||||
tokenizer: TokenizerModel,
|
||||
pre_processor: HunyuanVLProcessor,
|
||||
hunyuan_vl: HunyuanVLModel,
|
||||
device: Device,
|
||||
eos_token_id1: u32,
|
||||
eos_token_id2: u32,
|
||||
generation_config: HunyuanOCRGenerationConfig,
|
||||
model_name: String,
|
||||
}
|
||||
|
||||
impl<'a> HunyuanOCRGenerateModel<'a> {
|
||||
pub 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: HunYuanVLConfig = serde_json::from_slice(&std::fs::read(config_path)?)?;
|
||||
let device = get_device(device);
|
||||
let cfg_dtype = cfg.dtype.as_str();
|
||||
let dtype = get_dtype(dtype, cfg_dtype);
|
||||
let pre_processor = HunyuanVLProcessor::new(path, &device, dtype)?;
|
||||
let model_list = find_type_files(path, "safetensors")?;
|
||||
let vb = unsafe { VarBuilder::from_mmaped_safetensors(&model_list, dtype, &device)? };
|
||||
let hunyuan_vl = HunyuanVLModel::new(vb, cfg.clone())?;
|
||||
let generation_config_path = path.to_string() + "/generation_config.json";
|
||||
let generation_config: HunyuanOCRGenerationConfig =
|
||||
serde_json::from_slice(&std::fs::read(generation_config_path)?)?;
|
||||
Ok(Self {
|
||||
chat_template,
|
||||
tokenizer,
|
||||
pre_processor,
|
||||
hunyuan_vl,
|
||||
device,
|
||||
eos_token_id1: generation_config.eos_token_id[0] as u32,
|
||||
eos_token_id2: generation_config.eos_token_id[1] as u32,
|
||||
generation_config,
|
||||
model_name: "hunyuan_ocr".to_string(),
|
||||
})
|
||||
}
|
||||
}
|
||||
|
||||
impl<'a> GenerateModel for HunyuanOCRGenerateModel<'a> {
|
||||
fn generate(&mut self, mes: ChatCompletionParameters) -> Result<ChatCompletionResponse> {
|
||||
let temperature = match mes.temperature {
|
||||
None => self.generation_config.temperature,
|
||||
Some(tem) => tem,
|
||||
};
|
||||
let top_p = match mes.top_p {
|
||||
None => self.generation_config.top_p,
|
||||
Some(top_p) => top_p,
|
||||
};
|
||||
let top_k = self.generation_config.top_k;
|
||||
let seed = match mes.seed {
|
||||
None => 34562u64,
|
||||
Some(s) => s as u64,
|
||||
};
|
||||
let mut logit_processor =
|
||||
get_logit_processor(Some(temperature), Some(top_p), Some(top_k), seed);
|
||||
let mes_render = self.chat_template.apply_chat_template(&mes)?;
|
||||
let data = self
|
||||
.pre_processor
|
||||
.process_info(&mes, &self.tokenizer, &mes_render)?;
|
||||
let mut input_ids = data.input_ids;
|
||||
let mut position_ids = Some(&data.position_ids);
|
||||
let mut image_mask = Some(&data.image_mask);
|
||||
let mut pixel_values = data.pixel_values;
|
||||
let mut image_grid_thw = data.image_grid_thw;
|
||||
let mut seq_len = input_ids.dim(1)?;
|
||||
let mut seqlen_offset = 0;
|
||||
let mut generate: Vec<u32> = Vec::new();
|
||||
let sample_len = mes.max_tokens.unwrap_or(1024);
|
||||
for _ in 0..sample_len {
|
||||
let logits = self.hunyuan_vl.forward(
|
||||
&input_ids,
|
||||
pixel_values.as_ref(),
|
||||
image_grid_thw.as_ref(),
|
||||
image_mask,
|
||||
position_ids,
|
||||
seqlen_offset,
|
||||
)?;
|
||||
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 == self.eos_token_id1 || next_token == self.eos_token_id2 {
|
||||
break;
|
||||
}
|
||||
seqlen_offset += seq_len;
|
||||
seq_len = 1;
|
||||
input_ids = Tensor::from_vec(vec![next_token], (1, 1), &self.device)?;
|
||||
position_ids = None;
|
||||
image_mask = None;
|
||||
pixel_values = None;
|
||||
image_grid_thw = None;
|
||||
}
|
||||
let res = self.tokenizer.token_decode(generate)?;
|
||||
self.hunyuan_vl.clear_kv_cache();
|
||||
let response = build_completion_response(res, &self.model_name);
|
||||
Ok(response)
|
||||
}
|
||||
|
||||
fn generate_stream(
|
||||
&mut self,
|
||||
mes: ChatCompletionParameters,
|
||||
) -> Result<
|
||||
Box<
|
||||
dyn Stream<Item = Result<ChatCompletionChunkResponse, anyhow::Error>>
|
||||
+ Send
|
||||
+ Unpin
|
||||
+ '_,
|
||||
>,
|
||||
> {
|
||||
let temperature = match mes.temperature {
|
||||
None => self.generation_config.temperature,
|
||||
Some(tem) => tem,
|
||||
};
|
||||
let top_p = match mes.top_p {
|
||||
None => self.generation_config.top_p,
|
||||
Some(top_p) => top_p,
|
||||
};
|
||||
let top_k = self.generation_config.top_k;
|
||||
let seed = match mes.seed {
|
||||
None => 34562u64,
|
||||
Some(s) => s as u64,
|
||||
};
|
||||
let mut logit_processor =
|
||||
get_logit_processor(Some(temperature), Some(top_p), Some(top_k), seed);
|
||||
let mes_render = self.chat_template.apply_chat_template(&mes)?;
|
||||
let data = self
|
||||
.pre_processor
|
||||
.process_info(&mes, &self.tokenizer, &mes_render)?;
|
||||
|
||||
let mut seqlen_offset = 0;
|
||||
let sample_len = mes.max_tokens.unwrap_or(1024);
|
||||
let stream = stream! {
|
||||
let mut error_tokens = Vec::new();
|
||||
let mut input_ids = data.input_ids;
|
||||
let mut position_ids = Some(&data.position_ids);
|
||||
let mut image_mask = Some(&data.image_mask);
|
||||
let mut pixel_values = data.pixel_values;
|
||||
let mut image_grid_thw = data.image_grid_thw;
|
||||
let mut seq_len = input_ids.dim(1)?;
|
||||
for _ in 0..sample_len {
|
||||
let logits = self.hunyuan_vl.forward(
|
||||
&input_ids,
|
||||
pixel_values.as_ref(),
|
||||
image_grid_thw.as_ref(),
|
||||
image_mask,
|
||||
position_ids,
|
||||
seqlen_offset,
|
||||
)?;
|
||||
let logits = logits.squeeze(0)?.squeeze(0)?.to_dtype(DType::F32)?;
|
||||
let next_token = logit_processor.sample(&logits)?;
|
||||
let mut decode_ids = Vec::new();
|
||||
if !error_tokens.is_empty() {
|
||||
decode_ids.extend_from_slice(&error_tokens);
|
||||
}
|
||||
decode_ids.push(next_token);
|
||||
let decoded_token = self.tokenizer.token_decode(decode_ids).map_err(|e| anyhow!(format!("stream decode error{}", e)))?;
|
||||
if decoded_token.contains("�") {
|
||||
error_tokens.push(next_token);
|
||||
if error_tokens.len() > 3 {
|
||||
error_tokens.clear();
|
||||
}
|
||||
seqlen_offset += seq_len;
|
||||
seq_len = 1;
|
||||
input_ids = Tensor::from_vec(vec![next_token], (1, 1), &self.device)?;
|
||||
position_ids = None;
|
||||
image_mask = None;
|
||||
pixel_values = None;
|
||||
image_grid_thw = None;
|
||||
continue;
|
||||
}
|
||||
error_tokens.clear();
|
||||
let chunk = build_completion_chunk_response(decoded_token, &self.model_name, None, None);
|
||||
yield Ok(chunk);
|
||||
if next_token == self.eos_token_id1 || next_token == self.eos_token_id2 {
|
||||
break;
|
||||
}
|
||||
seqlen_offset += seq_len;
|
||||
seq_len = 1;
|
||||
input_ids = Tensor::from_vec(vec![next_token], (1, 1), &self.device)?;
|
||||
position_ids = None;
|
||||
image_mask = None;
|
||||
pixel_values = None;
|
||||
image_grid_thw = None;
|
||||
}
|
||||
self.hunyuan_vl.clear_kv_cache();
|
||||
};
|
||||
Ok(Box::new(Box::pin(stream)))
|
||||
}
|
||||
}
|
||||
@@ -0,0 +1,4 @@
|
||||
pub mod config;
|
||||
pub mod generate;
|
||||
pub mod model;
|
||||
pub mod processor;
|
||||
@@ -0,0 +1,621 @@
|
||||
use anyhow::{Result, anyhow};
|
||||
use candle_core::{D, IndexOp, Tensor};
|
||||
use candle_nn::{
|
||||
Conv2d, Embedding, Init, LayerNorm, Linear, Module, RmsNorm, VarBuilder, embedding, linear,
|
||||
linear_no_bias, rms_norm,
|
||||
};
|
||||
|
||||
use crate::{
|
||||
models::{
|
||||
common::{
|
||||
GateUpDownMLP, NaiveAttention, TwoLinearMLP, eager_attention_forward, get_conv2d,
|
||||
get_layer_norm,
|
||||
},
|
||||
hunyuan_ocr::config::{HunYuanVLConfig, HunYuanVLVisionConfig},
|
||||
},
|
||||
position_embed::rope::{RoPE, apply_rotary_pos_emb, get_xd_cos_sin},
|
||||
utils::tensor_utils::{
|
||||
interpolate_bilinear, masked_scatter_dim0, prepare_causal_attention_mask, split_tensor,
|
||||
},
|
||||
};
|
||||
|
||||
pub struct HunYuanVisionPatchEmbed {
|
||||
patch_embedding: Conv2d,
|
||||
// position_embedding: Embedding,
|
||||
num_channels: usize,
|
||||
patch_size: usize,
|
||||
// num_positions: usize,
|
||||
// position_edge: usize,
|
||||
embed_dim: usize,
|
||||
patch_pos_embed: Tensor,
|
||||
}
|
||||
|
||||
impl HunYuanVisionPatchEmbed {
|
||||
pub fn new(vb: VarBuilder, config: &HunYuanVLVisionConfig) -> Result<Self> {
|
||||
let patch_embedding = get_conv2d(
|
||||
vb.pp("patch_embedding"),
|
||||
config.num_channels,
|
||||
config.hidden_size,
|
||||
config.patch_size,
|
||||
0,
|
||||
config.patch_size,
|
||||
1,
|
||||
1,
|
||||
true,
|
||||
)?;
|
||||
let num_channels = config.num_channels;
|
||||
let patch_size = config.patch_size;
|
||||
let position_edge = config.max_image_size / patch_size;
|
||||
let num_positions = (position_edge).pow(2) + 1;
|
||||
let embed_dim = config.hidden_size;
|
||||
let position_embedding = embedding(num_positions, embed_dim, vb.pp("position_embedding"))?;
|
||||
let patch_pos_embed = position_embedding
|
||||
.embeddings()
|
||||
.i(1..)?
|
||||
.reshape((1, position_edge, position_edge, embed_dim))?
|
||||
.permute((0, 3, 1, 2))?;
|
||||
Ok(Self {
|
||||
patch_embedding,
|
||||
// position_embedding,
|
||||
num_channels,
|
||||
patch_size,
|
||||
// num_positions,
|
||||
// position_edge,
|
||||
embed_dim,
|
||||
patch_pos_embed,
|
||||
})
|
||||
}
|
||||
|
||||
pub fn forward(&self, pixel_values: &Tensor, grid_thw: &Tensor) -> Result<Tensor> {
|
||||
let (num_patches, _) = pixel_values.dims2()?;
|
||||
let pixel_values = pixel_values.reshape((
|
||||
num_patches,
|
||||
self.num_channels,
|
||||
self.patch_size,
|
||||
self.patch_size,
|
||||
))?;
|
||||
let patch_embeds = self.patch_embedding.forward(&pixel_values)?;
|
||||
let patch_embeds = patch_embeds
|
||||
.squeeze(D::Minus1)?
|
||||
.squeeze(D::Minus1)?
|
||||
.unsqueeze(0)?;
|
||||
let mut patch_pos_embed_list = vec![];
|
||||
let img_num = grid_thw.dim(0)?;
|
||||
for i in 0..img_num {
|
||||
let grid_i = grid_thw.i(i)?;
|
||||
let grid_h = grid_i.i(1)?.to_scalar::<u32>()? as usize;
|
||||
let grid_w = grid_i.i(2)?.to_scalar::<u32>()? as usize;
|
||||
let patch_pos_embed_ =
|
||||
interpolate_bilinear(&self.patch_pos_embed, (grid_h, grid_w), Some(false))?;
|
||||
let patch_pos_embed_ = patch_pos_embed_
|
||||
.reshape((self.embed_dim, ()))?
|
||||
.transpose(0, 1)?
|
||||
.unsqueeze(0)?;
|
||||
patch_pos_embed_list.push(patch_pos_embed_);
|
||||
}
|
||||
let patch_pos_embed = Tensor::cat(&patch_pos_embed_list, 1)?;
|
||||
let embedding = patch_embeds.add(&patch_pos_embed)?;
|
||||
Ok(embedding)
|
||||
}
|
||||
}
|
||||
pub struct HunYuanVisionBlock {
|
||||
self_attn: NaiveAttention,
|
||||
mlp: TwoLinearMLP,
|
||||
input_layernorm: LayerNorm,
|
||||
post_attention_layernorm: LayerNorm,
|
||||
}
|
||||
|
||||
impl HunYuanVisionBlock {
|
||||
pub fn new(vb: VarBuilder, config: &HunYuanVLVisionConfig) -> Result<Self> {
|
||||
let self_attn = NaiveAttention::new(
|
||||
vb.pp("self_attn"),
|
||||
config.hidden_size,
|
||||
config.num_attention_heads,
|
||||
config.num_attention_heads,
|
||||
true,
|
||||
)?;
|
||||
let mlp = TwoLinearMLP::new(
|
||||
vb.pp("mlp"),
|
||||
config.hidden_size,
|
||||
config.intermediate_size,
|
||||
config.hidden_act,
|
||||
true,
|
||||
"dense_h_to_4h",
|
||||
"dense_4h_to_h",
|
||||
)?;
|
||||
|
||||
let input_layernorm = get_layer_norm(
|
||||
vb.pp("input_layernorm"),
|
||||
config.rms_norm_eps,
|
||||
config.hidden_size,
|
||||
)?;
|
||||
let post_attention_layernorm = get_layer_norm(
|
||||
vb.pp("post_attention_layernorm"),
|
||||
config.rms_norm_eps,
|
||||
config.hidden_size,
|
||||
)?;
|
||||
Ok(Self {
|
||||
self_attn,
|
||||
mlp,
|
||||
input_layernorm,
|
||||
post_attention_layernorm,
|
||||
})
|
||||
}
|
||||
|
||||
pub fn forward(&self, xs: &Tensor) -> Result<Tensor> {
|
||||
let residual = xs.clone();
|
||||
let xs = self.input_layernorm.forward(xs)?;
|
||||
let xs = self.self_attn.forward(&xs, None, None, None, false)?;
|
||||
let residual = residual.add(&xs)?;
|
||||
let xs = self.post_attention_layernorm.forward(&residual)?;
|
||||
let xs = self.mlp.forward(&xs)?;
|
||||
let xs = residual.add(&xs)?;
|
||||
Ok(xs)
|
||||
}
|
||||
}
|
||||
|
||||
pub struct HunYuanVisionPatchMerger {
|
||||
proj_0: Conv2d,
|
||||
proj_2: Conv2d,
|
||||
mlp: Linear,
|
||||
image_newline: Tensor,
|
||||
image_begin: Tensor,
|
||||
image_end: Tensor,
|
||||
// image_sep: Tensor,
|
||||
before_rms: RmsNorm,
|
||||
after_rms: RmsNorm,
|
||||
}
|
||||
|
||||
impl HunYuanVisionPatchMerger {
|
||||
pub fn new(vb: VarBuilder, config: &HunYuanVLVisionConfig) -> Result<Self> {
|
||||
let proj_0 = get_conv2d(
|
||||
vb.pp("proj.0"),
|
||||
config.hidden_size,
|
||||
config.hidden_size * 2,
|
||||
config.spatial_merge_size,
|
||||
0,
|
||||
config.spatial_merge_size,
|
||||
1,
|
||||
1,
|
||||
true,
|
||||
)?;
|
||||
let proj_2 = get_conv2d(
|
||||
vb.pp("proj.2"),
|
||||
config.hidden_size * 2,
|
||||
config.hidden_size * 4,
|
||||
1,
|
||||
0,
|
||||
1,
|
||||
1,
|
||||
1,
|
||||
true,
|
||||
)?;
|
||||
let mlp = linear(config.hidden_size * 4, config.out_hidden_size, vb.pp("mlp"))?;
|
||||
let image_newline =
|
||||
vb.get_with_hints(config.hidden_size * 4, "image_newline", Init::Const(0.))?;
|
||||
let image_begin =
|
||||
vb.get_with_hints(config.out_hidden_size, "image_begin", Init::Const(0.))?;
|
||||
let image_end = vb.get_with_hints(config.out_hidden_size, "image_end", Init::Const(0.))?;
|
||||
// let image_sep = vb.get_with_hints(config.out_hidden_size, "image_sep", Init::Const(0.))?;
|
||||
let before_rms = rms_norm(config.hidden_size, config.rms_norm_eps, vb.pp("before_rms"))?;
|
||||
let after_rms = rms_norm(
|
||||
config.out_hidden_size,
|
||||
config.rms_norm_eps,
|
||||
vb.pp("after_rms"),
|
||||
)?;
|
||||
Ok(Self {
|
||||
proj_0,
|
||||
proj_2,
|
||||
mlp,
|
||||
image_newline,
|
||||
image_begin,
|
||||
image_end,
|
||||
// image_sep,
|
||||
before_rms,
|
||||
after_rms,
|
||||
})
|
||||
}
|
||||
pub fn forward(&self, xs: &Tensor, size: (usize, usize)) -> Result<Tensor> {
|
||||
let xs = self.before_rms.forward(xs)?;
|
||||
let (h, w) = size;
|
||||
let xs = xs.permute((0, 2, 1))?.reshape((xs.dim(0)?, (), h, w))?;
|
||||
let xs = self.proj_0.forward(&xs)?.gelu()?;
|
||||
let xs = self.proj_2.forward(&xs)?;
|
||||
let (b, c, h, _) = xs.dims4()?;
|
||||
let image_newline = self
|
||||
.image_newline
|
||||
.reshape((1, c, 1, 1))?
|
||||
.broadcast_as((b, c, h, 1))?
|
||||
.to_dtype(xs.dtype())?;
|
||||
let xs = Tensor::cat(&[xs, image_newline], D::Minus1)?;
|
||||
let xs = xs.reshape((b, c, ()))?.permute((0, 2, 1))?;
|
||||
let xs = self.mlp.forward(&xs)?;
|
||||
let begin = self
|
||||
.image_begin
|
||||
.reshape((1, 1, ()))?
|
||||
.broadcast_as((b, 1, xs.dim(D::Minus1)?))?
|
||||
.to_dtype(xs.dtype())?;
|
||||
let end = self
|
||||
.image_end
|
||||
.reshape((1, 1, ()))?
|
||||
.broadcast_as((b, 1, xs.dim(D::Minus1)?))?
|
||||
.to_dtype(xs.dtype())?;
|
||||
let xs = Tensor::cat(&[begin, xs, end], 1)?;
|
||||
let xs = self.after_rms.forward(&xs)?;
|
||||
Ok(xs)
|
||||
}
|
||||
}
|
||||
|
||||
pub struct HunYuanVisionTransformer {
|
||||
embeddings: HunYuanVisionPatchEmbed,
|
||||
layers: Vec<HunYuanVisionBlock>,
|
||||
perceive: HunYuanVisionPatchMerger,
|
||||
}
|
||||
|
||||
impl HunYuanVisionTransformer {
|
||||
pub fn new(vb: VarBuilder, config: &HunYuanVLVisionConfig) -> Result<Self> {
|
||||
let embeddings = HunYuanVisionPatchEmbed::new(vb.pp("embeddings"), config)?;
|
||||
let mut layers = vec![];
|
||||
let vb_layers = vb.pp("layers");
|
||||
for i in 0..config.num_hidden_layers {
|
||||
let layer_i = HunYuanVisionBlock::new(vb_layers.pp(i), config)?;
|
||||
layers.push(layer_i);
|
||||
}
|
||||
let perceive = HunYuanVisionPatchMerger::new(vb.pp("perceive"), config)?;
|
||||
Ok(Self {
|
||||
embeddings,
|
||||
layers,
|
||||
perceive,
|
||||
})
|
||||
}
|
||||
|
||||
pub fn forward(&self, xs: &Tensor, grid_thw: &Tensor) -> Result<Tensor> {
|
||||
let mut hidden_states = self.embeddings.forward(xs, grid_thw)?;
|
||||
for layer in &self.layers {
|
||||
hidden_states = layer.forward(&hidden_states)?;
|
||||
}
|
||||
let mut cu_seqlens = vec![];
|
||||
for i in 0..grid_thw.dim(0)? {
|
||||
let [_, h, w] = grid_thw.i(i)?.to_vec1::<u32>()?[..] else {
|
||||
return Err(anyhow!(format!("grid_thw Expected exactly 3 elements")));
|
||||
};
|
||||
cu_seqlens.push((h * w) as usize);
|
||||
}
|
||||
let split_items = split_tensor(&hidden_states, &cu_seqlens, 1)?;
|
||||
let mut processed_item = vec![];
|
||||
for i in 0..grid_thw.dim(0)? {
|
||||
let [_, h, w] = grid_thw.i(i)?.to_vec1::<u32>()?[..] else {
|
||||
return Err(anyhow!(format!("grid_thw Expected exactly 3 elements")));
|
||||
};
|
||||
let processed = self
|
||||
.perceive
|
||||
.forward(&split_items[i], (h as usize, w as usize))?;
|
||||
processed_item.push(processed);
|
||||
}
|
||||
let xs = Tensor::cat(&processed_item, 1)?;
|
||||
Ok(xs)
|
||||
}
|
||||
}
|
||||
|
||||
pub struct HunYuanVLAttention {
|
||||
q_proj: Linear,
|
||||
k_proj: Linear,
|
||||
v_proj: Linear,
|
||||
o_proj: Linear,
|
||||
query_layernorm: RmsNorm,
|
||||
key_layernorm: RmsNorm,
|
||||
num_attention_heads: usize,
|
||||
num_key_value_heads: usize,
|
||||
num_kv_groups: usize,
|
||||
head_dim: usize,
|
||||
scaling: f64,
|
||||
kv_cache: Option<(Tensor, Tensor)>,
|
||||
}
|
||||
|
||||
impl HunYuanVLAttention {
|
||||
pub fn new(
|
||||
vb: VarBuilder,
|
||||
hidden_size: usize,
|
||||
head_dim: usize,
|
||||
num_attention_heads: usize,
|
||||
num_key_value_heads: usize,
|
||||
attention_bias: bool,
|
||||
rms_norm_eps: f64,
|
||||
) -> Result<Self> {
|
||||
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 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 {
|
||||
q_proj,
|
||||
k_proj,
|
||||
v_proj,
|
||||
o_proj,
|
||||
query_layernorm,
|
||||
key_layernorm,
|
||||
num_attention_heads,
|
||||
num_key_value_heads,
|
||||
num_kv_groups,
|
||||
head_dim,
|
||||
scaling,
|
||||
kv_cache: None,
|
||||
})
|
||||
}
|
||||
|
||||
pub fn forward(
|
||||
&mut self,
|
||||
xs: &Tensor,
|
||||
cos: &Tensor,
|
||||
sin: &Tensor,
|
||||
attention_mask: Option<&Tensor>,
|
||||
) -> Result<Tensor> {
|
||||
let (b_sz, q_len, _) = xs.dims3()?;
|
||||
let query_states = self
|
||||
.q_proj
|
||||
.forward(xs)?
|
||||
.reshape((b_sz, q_len, self.num_attention_heads, self.head_dim))?
|
||||
.transpose(1, 2)?;
|
||||
|
||||
let key_states = self
|
||||
.k_proj
|
||||
.forward(xs)?
|
||||
.reshape((b_sz, q_len, self.num_key_value_heads, self.head_dim))?
|
||||
.transpose(1, 2)?;
|
||||
let value_states = self.v_proj.forward(xs)?;
|
||||
let value_states = value_states
|
||||
.reshape((b_sz, q_len, self.num_key_value_heads, self.head_dim))?
|
||||
.transpose(1, 2)?;
|
||||
let (query_states, key_states) =
|
||||
apply_rotary_pos_emb(&query_states, &key_states, cos, sin, false)?;
|
||||
let query_states = self.query_layernorm.forward(&query_states)?;
|
||||
let key_states = self.key_layernorm.forward(&key_states)?;
|
||||
let (key_states, value_states) = match &self.kv_cache {
|
||||
None => (key_states, value_states),
|
||||
Some((prev_k, prev_v)) => {
|
||||
let key_states = Tensor::cat(&[prev_k, &key_states], 2)?;
|
||||
let value_states = Tensor::cat(&[prev_v, &value_states], 2)?;
|
||||
(key_states, value_states)
|
||||
}
|
||||
};
|
||||
self.kv_cache = Some((key_states.clone(), value_states.clone()));
|
||||
let attn_output = eager_attention_forward(
|
||||
&query_states,
|
||||
&key_states,
|
||||
&value_states,
|
||||
Some(self.num_kv_groups),
|
||||
attention_mask,
|
||||
self.scaling,
|
||||
)?;
|
||||
let attn_output =
|
||||
attn_output.reshape((b_sz, q_len, self.num_attention_heads * self.head_dim))?;
|
||||
let attn_output = attn_output.apply(&self.o_proj)?;
|
||||
Ok(attn_output)
|
||||
}
|
||||
|
||||
pub fn clear_kv_cache(&mut self) {
|
||||
self.kv_cache = None
|
||||
}
|
||||
}
|
||||
|
||||
pub struct HunYuanVLDecoderLayer {
|
||||
self_attn: HunYuanVLAttention,
|
||||
mlp: GateUpDownMLP,
|
||||
input_layernorm: RmsNorm,
|
||||
post_attention_layernorm: RmsNorm,
|
||||
}
|
||||
|
||||
impl HunYuanVLDecoderLayer {
|
||||
pub fn new(config: &HunYuanVLConfig, vb: VarBuilder) -> Result<Self> {
|
||||
let self_attn = HunYuanVLAttention::new(
|
||||
vb.pp("self_attn"),
|
||||
config.hidden_size,
|
||||
config.head_dim,
|
||||
config.num_attention_heads,
|
||||
config.num_key_value_heads,
|
||||
config.attention_bias,
|
||||
config.rms_norm_eps,
|
||||
)?;
|
||||
let mlp = GateUpDownMLP::new(
|
||||
vb.pp("mlp"),
|
||||
config.hidden_size,
|
||||
config.intermediate_size,
|
||||
config.hidden_act,
|
||||
false,
|
||||
)?;
|
||||
let input_layernorm = rms_norm(
|
||||
config.hidden_size,
|
||||
config.rms_norm_eps,
|
||||
vb.pp("input_layernorm"),
|
||||
)?;
|
||||
let post_attention_layernorm = rms_norm(
|
||||
config.hidden_size,
|
||||
config.rms_norm_eps,
|
||||
vb.pp("post_attention_layernorm"),
|
||||
)?;
|
||||
Ok(Self {
|
||||
self_attn,
|
||||
mlp,
|
||||
input_layernorm,
|
||||
post_attention_layernorm,
|
||||
})
|
||||
}
|
||||
|
||||
pub fn forward(
|
||||
&mut self,
|
||||
xs: &Tensor,
|
||||
cos: &Tensor,
|
||||
sin: &Tensor,
|
||||
attention_mask: Option<&Tensor>,
|
||||
) -> Result<Tensor> {
|
||||
let residual = xs.clone();
|
||||
let xs = self.input_layernorm.forward(xs)?;
|
||||
let xs = self.self_attn.forward(&xs, cos, sin, attention_mask)?;
|
||||
let xs = residual.add(&xs)?;
|
||||
let residual = xs.clone();
|
||||
let xs = self.post_attention_layernorm.forward(&xs)?;
|
||||
let xs = self.mlp.forward(&xs)?;
|
||||
let xs = residual.add(&xs)?;
|
||||
Ok(xs)
|
||||
}
|
||||
|
||||
pub fn clear_kv_cache(&mut self) {
|
||||
self.self_attn.clear_kv_cache();
|
||||
}
|
||||
}
|
||||
|
||||
pub struct HunYuanVLTextModel {
|
||||
embed_tokens: Embedding,
|
||||
layers: Vec<HunYuanVLDecoderLayer>,
|
||||
norm: RmsNorm,
|
||||
rope: RoPE,
|
||||
xdrope_section: Vec<usize>,
|
||||
}
|
||||
|
||||
impl HunYuanVLTextModel {
|
||||
pub fn new(vb: VarBuilder, config: &HunYuanVLConfig) -> Result<Self> {
|
||||
let embed_tokens = embedding(config.vocab_size, config.hidden_size, vb.pp("embed_tokens"))?;
|
||||
let mut layers = vec![];
|
||||
let vb_layers = vb.pp("layers");
|
||||
for i in 0..config.num_hidden_layers {
|
||||
let layer = HunYuanVLDecoderLayer::new(config, vb_layers.pp(i))?;
|
||||
layers.push(layer);
|
||||
}
|
||||
let norm = rms_norm(config.hidden_size, config.rms_norm_eps, vb.pp("norm"))?;
|
||||
let base = config.rope_theta
|
||||
* config
|
||||
.rope_scaling
|
||||
.alpha
|
||||
.powf(config.head_dim as f64 / (config.head_dim - 2) as f64);
|
||||
let rope = RoPE::new(config.head_dim, base as f32, vb.device())?;
|
||||
let xdrope_section = config.rope_scaling.xdrope_section.clone();
|
||||
Ok(Self {
|
||||
embed_tokens,
|
||||
layers,
|
||||
norm,
|
||||
rope,
|
||||
xdrope_section,
|
||||
})
|
||||
}
|
||||
|
||||
pub fn forward(
|
||||
&mut self,
|
||||
inputs_embeds: &Tensor,
|
||||
position_ids: Option<&Tensor>,
|
||||
seqlen_offset: usize,
|
||||
) -> Result<Tensor> {
|
||||
let (b_size, seq_len, _) = inputs_embeds.dims3()?;
|
||||
|
||||
// let position_ids = match position_ids {
|
||||
// Some(ids) => ids.clone(),
|
||||
// None => Tensor::arange(
|
||||
// seqlen_offset as u32,
|
||||
// (seq_len + seqlen_offset) as u32,
|
||||
// inputs_embeds.device(),
|
||||
// )?
|
||||
// .unsqueeze(0)?,
|
||||
// };
|
||||
let attention_mask: Option<&Tensor> = {
|
||||
if seq_len <= 1 {
|
||||
None
|
||||
} else {
|
||||
Some(&prepare_causal_attention_mask(
|
||||
b_size,
|
||||
seq_len,
|
||||
0,
|
||||
inputs_embeds.device(),
|
||||
)?)
|
||||
}
|
||||
};
|
||||
|
||||
let (cos, sin) = self
|
||||
.rope
|
||||
.forward(seqlen_offset, seq_len, inputs_embeds.device())?;
|
||||
let mut xs = inputs_embeds.clone();
|
||||
for (i, layer) in self.layers.iter_mut().enumerate() {
|
||||
if i == 0
|
||||
&& let Some(position_ids) = position_ids
|
||||
{
|
||||
let (cos, sin) =
|
||||
get_xd_cos_sin(&cos, &sin, position_ids, self.xdrope_section.clone())?;
|
||||
xs = layer.forward(&xs, &cos, &sin, attention_mask)?;
|
||||
} else {
|
||||
xs = layer.forward(&xs, &cos, &sin, attention_mask)?;
|
||||
}
|
||||
}
|
||||
let xs = self.norm.forward(&xs)?;
|
||||
Ok(xs)
|
||||
}
|
||||
|
||||
pub fn clear_kv_cache(&mut self) {
|
||||
for layer in self.layers.iter_mut() {
|
||||
layer.clear_kv_cache()
|
||||
}
|
||||
}
|
||||
}
|
||||
|
||||
pub struct HunyuanVLModel {
|
||||
// config: HunYuanVLConfig,
|
||||
vit: HunYuanVisionTransformer,
|
||||
model: HunYuanVLTextModel,
|
||||
lm_head: Linear,
|
||||
}
|
||||
|
||||
impl HunyuanVLModel {
|
||||
pub fn new(vb: VarBuilder, config: HunYuanVLConfig) -> Result<Self> {
|
||||
let vit = HunYuanVisionTransformer::new(vb.pp("vit"), &config.vision_config)?;
|
||||
let model = HunYuanVLTextModel::new(vb.pp("model"), &config)?;
|
||||
let lm_head = Linear::new(model.embed_tokens.embeddings().clone(), None);
|
||||
Ok(Self {
|
||||
// config,
|
||||
vit,
|
||||
model,
|
||||
lm_head,
|
||||
})
|
||||
}
|
||||
pub fn forward(
|
||||
&mut self,
|
||||
input_ids: &Tensor,
|
||||
pixel_values: Option<&Tensor>,
|
||||
image_grid_thw: Option<&Tensor>,
|
||||
image_mask: Option<&Tensor>,
|
||||
position_ids: Option<&Tensor>,
|
||||
seqlen_offset: usize,
|
||||
) -> Result<Tensor> {
|
||||
let mut inputs_embeds = self.model.embed_tokens.forward(input_ids)?;
|
||||
if let Some(pixel_values) = pixel_values
|
||||
&& let Some(grid_thw) = image_grid_thw
|
||||
&& let Some(image_mask) = image_mask
|
||||
{
|
||||
let image_embeds = self.vit.forward(pixel_values, grid_thw)?.squeeze(0)?;
|
||||
inputs_embeds = masked_scatter_dim0(&inputs_embeds, &image_embeds, image_mask)?;
|
||||
}
|
||||
let outputs = self
|
||||
.model
|
||||
.forward(&inputs_embeds, position_ids, seqlen_offset)?;
|
||||
let seq_len = outputs.dim(1)?;
|
||||
let hidden_state = outputs.narrow(1, seq_len - 1, 1)?;
|
||||
let logits = self.lm_head.forward(&hidden_state)?;
|
||||
Ok(logits)
|
||||
}
|
||||
|
||||
pub fn clear_kv_cache(&mut self) {
|
||||
self.model.clear_kv_cache();
|
||||
}
|
||||
}
|
||||
@@ -0,0 +1,236 @@
|
||||
use aha_openai_dive::v1::resources::chat::ChatCompletionParameters;
|
||||
use anyhow::Result;
|
||||
use candle_core::{DType, Device, IndexOp, Shape, Tensor};
|
||||
use image::DynamicImage;
|
||||
|
||||
use crate::{
|
||||
models::hunyuan_ocr::config::HunyuanOCRPreprocessorConfig,
|
||||
tokenizer::TokenizerModel,
|
||||
utils::{
|
||||
img_utils::{extract_images, img_smart_resize, img_transform},
|
||||
tensor_utils::{get_eq_indices, get_equal_mask},
|
||||
},
|
||||
};
|
||||
|
||||
pub struct HunyuanVLProcessor {
|
||||
image_token_id: u32,
|
||||
image_token: String,
|
||||
placeholder_token: String,
|
||||
process_cfg: HunyuanOCRPreprocessorConfig,
|
||||
device: Device,
|
||||
dtype: DType,
|
||||
}
|
||||
|
||||
impl HunyuanVLProcessor {
|
||||
pub fn new(path: &str, device: &Device, dtype: DType) -> Result<Self> {
|
||||
let path = path.to_string();
|
||||
assert!(
|
||||
std::path::Path::new(&path).exists(),
|
||||
"model path file not exists"
|
||||
);
|
||||
let process_cfg_file = path.clone() + "/preprocessor_config.json";
|
||||
assert!(
|
||||
std::path::Path::new(&process_cfg_file).exists(),
|
||||
"preprocessor_config.json not exists in model path"
|
||||
);
|
||||
let process_cfg: HunyuanOCRPreprocessorConfig =
|
||||
serde_json::from_slice(&std::fs::read(process_cfg_file)?)?;
|
||||
let image_token_id = 120120u32;
|
||||
let image_token = "<|hy_place▁holder▁no▁102|>".to_string();
|
||||
let placeholder_token = "<|hy_place▁holder▁no▁799|>".to_string();
|
||||
// let pad_id = 120002u32;
|
||||
Ok(Self {
|
||||
image_token_id,
|
||||
image_token,
|
||||
placeholder_token,
|
||||
process_cfg,
|
||||
device: device.clone(),
|
||||
dtype,
|
||||
})
|
||||
}
|
||||
|
||||
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成 32的倍数
|
||||
let (resize_h, resize_w) = img_smart_resize(
|
||||
img_h,
|
||||
img_w,
|
||||
(self.process_cfg.patch_size * self.process_cfg.merge_size) as u32,
|
||||
self.process_cfg.min_pixels as u32,
|
||||
self.process_cfg.max_pixels as u32,
|
||||
)?;
|
||||
let img = img.resize_exact(resize_w, resize_h, image::imageops::FilterType::CatmullRom);
|
||||
let img_tensor = img_transform(&img, img_mean, img_std, &self.device, self.dtype)?;
|
||||
// (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)?;
|
||||
// img_temsor.dim[0] = 1, temporal_patch_size = 1, grid_t = 1
|
||||
let grid_t = img_tensor.dim(0)? / self.process_cfg.temporal_patch_size;
|
||||
let grid_h = img_tensor.dim(2)? / self.process_cfg.patch_size;
|
||||
let grid_w = img_tensor.dim(3)? / self.process_cfg.patch_size;
|
||||
let shape = Shape::from(vec![
|
||||
grid_t,
|
||||
channel,
|
||||
grid_h / self.process_cfg.merge_size,
|
||||
self.process_cfg.merge_size,
|
||||
self.process_cfg.patch_size,
|
||||
grid_w / self.process_cfg.merge_size,
|
||||
self.process_cfg.merge_size,
|
||||
self.process_cfg.patch_size,
|
||||
]);
|
||||
let img_tensor = img_tensor.reshape(shape)?;
|
||||
// shape to // grid_t,
|
||||
// grid_h / merge_size,
|
||||
// merge_size,
|
||||
// grid_w / merge_size,
|
||||
// merge_size,
|
||||
// channel,
|
||||
// patch_size,
|
||||
// patch_size,
|
||||
let img_tensor = img_tensor.permute(vec![0, 2, 3, 5, 6, 1, 4, 7])?;
|
||||
let img_tensor = img_tensor
|
||||
.reshape((
|
||||
grid_t * grid_h * grid_w,
|
||||
channel * self.process_cfg.patch_size * self.process_cfg.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<(Tensor, Tensor)> {
|
||||
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, 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((pixel_values, vision_grid_thws))
|
||||
}
|
||||
|
||||
pub fn process_info(
|
||||
&self,
|
||||
messages: &ChatCompletionParameters,
|
||||
tokenizer: &TokenizerModel,
|
||||
text: &str,
|
||||
) -> Result<HunyuanData> {
|
||||
let imgs = extract_images(messages)?;
|
||||
let img_mean = Tensor::from_slice(&self.process_cfg.image_mean, (3, 1, 1), &self.device)?
|
||||
.to_dtype(self.dtype)?;
|
||||
let img_std = Tensor::from_slice(&self.process_cfg.image_std, (3, 1, 1), &self.device)?
|
||||
.to_dtype(self.dtype)?;
|
||||
let (pixel_values, image_grid_thw) = if !imgs.is_empty() {
|
||||
let (pixel_values, image_grid_thw) = self.process_images(&imgs, &img_mean, &img_std)?;
|
||||
(Some(pixel_values), Some(image_grid_thw))
|
||||
} else {
|
||||
(None, None)
|
||||
};
|
||||
|
||||
let mut image_tokens_cumsum = vec![0];
|
||||
let mut text = text.to_string();
|
||||
if !imgs.is_empty()
|
||||
&& let Some(grid_thw) = image_grid_thw.as_ref()
|
||||
{
|
||||
let mut index = 0;
|
||||
while text.contains(&self.image_token) {
|
||||
let grid_i = grid_thw.i(index)?;
|
||||
let grid_h = grid_i.i(1)?.to_scalar::<u32>()?;
|
||||
let grid_w = grid_i.i(2)?.to_scalar::<u32>()?;
|
||||
let patch_h = grid_h / self.process_cfg.merge_size as u32;
|
||||
let patch_w = grid_w / self.process_cfg.merge_size as u32;
|
||||
let num_image_tokens = patch_h * (patch_w + 1) + 2;
|
||||
let num_id = image_tokens_cumsum[image_tokens_cumsum.len() - 1] + num_image_tokens;
|
||||
image_tokens_cumsum.push(num_id);
|
||||
let replace = self.placeholder_token.repeat(num_image_tokens as usize);
|
||||
text = text.replacen(&self.image_token, &replace, 1);
|
||||
index += 1;
|
||||
}
|
||||
}
|
||||
|
||||
text = text.replace(&self.placeholder_token, &self.image_token);
|
||||
let input_ids = tokenizer.text_encode(text, &self.device)?;
|
||||
let seq_len = input_ids.dim(1)?;
|
||||
let position_ids = Tensor::arange(0, seq_len as u32, &self.device)?;
|
||||
let mut position_ids_w = Tensor::arange(0, seq_len as u32, &self.device)?;
|
||||
let mut position_ids_h = Tensor::arange(0, seq_len as u32, &self.device)?;
|
||||
let mut position_ids_t = Tensor::arange(0, seq_len as u32, &self.device)?;
|
||||
if !imgs.is_empty()
|
||||
&& let Some(grid_thw) = image_grid_thw.as_ref()
|
||||
{
|
||||
let image_token_pos_indices = get_eq_indices(&input_ids.i(0)?, self.image_token_id)?;
|
||||
for i in 0..grid_thw.dim(0)? {
|
||||
let grid_i = grid_thw.i(i)?;
|
||||
let grid_h = grid_i.i(1)?.to_scalar::<u32>()?;
|
||||
let grid_w = grid_i.i(2)?.to_scalar::<u32>()?;
|
||||
let patch_h = grid_h / self.process_cfg.merge_size as u32;
|
||||
let patch_w = grid_w / self.process_cfg.merge_size as u32;
|
||||
let start_pos = image_token_pos_indices
|
||||
.i(image_tokens_cumsum[i] as usize)?
|
||||
.to_scalar::<u32>()? as usize
|
||||
+ 1;
|
||||
let replace_num = ((patch_w + 1) * patch_h) as usize;
|
||||
let pos_w: Vec<u32> = (0..patch_h).flat_map(|_| 0u32..patch_w + 1).collect();
|
||||
position_ids_w = position_ids_w.slice_assign(
|
||||
&[start_pos..start_pos + replace_num],
|
||||
&Tensor::new(pos_w, &self.device)?,
|
||||
)?;
|
||||
let pos_h: Vec<u32> = (0..patch_h)
|
||||
.flat_map(|h| vec![h; (patch_w + 1) as usize])
|
||||
.collect();
|
||||
position_ids_h = position_ids_h.slice_assign(
|
||||
&[start_pos..start_pos + replace_num],
|
||||
&Tensor::new(pos_h, &self.device)?,
|
||||
)?;
|
||||
position_ids_t = position_ids_t.slice_assign(
|
||||
&[start_pos..start_pos + replace_num],
|
||||
&Tensor::new(vec![0u32; replace_num], &self.device)?,
|
||||
)?;
|
||||
}
|
||||
}
|
||||
let position_ids = Tensor::stack(
|
||||
&[position_ids, position_ids_h, position_ids_w, position_ids_t],
|
||||
0,
|
||||
)?
|
||||
.unsqueeze(0)?;
|
||||
let image_mask = get_equal_mask(&input_ids, self.image_token_id)?;
|
||||
let data = HunyuanData {
|
||||
input_ids,
|
||||
position_ids,
|
||||
image_mask,
|
||||
pixel_values,
|
||||
image_grid_thw,
|
||||
};
|
||||
Ok(data)
|
||||
}
|
||||
}
|
||||
|
||||
pub struct HunyuanData {
|
||||
pub input_ids: Tensor,
|
||||
pub position_ids: Tensor,
|
||||
pub image_mask: Tensor,
|
||||
pub pixel_values: Option<Tensor>,
|
||||
pub image_grid_thw: Option<Tensor>,
|
||||
}
|
||||
Reference in New Issue
Block a user