update readme and fmt
This commit is contained in:
@@ -112,9 +112,13 @@ cargo run -F cuda -- [参数]
|
|||||||
-m, --model <MODEL>
|
-m, --model <MODEL>
|
||||||
* 指定要加载的模型类型
|
* 指定要加载的模型类型
|
||||||
* 可选值:
|
* 可选值:
|
||||||
* minicpm4-0.5b:MiniCPM4-0.5B 模型
|
* minicpm4-0.5b:OpenBMB/MiniCPM4-0.5B 模型
|
||||||
* qwen2.5vl-3b:Qwen2.5-VL-3B 模型
|
* qwen2.5vl-3b:Qwen/Qwen2.5-VL-3B-Instruct 模型
|
||||||
* qwen3vl-2b:Qwen3-VL-2B 模型
|
* qwen2.5vl-7b:Qwen/Qwen2.5-VL-7B-Instruct 模型
|
||||||
|
* qwen3vl-2b:Qwen/Qwen3-VL-2B-Instruct 模型
|
||||||
|
* qwen3vl-4b:Qwen/Qwen3-VL-4B-Instruct 模型
|
||||||
|
* qwen3vl-8b:Qwen/Qwen3-VL-8B-Instruct 模型
|
||||||
|
* qwen3vl-32b:Qwen/Qwen3-VL-32B-Instruct 模型
|
||||||
* 示例:--model minicpm4-0.5b 或 -m qwen3vl-2b
|
* 示例:--model minicpm4-0.5b 或 -m qwen3vl-2b
|
||||||
|
|
||||||
3. 权重路径
|
3. 权重路径
|
||||||
|
|||||||
+5
-5
@@ -77,12 +77,12 @@ async fn main() -> anyhow::Result<()> {
|
|||||||
let args = Args::parse();
|
let args = Args::parse();
|
||||||
let model_id = match args.model {
|
let model_id = match args.model {
|
||||||
WhichModel::MiniCPM4_0_5B => "OpenBMB/MiniCPM4-0.5B",
|
WhichModel::MiniCPM4_0_5B => "OpenBMB/MiniCPM4-0.5B",
|
||||||
WhichModel::Qwen2_5vl3B => "Qwen/Qwen2.5-VL-3B-Instruct",
|
WhichModel::Qwen2_5vl3B => "Qwen/Qwen2.5-VL-3B-Instruct",
|
||||||
WhichModel::Qwen2_5vl7B => "Qwen/Qwen2.5-VL-7B-Instruct",
|
WhichModel::Qwen2_5vl7B => "Qwen/Qwen2.5-VL-7B-Instruct",
|
||||||
WhichModel::Qwen3vl2B => "Qwen/Qwen3-VL-2B-Instruct",
|
WhichModel::Qwen3vl2B => "Qwen/Qwen3-VL-2B-Instruct",
|
||||||
WhichModel::Qwen3vl4B => "Qwen/Qwen3-VL-4B-Instruct",
|
WhichModel::Qwen3vl4B => "Qwen/Qwen3-VL-4B-Instruct",
|
||||||
WhichModel::Qwen3vl8B => "Qwen/Qwen3-VL-8B-Instruct",
|
WhichModel::Qwen3vl8B => "Qwen/Qwen3-VL-8B-Instruct",
|
||||||
WhichModel::Qwen3vl32B => "Qwen/Qwen3-VL-32B-Instruct"
|
WhichModel::Qwen3vl32B => "Qwen/Qwen3-VL-32B-Instruct",
|
||||||
};
|
};
|
||||||
let model_path = match args.weight_path {
|
let model_path = match args.weight_path {
|
||||||
Some(path) => path,
|
Some(path) => path,
|
||||||
|
|||||||
+3
-3
@@ -99,15 +99,15 @@ pub fn load_model(model_type: WhichModel, path: &str) -> Result<ModelInstance<'_
|
|||||||
WhichModel::Qwen3vl2B => {
|
WhichModel::Qwen3vl2B => {
|
||||||
let model = Qwen3VLGenerateModel::init(path, None, None)?;
|
let model = Qwen3VLGenerateModel::init(path, None, None)?;
|
||||||
ModelInstance::Qwen3VL(model)
|
ModelInstance::Qwen3VL(model)
|
||||||
}
|
}
|
||||||
WhichModel::Qwen3vl4B => {
|
WhichModel::Qwen3vl4B => {
|
||||||
let model = Qwen3VLGenerateModel::init(path, None, None)?;
|
let model = Qwen3VLGenerateModel::init(path, None, None)?;
|
||||||
ModelInstance::Qwen3VL(model)
|
ModelInstance::Qwen3VL(model)
|
||||||
}
|
}
|
||||||
WhichModel::Qwen3vl8B => {
|
WhichModel::Qwen3vl8B => {
|
||||||
let model = Qwen3VLGenerateModel::init(path, None, None)?;
|
let model = Qwen3VLGenerateModel::init(path, None, None)?;
|
||||||
ModelInstance::Qwen3VL(model)
|
ModelInstance::Qwen3VL(model)
|
||||||
}
|
}
|
||||||
WhichModel::Qwen3vl32B => {
|
WhichModel::Qwen3vl32B => {
|
||||||
let model = Qwen3VLGenerateModel::init(path, None, None)?;
|
let model = Qwen3VLGenerateModel::init(path, None, None)?;
|
||||||
ModelInstance::Qwen3VL(model)
|
ModelInstance::Qwen3VL(model)
|
||||||
|
|||||||
@@ -559,7 +559,6 @@ pub struct Qwen3VLTextAttention {
|
|||||||
num_key_value_heads: usize,
|
num_key_value_heads: usize,
|
||||||
num_kv_groups: usize,
|
num_kv_groups: usize,
|
||||||
head_dim: usize,
|
head_dim: usize,
|
||||||
hidden_size: usize,
|
|
||||||
scaling: f64,
|
scaling: f64,
|
||||||
kv_cache: Option<(Tensor, Tensor)>,
|
kv_cache: Option<(Tensor, Tensor)>,
|
||||||
}
|
}
|
||||||
@@ -585,7 +584,8 @@ impl Qwen3VLTextAttention {
|
|||||||
linear_no_bias(hidden_size, num_key_value_heads * head_dim, vb.pp("k_proj"))?;
|
linear_no_bias(hidden_size, num_key_value_heads * head_dim, vb.pp("k_proj"))?;
|
||||||
let v_proj =
|
let v_proj =
|
||||||
linear_no_bias(hidden_size, num_key_value_heads * head_dim, vb.pp("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"))?;
|
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)
|
(q_proj, k_proj, v_proj, o_proj)
|
||||||
};
|
};
|
||||||
let q_norm = rms_norm(head_dim, config.rms_norm_eps, vb.pp("q_norm"))?;
|
let q_norm = rms_norm(head_dim, config.rms_norm_eps, vb.pp("q_norm"))?;
|
||||||
@@ -601,7 +601,6 @@ impl Qwen3VLTextAttention {
|
|||||||
num_key_value_heads,
|
num_key_value_heads,
|
||||||
num_kv_groups,
|
num_kv_groups,
|
||||||
head_dim,
|
head_dim,
|
||||||
hidden_size,
|
|
||||||
scaling,
|
scaling,
|
||||||
kv_cache: None,
|
kv_cache: None,
|
||||||
})
|
})
|
||||||
@@ -652,7 +651,8 @@ impl Qwen3VLTextAttention {
|
|||||||
attention_mask,
|
attention_mask,
|
||||||
self.scaling,
|
self.scaling,
|
||||||
)?;
|
)?;
|
||||||
let attn_output = attn_output.reshape((b_sz, q_len, self.num_attention_heads*self.head_dim))?;
|
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)?;
|
let attn_output = attn_output.apply(&self.o_proj)?;
|
||||||
Ok(attn_output)
|
Ok(attn_output)
|
||||||
}
|
}
|
||||||
|
|||||||
+1
-1
@@ -21,7 +21,7 @@ pub fn get_device(device: Option<&Device>) -> Device {
|
|||||||
None => {
|
None => {
|
||||||
#[cfg(feature = "cuda")]
|
#[cfg(feature = "cuda")]
|
||||||
{
|
{
|
||||||
Device::new_cuda(6).unwrap_or(Device::Cpu)
|
Device::new_cuda(0).unwrap_or(Device::Cpu)
|
||||||
}
|
}
|
||||||
#[cfg(not(feature = "cuda"))]
|
#[cfg(not(feature = "cuda"))]
|
||||||
{
|
{
|
||||||
|
|||||||
@@ -15,7 +15,10 @@ pub fn prepare_causal_attention_mask(
|
|||||||
let arange = Tensor::arange(0u32, tgt_len as u32, device)?;
|
let arange = Tensor::arange(0u32, tgt_len as u32, device)?;
|
||||||
let arange = arange.unsqueeze(1)?.broadcast_as((tgt_len, tgt_len))?;
|
let arange = arange.unsqueeze(1)?.broadcast_as((tgt_len, tgt_len))?;
|
||||||
let upper_triangle = arange.t()?.gt(&arange)?;
|
let upper_triangle = arange.t()?.gt(&arange)?;
|
||||||
let mask = upper_triangle.where_cond(&Tensor::new(f32::NEG_INFINITY, device)?.broadcast_as(arange.shape())?, &Tensor::new(0f32, device)?.broadcast_as(arange.shape())?)?;
|
let mask = upper_triangle.where_cond(
|
||||||
|
&Tensor::new(f32::NEG_INFINITY, device)?.broadcast_as(arange.shape())?,
|
||||||
|
&Tensor::new(0f32, device)?.broadcast_as(arange.shape())?,
|
||||||
|
)?;
|
||||||
|
|
||||||
let mask = if seqlen_offset > 0 {
|
let mask = if seqlen_offset > 0 {
|
||||||
let mask0 = Tensor::zeros((tgt_len, seqlen_offset), DType::F32, device)?;
|
let mask0 = Tensor::zeros((tgt_len, seqlen_offset), DType::F32, device)?;
|
||||||
|
|||||||
Reference in New Issue
Block a user