stash save

This commit is contained in:
jhqxxx
2026-05-09 12:18:16 +08:00
parent 5df1f30f83
commit e8bf980b38
39 changed files with 2505 additions and 316 deletions
Generated
+42 -3
View File
@@ -16,7 +16,7 @@ checksum = "b169f7a6d4742236a0a00c541b845991d0ac43e546831af1249753ab4c3aa3a0"
dependencies = [
"cfg-if",
"cipher",
"cpufeatures",
"cpufeatures 0.2.17",
]
[[package]]
@@ -42,6 +42,7 @@ dependencies = [
"minijinja",
"modelscope",
"num",
"rand 0.10.1",
"rayon",
"realfft",
"reqwest",
@@ -662,6 +663,17 @@ version = "0.2.1"
source = "registry+https://github.com/rust-lang/crates.io-index"
checksum = "613afe47fcd5fac7ccf1db93babcb082c5994d996f20b8b159f2ad1658eb5724"
[[package]]
name = "chacha20"
version = "0.10.0"
source = "registry+https://github.com/rust-lang/crates.io-index"
checksum = "6f8d983286843e49675a4b7a2d174efe136dc93a18d69130dd18198a6c167601"
dependencies = [
"cfg-if",
"cpufeatures 0.3.0",
"rand_core 0.10.1",
]
[[package]]
name = "cipher"
version = "0.4.4"
@@ -871,6 +883,15 @@ dependencies = [
"libc",
]
[[package]]
name = "cpufeatures"
version = "0.3.0"
source = "registry+https://github.com/rust-lang/crates.io-index"
checksum = "8b2a41393f66f16b0823bb79094d54ac5fbd34ab292ddafb9a0456ac9f87d201"
dependencies = [
"libc",
]
[[package]]
name = "crc"
version = "3.3.0"
@@ -1818,6 +1839,7 @@ dependencies = [
"cfg-if",
"libc",
"r-efi 6.0.0",
"rand_core 0.10.1",
"wasip2",
"wasip3",
]
@@ -3400,6 +3422,17 @@ dependencies = [
"rand_core 0.9.5",
]
[[package]]
name = "rand"
version = "0.10.1"
source = "registry+https://github.com/rust-lang/crates.io-index"
checksum = "d2e8e8bcc7961af1fdac401278c6a831614941f6164ee3bf4ce61b7edb162207"
dependencies = [
"chacha20",
"getrandom 0.4.2",
"rand_core 0.10.1",
]
[[package]]
name = "rand_chacha"
version = "0.3.1"
@@ -3438,6 +3471,12 @@ dependencies = [
"getrandom 0.3.4",
]
[[package]]
name = "rand_core"
version = "0.10.1"
source = "registry+https://github.com/rust-lang/crates.io-index"
checksum = "63b8176103e19a2643978565ca18b50549f6101881c443590420e4dc998a3c69"
[[package]]
name = "rand_distr"
version = "0.4.3"
@@ -4093,7 +4132,7 @@ source = "registry+https://github.com/rust-lang/crates.io-index"
checksum = "e3bf829a2d51ab4a5ddf1352d8470c140cadc8301b2ae1789db023f01cedd6ba"
dependencies = [
"cfg-if",
"cpufeatures",
"cpufeatures 0.2.17",
"digest",
]
@@ -4104,7 +4143,7 @@ source = "registry+https://github.com/rust-lang/crates.io-index"
checksum = "a7507d819769d01a365ab707794a4084392c824f54a7a6a7862f8c3d0892b283"
dependencies = [
"cfg-if",
"cpufeatures",
"cpufeatures 0.2.17",
"digest",
]
+1
View File
@@ -45,6 +45,7 @@ sentencepiece = "0.13.1"
ahash = "0.8.12"
derive_builder = "0.20.2"
kaldi-native-fbank = "0.1.0"
rand = "0.10.1"
[patch.crates-io]
esaxx-rs = { path = "vendor/esaxx-rs" }
Binary file not shown.
Binary file not shown.
Binary file not shown.
Binary file not shown.
+7 -1
View File
@@ -186,6 +186,8 @@ impl AMPBlock1 {
1,
1,
true,
None,
None,
)?;
convs1.push(layer);
}
@@ -203,6 +205,8 @@ impl AMPBlock1 {
1,
1,
true,
None,
None,
)?;
convs2.push(layer);
}
@@ -261,6 +265,8 @@ impl BigVGAN {
1,
1,
true,
None,
None,
)?;
let vb_ups = vb.pp("ups");
@@ -296,7 +302,7 @@ impl BigVGAN {
}
}
let activation_post = TorchActivation1d::new(vb.pp("activation_post"), 2, 2, 12, 12, ch)?;
let conv_post = WNConv1d::new(vb.pp("conv_post"), ch, 1, 7, 1, 3, 1, 1, false)?;
let conv_post = WNConv1d::new(vb.pp("conv_post"), ch, 1, 7, 1, 3, 1, 1, false, None, None)?;
Ok(Self {
num_kernels,
+8 -44
View File
@@ -1,12 +1,12 @@
use anyhow::Result;
use candle_core::{DType, Device, Tensor};
use candle_transformers::generation::{LogitsProcessor, Sampling};
use candle_transformers::generation::{LogitsProcessor};
use rocket::async_stream::stream;
use rocket::futures::Stream;
use std::time::Instant;
use crate::{
models::common::{InferenceModel, MultiModalData},
models::common::{InferenceModel, MultiModalData, sample::{use_repeat_penalty, get_logit_processor}},
params::chat::{ChatCompletionChunkResponse, ChatCompletionResponse},
tokenizer::TokenizerModel,
utils::response_utils::{
@@ -14,38 +14,6 @@ use crate::{
build_completion_chunk_response, build_completion_response_with_time,
},
};
pub fn get_logit_processor(
temperature: Option<f32>,
top_p: Option<f32>,
top_k: Option<usize>,
seed: u64,
) -> LogitsProcessor {
let temperature = temperature.and_then(|v| if v < 1e-7 { None } else { Some(v) });
match top_k {
None => LogitsProcessor::new(
seed,
temperature.map(|temp| temp as f64),
top_p.map(|tp| tp as f64),
),
Some(k) => {
let sampling = match temperature {
None => Sampling::ArgMax,
Some(temperature) => match top_p {
None => Sampling::TopK {
k,
temperature: temperature as f64,
},
Some(p) => Sampling::TopKThenTopP {
k,
p: p as f64,
temperature: temperature as f64,
},
},
};
LogitsProcessor::from_sampling(seed, sampling)
}
}
}
pub struct GenerationContext {
pub logit_processor: LogitsProcessor,
@@ -103,16 +71,12 @@ fn sample_and_push(
) -> Result<u32> {
let logits = logits.squeeze(0)?.squeeze(0)?.to_dtype(DType::F32)?;
// 重复惩罚
let logits = if ctx.repeat_penalty == 1. || ctx.repeat_last_n == 0 {
logits
} else {
let start_at = generated.len().saturating_sub(ctx.repeat_last_n);
candle_transformers::utils::apply_repeat_penalty(
&logits,
ctx.repeat_penalty,
&generated[start_at..],
)?
};
let logits = use_repeat_penalty(
ctx.repeat_penalty,
Some(ctx.repeat_last_n),
&logits,
generated,
)?;
let token = ctx.logit_processor.sample(&logits)?;
generated.push(token);
Ok(token)
+1
View File
@@ -6,6 +6,7 @@ pub mod gguf;
pub mod model_mapping;
pub mod modules;
pub mod reranker;
pub mod sample;
/// 多模态模型的特征数据
/// 每个模型数据不一样
+18 -170
View File
@@ -2,14 +2,14 @@ use anyhow::{Result, anyhow};
use candle_core::{D, DType, IndexOp, Tensor};
use candle_nn::{
Activation, BatchNorm, BatchNormConfig, Conv1d, Conv1dConfig, Conv2d, Conv2dConfig,
ConvTranspose1d, ConvTranspose1dConfig, Embedding, LayerNorm, LayerNormConfig, Linear, Module,
ModuleT, RmsNorm, VarBuilder, batch_norm, conv1d, conv1d_no_bias, conv2d, conv2d_no_bias,
embedding, layer_norm, linear_b, linear_no_bias, ops::sigmoid, rms_norm,
ConvTranspose1d, ConvTranspose1dConfig, LayerNorm, LayerNormConfig, Linear, Module, ModuleT,
RmsNorm, VarBuilder, batch_norm, conv1d, conv1d_no_bias, conv2d, conv2d_no_bias, layer_norm,
linear_b, ops::sigmoid, rms_norm,
};
use crate::{
position_embed::rope::{RoPE, apply_rotary_pos_emb, apply_rotary_pos_emb_roformer},
utils::tensor_utils::{pad_replicate_last_dim, prepare_causal_attention_mask, repeat_kv},
position_embed::rope::{apply_rotary_pos_emb, apply_rotary_pos_emb_roformer},
utils::tensor_utils::{pad_replicate_last_dim, repeat_kv},
};
#[derive(Debug)]
@@ -371,7 +371,7 @@ impl QKVCatAttention {
&& let Some(sin) = sin
{
if use_roformer {
apply_rotary_pos_emb_roformer(&query_states, &key_states, cos, sin, tof32)?
apply_rotary_pos_emb_roformer(&query_states, &key_states, cos, sin)?
} else {
apply_rotary_pos_emb(&query_states, &key_states, cos, sin, tof32)?
}
@@ -412,7 +412,7 @@ impl QKVCatAttention {
let key_states = qkv.i(1)?.contiguous()?;
let value_states = qkv.i(2)?.contiguous()?;
let (query_states, key_states) = if use_roformer {
apply_rotary_pos_emb_roformer(&query_states, &key_states, cos, sin, tof32)?
apply_rotary_pos_emb_roformer(&query_states, &key_states, cos, sin)?
} else {
apply_rotary_pos_emb(&query_states, &key_states, cos, sin, tof32)?
};
@@ -765,11 +765,11 @@ pub fn eager_attention_forward(
// input q shape:(b, num_head, seq_len, dim)
// input k/v shape:(b, num_kv_head, seq_len, dim)
let key_states = match num_key_value_groups {
Some(g) => repeat_kv(key_states.clone(), g)?.contiguous()?,
Some(g) => repeat_kv(key_states.clone(), g)?,
None => key_states.clone(),
};
let value_states = match num_key_value_groups {
Some(g) => repeat_kv(value_states.clone(), g)?.contiguous()?,
Some(g) => repeat_kv(value_states.clone(), g)?,
None => value_states.clone(),
};
let query_states = query_states.contiguous()?;
@@ -778,13 +778,14 @@ pub fn eager_attention_forward(
let attn_output = {
#[cfg(not(feature = "flash-attn"))]
{
let attn_weights = query_states.matmul(&key_states.transpose(D::Minus2, D::Minus1)?)?;
let attn_weights =
query_states.matmul(&key_states.transpose(D::Minus2, D::Minus1)?.contiguous()?)?;
let attn_weights = (attn_weights * scaling)?;
let attn_weights = match attention_mask {
None => attn_weights,
Some(mask) => attn_weights.broadcast_add(&mask.to_dtype(attn_weights.dtype())?)?,
};
let attn_weights = candle_nn::ops::softmax_last_dim(&attn_weights)?;
let attn_weights = candle_nn::ops::softmax_last_dim(&attn_weights)?.contiguous()?;
attn_weights.matmul(&value_states)?
}
#[cfg(feature = "flash-attn")]
@@ -987,163 +988,6 @@ pub fn deform_conv2d_kernel(
Ok(out)
}
pub struct LlamaModel {
pub embed_tokens: Embedding,
layers: Vec<NaiveAttnGateUpDownMLPBlock>,
norm: RmsNorm,
rotary_emb: RoPE,
}
impl LlamaModel {
pub fn new(
vb: VarBuilder,
vocab_size: usize,
hidden_size: usize,
num_hidden_layers: usize,
num_attention_heads: usize,
num_key_value_heads: Option<usize>,
head_dim: Option<usize>,
attn_bias: bool,
attn_pp_name: &str,
o_proj_pp_name: Option<&str>,
intermediate_size: usize,
hidden_act: Activation,
mlp_bias: bool,
mlp_pp_name: &str,
norm_eps: f64,
input_norm_pp_name: &str,
post_norm_pp_name: &str,
rope_theta_base: f32,
) -> Result<Self> {
let embed_tokens = embedding(vocab_size, hidden_size, vb.pp("embed_tokens"))?;
let mut layers = vec![];
let vb_layers = vb.pp("layers");
for i in 0..num_hidden_layers {
let layers_i = NaiveAttnGateUpDownMLPBlock::new(
vb_layers.pp(i),
hidden_size,
num_attention_heads,
num_key_value_heads,
head_dim,
attn_bias,
attn_pp_name,
o_proj_pp_name,
intermediate_size,
hidden_act,
mlp_bias,
mlp_pp_name,
norm_eps,
input_norm_pp_name,
post_norm_pp_name,
)?;
layers.push(layers_i);
}
let norm = rms_norm(hidden_size, norm_eps, vb.pp("norm"))?;
let head_dim = head_dim.unwrap_or(hidden_size / num_attention_heads);
let rotary_emb = RoPE::new(head_dim, rope_theta_base, vb.device())?;
Ok(Self {
embed_tokens,
layers,
norm,
rotary_emb,
})
}
pub fn forward(&mut self, inputs_embeds: &Tensor, seqlen_offset: usize) -> Result<Tensor> {
let (b_size, seq_len, _) = inputs_embeds.dims3()?;
let (cos, sin) = self
.rotary_emb
.forward(seqlen_offset, seq_len, inputs_embeds.device())?;
let mut xs = inputs_embeds.clone();
let attention_mask: Option<Tensor> = {
if seq_len <= 1 {
None
} else {
Some(prepare_causal_attention_mask(
b_size,
seq_len,
0,
xs.device(),
)?)
}
};
for layer in self.layers.iter_mut() {
xs = layer.forward(&xs, &cos, &sin, attention_mask.as_ref())?;
}
let xs = xs.apply(&self.norm)?;
Ok(xs)
}
pub fn clear_kv_cache(&mut self) {
for layer in self.layers.iter_mut() {
layer.clear_kv_cache()
}
}
}
pub struct LlamaForCausalLM {
pub model: LlamaModel,
lm_head: Linear,
}
impl LlamaForCausalLM {
pub fn new(
vb: VarBuilder,
vocab_size: usize,
hidden_size: usize,
num_hidden_layers: usize,
num_attention_heads: usize,
num_key_value_heads: Option<usize>,
head_dim: Option<usize>,
attn_bias: bool,
attn_pp_name: &str,
o_proj_pp_name: Option<&str>,
intermediate_size: usize,
hidden_act: Activation,
mlp_bias: bool,
mlp_pp_name: &str,
norm_eps: f64,
input_norm_pp_name: &str,
post_norm_pp_name: &str,
rope_theta_base: f32,
) -> Result<Self> {
let model = LlamaModel::new(
vb.pp("model"),
vocab_size,
hidden_size,
num_hidden_layers,
num_attention_heads,
num_key_value_heads,
head_dim,
attn_bias,
attn_pp_name,
o_proj_pp_name,
intermediate_size,
hidden_act,
mlp_bias,
mlp_pp_name,
norm_eps,
input_norm_pp_name,
post_norm_pp_name,
rope_theta_base,
)?;
let lm_head = linear_no_bias(hidden_size, vocab_size, vb.pp("lm_head"))?;
Ok(Self { model, lm_head })
}
pub fn forward(&mut self, inputs_embeds: &Tensor, seqlen_offset: usize) -> Result<Tensor> {
let outputs = self.model.forward(inputs_embeds, 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();
}
}
pub struct GLU {
dim: usize,
}
@@ -1190,10 +1034,14 @@ impl WNConv1d {
groups: usize,
stride: usize,
bias: bool,
weight_g_pp_name: Option<&str>,
weight_v_pp_name: Option<&str>,
) -> Result<Self> {
let in_c = in_c / groups;
let weight_g = vb.get((out_c, 1, 1), "weight_g")?;
let weight_v = vb.get((out_c, in_c, kernel_size), "weight_v")?;
let weight_g_pp_name = weight_g_pp_name.unwrap_or("weight_g");
let weight_v_pp_name = weight_v_pp_name.unwrap_or("weight_v");
let weight_g = vb.get((out_c, 1, 1), weight_g_pp_name)?;
let weight_v = vb.get((out_c, in_c, kernel_size), weight_v_pp_name)?;
// let bias = vb.get(out_c, "bias").ok();
let bias = if bias {
vb.get(out_c, "bias").ok()
+140
View File
@@ -0,0 +1,140 @@
use anyhow::{Result, anyhow};
use candle_core::{IndexOp, Tensor};
use candle_nn::ops::softmax;
use candle_transformers::generation::{LogitsProcessor, Sampling};
use rand::{SeedableRng, distr::Distribution};
pub fn get_logit_processor(
temperature: Option<f32>,
top_p: Option<f32>,
top_k: Option<usize>,
seed: u64,
) -> LogitsProcessor {
let temperature = temperature.and_then(|v| if v < 1e-7 { None } else { Some(v) });
match top_k {
None => LogitsProcessor::new(
seed,
temperature.map(|temp| temp as f64),
top_p.map(|tp| tp as f64),
),
Some(k) => {
let sampling = match temperature {
None => Sampling::ArgMax,
Some(temperature) => match top_p {
None => Sampling::TopK {
k,
temperature: temperature as f64,
},
Some(p) => Sampling::TopKThenTopP {
k,
p: p as f64,
temperature: temperature as f64,
},
},
};
LogitsProcessor::from_sampling(seed, sampling)
}
}
}
pub fn use_repeat_penalty(
repeat_penalty: f32,
repeat_last_n: Option<usize>,
logits: &Tensor,
context: &[u32],
) -> Result<Tensor> {
if repeat_penalty == 1.0 || repeat_last_n.map_or(false, |n| n == 0) {
Ok(logits.clone())
} else {
let start_at = if let Some(last_n) = repeat_last_n {
context.len().saturating_sub(last_n)
} else {
0
};
Ok(candle_transformers::utils::apply_repeat_penalty(
&logits,
repeat_penalty,
&context[start_at..],
)?)
}
}
/// logits shape: (dim)
pub fn simple_sample(
logits: &Tensor,
do_sample: bool,
temperature: Option<f64>,
top_k: Option<usize>,
top_p: Option<f32>,
previous_token_ids: Option<&[u32]>,
repeat_penalty: f32,
seed: Option<u64>,
) -> Result<u32> {
if logits.rank() != 1 {
return Err(anyhow!("simple_sample logits need rank = 1"));
}
let mut logits = if repeat_penalty != 1.0
&& let Some(tokens) = previous_token_ids
{
use_repeat_penalty(repeat_penalty, None, logits, tokens)?
} else {
logits.clone()
};
if !do_sample {
Ok(logits.argmax(0)?.to_scalar::<u32>()?)
} else {
if let Some(temp) = temperature
&& temp > 0.0
{
logits = logits.affine(1.0 / temp, 0.0)?;
}
if let Some(top_k) = top_k
&& top_k > 0
&& top_k > logits.dim(0)?
{
let sorted_indices = logits.arg_sort_last_dim(false)?;
let top_k_indices = sorted_indices.narrow(0, 0, top_k)?;
let top_k_logits = logits.gather(&top_k_indices, 0)?;
let threshold = top_k_logits.min_all()?;
let mask = logits.broadcast_lt(&threshold)?;
let on_true = Tensor::new(f32::NEG_INFINITY, logits.device())?
.to_dtype(logits.dtype())?
.broadcast_as(mask.shape())?;
logits = mask.where_cond(&on_true, &logits)?;
}
if let Some(top_p) = top_p
&& top_p > 0.0
&& top_p < 1.0
{
let sorted_indices = logits.arg_sort_last_dim(false)?;
let sorted_logits = logits.gather(&sorted_indices, 0)?;
let sorted_probs = softmax(&sorted_logits, 0)?;
let sorted_cumsum = sorted_probs.cumsum(0)?;
let mut mask = sorted_cumsum
.broadcast_gt(&Tensor::new(top_p, logits.device())?.to_dtype(logits.dtype())?)?;
// 保证数据不会被全部置为-inf
if mask.i(0)?.to_scalar::<u8>()? == 1 {
mask = mask.slice_scatter(&Tensor::new(0u32, logits.device())?, 0, 0)?;
}
let on_true = Tensor::new(f32::NEG_INFINITY, logits.device())?
.to_dtype(logits.dtype())?
.broadcast_as(mask.shape())?;
let new_logits = mask.where_cond(&on_true, &sorted_logits)?;
logits = logits.scatter(&sorted_indices, &new_logits, 0)?;
}
let probs = softmax(&logits, 0)?
.to_dtype(candle_core::DType::F32)?
.to_vec1::<f32>()?;
let distr = rand::distr::weighted::WeightedIndex::new(probs).map_err(|e| {
anyhow!(format!(
"simple_sampel new rand::distr::weighted::WeightedIndex Failed: {}",
e
))
})?;
let seed = seed.unwrap_or(34567);
let mut rng = rand::rngs::StdRng::seed_from_u64(seed);
let next_token = distr.sample(&mut rng) as u32;
Ok(next_token)
}
}
+2 -1
View File
@@ -133,7 +133,8 @@ impl AudioFeat {
}
pub fn extract_file(&self, audio_path: &str, device: &Device) -> Result<(Tensor, f32)> {
let wave_tensor = load_audio_with_resample(audio_path, device, Some(16000))?.squeeze(0)?;
let wave_tensor =
load_audio_with_resample(audio_path, device, Some(16000), Some(1))?.squeeze(0)?;
// fire_red_vad need i16 type data
let wave_tensor = wave_tensor.affine(32768.0, 0.0)?;
let dur = wave_tensor.dim(0)? as f32 / 16000.0;
+1 -1
View File
@@ -191,7 +191,7 @@ impl FireRedVad {
return Err(anyhow!("only stream model support detect_frame"));
}
let audio_frame =
resample_audio_from_bytes(audio_bytes, &self.device, Some(16000))?.squeeze(0)?;
resample_audio_from_bytes(audio_bytes, &self.device, Some(16000), 1)?.squeeze(0)?;
self.detect_frame(&audio_frame)
}
+1 -1
View File
@@ -91,7 +91,7 @@ impl FunAsrNanoProcessor {
let sub_token = tokenizer.text_encode_vec(sub_prompt, true)?;
source_ids.extend_from_slice(&sub_token);
fbank_mask.extend_from_slice(&vec![0u32; sub_token.len()]);
let audio_tensors = extract_audios(mes, &self.device, Some(self.fronted_conf.fs))?;
let audio_tensors = extract_audios(mes, &self.device, Some(self.fronted_conf.fs), Some(1))?;
if audio_tensors.is_empty() {
return Err(anyhow!("FunASRNano need audio input"));
}
+2 -2
View File
@@ -7,10 +7,10 @@ use crate::{
common::{
InferenceModel,
modules::{
LlamaForCausalLM, TwoLinearMLP, eager_attention_forward, get_conv1d, get_layer_norm,
TwoLinearMLP, eager_attention_forward, get_conv1d, get_layer_norm,
},
},
glm_asr_nano::config::{GlmAsrAudioConfig, GlmAsrNanoConfig},
glm_asr_nano::config::{GlmAsrAudioConfig, GlmAsrNanoConfig}, llama::LlamaForCausalLM,
},
position_embed::rope::{RoPE, glm_asr_apply_rotary_pos_emb},
utils::tensor_utils::{get_equal_mask, masked_scatter_dim0},
+1 -1
View File
@@ -205,7 +205,7 @@ impl GlmAsrNanoProcessor {
mes: &ChatCompletionParameters,
render_text: &str,
) -> Result<(Tensor, Vec<u32>, String)> {
let audio_tensors = extract_audios(mes, &self.device, Some(self.sampling_rate))?;
let audio_tensors = extract_audios(mes, &self.device, Some(self.sampling_rate), Some(1))?;
if audio_tensors.is_empty() {
return Err(anyhow::anyhow!("GlmASRNano need audio input"));
}
+81
View File
@@ -0,0 +1,81 @@
use serde::Deserialize;
#[derive(Clone, Debug, Deserialize)]
pub struct GPT2Config {
pub activation_function: String,
pub add_cross_attention: bool,
pub attn_pdrop: f64,
pub bad_words_ids: Option<Vec<u32>>,
pub begin_suppress_tokens: Option<Vec<u32>>,
pub bos_token_id: u32,
pub chunk_size_feed_forward: usize,
pub cross_attention_hidden_size: Option<usize>,
pub decoder_start_token_id: Option<u32>,
pub diversity_penalty: f64,
pub do_sample: bool,
pub dtype: Option<String>,
pub early_stopping: bool,
pub embd_pdrop: f64,
pub encoder_no_repeat_ngram_size: usize,
pub eos_token_id: u32,
pub exponential_decay_length_penalty: Option<f64>,
pub finetuning_task: Option<String>,
pub forced_bos_token_id: Option<u32>,
pub forced_eos_token_id: Option<u32>,
pub id2label: std::collections::HashMap<usize, String>,
pub initializer_range: f64,
pub is_decoder: bool,
pub is_encoder_decoder: bool,
pub label2id: std::collections::HashMap<String, usize>,
pub layer_norm_epsilon: f64,
pub length_penalty: f64,
pub max_length: usize,
pub min_length: usize,
pub model_type: String,
pub n_ctx: usize,
pub n_embd: usize,
pub n_head: usize,
pub n_inner: usize,
pub n_layer: usize,
pub n_positions: usize,
pub no_repeat_ngram_size: usize,
pub num_beam_groups: usize,
pub num_beams: usize,
pub num_return_sequences: usize,
pub output_attentions: bool,
pub output_hidden_states: bool,
pub output_scores: bool,
pub pad_token_id: u32,
pub position_embedding_type: String,
pub prefix: Option<String>,
pub problem_type: Option<String>,
pub remove_invalid_values: bool,
pub reorder_and_upcast_attn: bool,
pub repetition_penalty: f64,
pub resid_pdrop: f64,
pub return_dict: bool,
pub return_dict_in_generate: bool,
pub rope_base: f64,
pub scale_attn_by_inverse_layer_idx: bool,
pub scale_attn_weights: bool,
pub sep_token_id: Option<u32>,
pub summary_activation: Option<String>,
pub summary_first_dropout: f64,
pub summary_proj_to_labels: bool,
pub summary_type: String,
pub summary_use_proj: bool,
pub suppress_tokens: Option<Vec<u32>>,
pub task_specific_params: Option<serde_json::Value>,
pub temperature: f64,
pub tf_legacy_loss: bool,
pub tie_encoder_decoder: bool,
pub tie_word_embeddings: bool,
pub tokenizer_class: Option<String>,
pub top_k: usize,
pub top_p: f64,
pub torchscript: bool,
pub typical_p: f64,
pub use_bfloat16: bool,
pub use_cache: bool,
pub vocab_size: usize,
}
+311
View File
@@ -0,0 +1,311 @@
use anyhow::Result;
use candle_core::Tensor;
use candle_nn::{
Activation, Embedding, Init, LayerNorm, Linear, Module, VarBuilder, embedding, linear_b,
};
use crate::{
models::common::modules::{TwoLinearMLP, eager_attention_forward, get_layer_norm},
position_embed::rope::{RoPE, apply_rotary_pos_emb_interleave},
utils::tensor_utils::prepare_causal_attention_mask,
};
pub mod config;
pub struct GPT2Attention {
num_heads: usize,
head_dim: usize,
c_attn: Linear,
c_proj: Linear,
kv_cache: Option<(Tensor, Tensor)>,
}
impl GPT2Attention {
pub fn new(vb: VarBuilder, hidden_size: usize, num_heads: usize) -> Result<Self> {
let c_attn = linear_b(hidden_size, 3 * hidden_size, true, vb.pp("c_attn"))?;
let c_proj = linear_b(hidden_size, hidden_size, true, vb.pp("c_proj"))?;
let head_dim = hidden_size / num_heads;
Ok(Self {
num_heads,
head_dim,
c_attn,
c_proj,
kv_cache: None,
})
}
pub fn forward(
&mut self,
xs: &Tensor,
cos: Option<&Tensor>,
sin: Option<&Tensor>,
attention_mask: Option<&Tensor>,
) -> Result<Tensor> {
let (b, seq_len, _) = xs.dims3()?;
let xs = self.c_attn.forward(xs)?;
let xs_splits = xs.chunk(3, 2)?;
let query_states = xs_splits[0]
.as_ref()
.reshape((b, seq_len, self.num_heads, self.head_dim))?
.transpose(1, 2)?;
let key_states = xs_splits[1]
.as_ref()
.reshape((b, seq_len, self.num_heads, self.head_dim))?
.transpose(1, 2)?;
let value_states = xs_splits[2]
.as_ref()
.reshape((b, seq_len, self.num_heads, self.head_dim))?
.transpose(1, 2)?;
let (query_states, key_states) = if let Some(cos) = cos
&& let Some(sin) = sin
{
apply_rotary_pos_emb_interleave(&query_states, &key_states, cos, sin, false)?
} else {
(query_states, 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 scale = 1f64 / f64::sqrt(self.head_dim as f64);
let attn_output = eager_attention_forward(
&query_states,
&key_states,
&value_states,
None,
attention_mask,
scale,
)?;
let attn_output = attn_output.reshape((b, seq_len, self.num_heads * self.head_dim))?;
let attn_output = attn_output.apply(&self.c_proj)?;
Ok(attn_output)
}
pub fn clear_kv_cache(&mut self) {
self.kv_cache = None
}
}
pub struct GPT2MLP {
linear1: Linear,
linear2: Linear,
act: Activation,
}
impl GPT2MLP {
pub fn new(
vb: VarBuilder,
in_dim: usize,
middle_dim: usize,
out_dim: usize,
act: Activation,
) -> Result<Self> {
let c_fc_weight = vb
.get_with_hints((in_dim, middle_dim), "c_fc.weight", Init::Const(1.0))?
.t()?;
let c_fc_bias = vb.get_with_hints(middle_dim, "c_fc.bias", Init::Const(0.0))?;
let c_fc = Linear::new(c_fc_weight, Some(c_fc_bias));
let c_proj_weight = vb
.get_with_hints((middle_dim, out_dim), "c_proj.weight", Init::Const(1.0))?
.t()?;
let c_proj_bias = vb.get_with_hints(out_dim, "c_proj.bias", Init::Const(0.0))?;
let c_proj = Linear::new(c_proj_weight, Some(c_proj_bias));
Ok(Self {
linear1: c_fc,
linear2: c_proj,
act,
})
}
pub fn forward(&self, xs: &Tensor) -> Result<Tensor> {
let xs = xs
.apply(&self.linear1)?
.apply(&self.act)?
.apply(&self.linear2)?;
Ok(xs)
}
}
pub struct GPT2Block {
ln_1: LayerNorm,
attn: GPT2Attention,
ln_2: LayerNorm,
mlp: TwoLinearMLP,
}
impl GPT2Block {
pub fn new(
vb: VarBuilder,
hidden_size: usize,
num_heads: usize,
inner_dim: Option<usize>,
) -> Result<Self> {
let inner_dim = inner_dim.unwrap_or(4 * hidden_size);
let ln_1 = get_layer_norm(vb.pp("ln_1"), 1e-5, hidden_size, true)?;
let attn = GPT2Attention::new(vb.pp("attn"), hidden_size, num_heads)?;
let ln_2 = get_layer_norm(vb.pp("ln_2"), 1e-5, hidden_size, true)?;
let mlp = TwoLinearMLP::new(
vb.pp("mlp"),
hidden_size,
inner_dim,
hidden_size,
Activation::NewGelu,
true,
"fc_in",
"fc_out",
)?;
Ok(Self {
ln_1,
attn,
ln_2,
mlp,
})
}
pub fn forward(
&mut self,
xs: &Tensor,
cos: Option<&Tensor>,
sin: Option<&Tensor>,
attention_mask: Option<&Tensor>,
) -> Result<Tensor> {
let residual = xs.clone();
let xs = self.ln_1.forward(xs)?;
let xs = self.attn.forward(&xs, cos, sin, attention_mask)?;
let residual = xs.add(&residual)?;
let xs = self.ln_2.forward(&residual)?;
let xs = self.mlp.forward(&xs)?;
let xs = xs.add(&residual)?;
Ok(xs)
}
pub fn clear_kv_cache(&mut self) {
self.attn.clear_kv_cache()
}
}
#[allow(unused)]
pub struct GPT2Model {
pub wte: Option<Embedding>,
// wpe: Embedding, //rope not need
h: Vec<GPT2Block>,
ln_f: LayerNorm,
rope: Option<RoPE>,
}
#[allow(unused)]
impl GPT2Model {
pub fn new(
vb: VarBuilder,
hidden_size: usize,
num_heads: usize,
num_hidden_layers: usize,
vocab_size: usize,
// n_positions: usize,
) -> Result<Self> {
let wte = Some(embedding(vocab_size, hidden_size, vb.pp("wte"))?);
let vb_layers = vb.pp("h");
let mut h = vec![];
for i in 0..num_hidden_layers {
let block = GPT2Block::new(vb_layers.pp(i), hidden_size, num_heads, None)?;
h.push(block);
}
let ln_f = get_layer_norm(vb.pp("ln_f"), 1e-5, hidden_size, true)?;
let head_dim = hidden_size / num_heads;
let rope = RoPE::new(head_dim, 10000.0, vb.device())?;
Ok(Self {
wte,
h,
ln_f,
rope: Some(rope),
})
}
pub fn new_without_wte(
vb: VarBuilder,
hidden_size: usize,
num_heads: usize,
num_hidden_layers: usize,
vocab_size: usize,
// n_positions: usize,
) -> Result<Self> {
let wte = None;
let vb_layers = vb.pp("h");
let mut h = vec![];
for i in 0..num_hidden_layers {
let block = GPT2Block::new(vb_layers.pp(i), hidden_size, num_heads, None)?;
h.push(block);
}
let ln_f = get_layer_norm(vb.pp("ln_f"), 1e-5, hidden_size, true)?;
let head_dim = hidden_size / num_heads;
let rope = RoPE::new(head_dim, 10000.0, vb.device())?;
Ok(Self {
wte,
h,
ln_f,
rope: Some(rope),
})
}
pub fn new_with_wte(
vb: VarBuilder,
hidden_size: usize,
num_heads: usize,
num_hidden_layers: usize,
wte_embeddings: &Tensor,
) -> Result<Self> {
let wte = Some(Embedding::new(wte_embeddings.clone(), hidden_size));
let vb_layers = vb.pp("h");
let mut h = vec![];
for i in 0..num_hidden_layers {
let block = GPT2Block::new(vb_layers.pp(i), hidden_size, num_heads, None)?;
h.push(block);
}
let ln_f = get_layer_norm(vb.pp("ln_f"), 1e-5, hidden_size, true)?;
Ok(Self {
wte,
h,
ln_f,
rope: None,
})
}
pub fn forward(&mut self, inputs_embeds: &Tensor, seqlen_offset: usize) -> Result<Tensor> {
let (b_size, seq_len, _) = inputs_embeds.dims3()?;
let (cos, sin) = if let Some(rope) = &self.rope {
let (cos, sin) = rope.forward_repeat_interleave(seqlen_offset, seq_len, inputs_embeds.device())?;
(Some(cos), Some(sin))
} else {
(None, None)
};
let mut xs = inputs_embeds.clone();
let attention_mask: Option<Tensor> = {
if seq_len <= 1 {
None
} else {
Some(prepare_causal_attention_mask(
b_size,
seq_len,
0,
xs.device(),
)?)
}
};
for block in &mut self.h {
xs = block.forward(&xs, cos.as_ref(), sin.as_ref(), attention_mask.as_ref())?;
}
xs = self.ln_f.forward(&xs)?;
Ok(xs)
}
pub fn clear_kv_cache(&mut self) {
for layer in self.h.iter_mut() {
layer.clear_kv_cache()
}
}
}
+166
View File
@@ -0,0 +1,166 @@
use crate::{
models::common::modules::NaiveAttnGateUpDownMLPBlock, position_embed::rope::RoPE,
utils::tensor_utils::prepare_causal_attention_mask,
};
use anyhow::Result;
use candle_core::Tensor;
use candle_nn::{
Activation, Embedding, Linear, Module, RmsNorm, VarBuilder, embedding, linear_no_bias, rms_norm,
};
pub struct LlamaModel {
pub embed_tokens: Embedding,
layers: Vec<NaiveAttnGateUpDownMLPBlock>,
norm: RmsNorm,
rotary_emb: RoPE,
}
impl LlamaModel {
pub fn new(
vb: VarBuilder,
vocab_size: usize,
hidden_size: usize,
num_hidden_layers: usize,
num_attention_heads: usize,
num_key_value_heads: Option<usize>,
head_dim: Option<usize>,
attn_bias: bool,
attn_pp_name: &str,
o_proj_pp_name: Option<&str>,
intermediate_size: usize,
hidden_act: Activation,
mlp_bias: bool,
mlp_pp_name: &str,
norm_eps: f64,
input_norm_pp_name: &str,
post_norm_pp_name: &str,
rope_theta_base: f32,
) -> Result<Self> {
let embed_tokens = embedding(vocab_size, hidden_size, vb.pp("embed_tokens"))?;
let mut layers = vec![];
let vb_layers = vb.pp("layers");
for i in 0..num_hidden_layers {
let layers_i = NaiveAttnGateUpDownMLPBlock::new(
vb_layers.pp(i),
hidden_size,
num_attention_heads,
num_key_value_heads,
head_dim,
attn_bias,
attn_pp_name,
o_proj_pp_name,
intermediate_size,
hidden_act,
mlp_bias,
mlp_pp_name,
norm_eps,
input_norm_pp_name,
post_norm_pp_name,
)?;
layers.push(layers_i);
}
let norm = rms_norm(hidden_size, norm_eps, vb.pp("norm"))?;
let head_dim = head_dim.unwrap_or(hidden_size / num_attention_heads);
let rotary_emb = RoPE::new(head_dim, rope_theta_base, vb.device())?;
Ok(Self {
embed_tokens,
layers,
norm,
rotary_emb,
})
}
pub fn forward(&mut self, inputs_embeds: &Tensor, seqlen_offset: usize) -> Result<Tensor> {
let (b_size, seq_len, _) = inputs_embeds.dims3()?;
let (cos, sin) = self
.rotary_emb
.forward(seqlen_offset, seq_len, inputs_embeds.device())?;
let mut xs = inputs_embeds.clone();
let attention_mask: Option<Tensor> = {
if seq_len <= 1 {
None
} else {
Some(prepare_causal_attention_mask(
b_size,
seq_len,
0,
xs.device(),
)?)
}
};
for layer in self.layers.iter_mut() {
xs = layer.forward(&xs, &cos, &sin, attention_mask.as_ref())?;
}
let xs = xs.apply(&self.norm)?;
Ok(xs)
}
pub fn clear_kv_cache(&mut self) {
for layer in self.layers.iter_mut() {
layer.clear_kv_cache()
}
}
}
pub struct LlamaForCausalLM {
pub model: LlamaModel,
lm_head: Linear,
}
impl LlamaForCausalLM {
pub fn new(
vb: VarBuilder,
vocab_size: usize,
hidden_size: usize,
num_hidden_layers: usize,
num_attention_heads: usize,
num_key_value_heads: Option<usize>,
head_dim: Option<usize>,
attn_bias: bool,
attn_pp_name: &str,
o_proj_pp_name: Option<&str>,
intermediate_size: usize,
hidden_act: Activation,
mlp_bias: bool,
mlp_pp_name: &str,
norm_eps: f64,
input_norm_pp_name: &str,
post_norm_pp_name: &str,
rope_theta_base: f32,
) -> Result<Self> {
let model = LlamaModel::new(
vb.pp("model"),
vocab_size,
hidden_size,
num_hidden_layers,
num_attention_heads,
num_key_value_heads,
head_dim,
attn_bias,
attn_pp_name,
o_proj_pp_name,
intermediate_size,
hidden_act,
mlp_bias,
mlp_pp_name,
norm_eps,
input_norm_pp_name,
post_norm_pp_name,
rope_theta_base,
)?;
let lm_head = linear_no_bias(hidden_size, vocab_size, vb.pp("lm_head"))?;
Ok(Self { model, lm_head })
}
pub fn forward(&mut self, inputs_embeds: &Tensor, seqlen_offset: usize) -> Result<Tensor> {
let outputs = self.model.forward(inputs_embeds, 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();
}
}
+4
View File
@@ -128,6 +128,8 @@ impl FactorizedVectorQuantize {
1,
1,
true,
None,
None,
)?;
let out_project = WNConv1d::new(
vb.pp("out_project"),
@@ -139,6 +141,8 @@ impl FactorizedVectorQuantize {
1,
1,
true,
None,
None,
)?;
(Some(in_project), Some(out_project))
} else {
+3
View File
@@ -12,6 +12,7 @@ pub mod lfm2;
pub mod lfm2vl;
pub mod mask_gct;
pub mod minicpm4;
pub mod moss;
pub mod paddleocr_vl;
pub mod qwen2;
pub mod qwen2_5vl;
@@ -27,6 +28,8 @@ pub mod voxcpm_refact;
pub mod w2v_bert_2_0;
// pub mod sam3;
pub mod fire_red_vad;
pub mod gpt2;
pub mod llama;
use crate::{
models::{
+669
View File
@@ -0,0 +1,669 @@
use anyhow::{Result, anyhow};
use candle_core::{D, IndexOp, Tensor};
use candle_nn::{Embedding, LayerNorm, Linear, Module, VarBuilder, embedding, linear_no_bias};
use crate::{
models::{
common::modules::{
TwoLinearMLP, WNConv1d, eager_attention_forward, get_layer_norm, l2_normalize,
},
moss::config::{
MossAudioTokenizerConfig, MossAudioTokenizerModuleConfig,
MossAudioTokenizerQuantizerKwargs,
},
},
position_embed::rope::{RoPE, apply_rotary_pos_emb_roformer},
};
pub struct MossAudioTokenizerPatchedPretransform {
patch_size: usize,
is_downsample: bool,
}
impl MossAudioTokenizerPatchedPretransform {
pub fn new(patch_size: usize, is_downsample: bool) -> Self {
Self {
patch_size,
is_downsample,
}
}
pub fn encode(&self, x: &Tensor, input_lengths: &Tensor) -> Result<(Tensor, Tensor)> {
let (b, d, _) = x.dims3()?;
let x = x
.reshape((b, d, (), self.patch_size))?
.permute((0, 1, 3, 2))?
.reshape((b, d * self.patch_size, ()))?;
let out_lengths = input_lengths
.affine(1.0 / self.patch_size as f64, 0.0)?
.floor()?;
Ok((x, out_lengths))
}
pub fn decode(&self, x: &Tensor, input_lengths: &Tensor) -> Result<(Tensor, Tensor)> {
let (b, dh, l) = x.dims3()?;
let d = dh / self.patch_size;
let x = x
.reshape((b, d, self.patch_size, l))?
.permute((0, 1, 3, 2))?
.reshape((b, d, l * self.patch_size))?;
let out_lengths = (input_lengths * self.patch_size as f64)?;
Ok((x, out_lengths))
}
pub fn forward(&self, x: &Tensor, input_lengths: &Tensor) -> Result<(Tensor, Tensor)> {
if self.is_downsample {
self.encode(x, input_lengths)
} else {
self.decode(x, input_lengths)
}
}
}
pub struct MossAudioTokenizerMultiheadAttention {
num_heads: usize,
scale: f64,
in_proj: Linear,
out_proj: Linear,
}
impl MossAudioTokenizerMultiheadAttention {
pub fn new(vb: VarBuilder, embed_dim: usize, num_heads: usize) -> Result<Self> {
let head_dim = embed_dim / num_heads;
let scale = 1f64 / f64::sqrt(head_dim as f64);
let in_proj = linear_no_bias(embed_dim, 3 * embed_dim, vb.pp("in_proj"))?;
let out_proj = linear_no_bias(embed_dim, embed_dim, vb.pp("out_proj"))?;
Ok(Self {
num_heads,
scale,
in_proj,
out_proj,
})
}
pub fn forward(
&self,
xs: &Tensor,
cos: &Tensor,
sin: &Tensor,
mask: &Tensor,
input_lengths: &Tensor,
) -> Result<Tensor> {
let (bs, max_seqlen, _) = xs.dims3()?;
let projected = self
.in_proj
.forward(xs)?
.reshape((bs, max_seqlen, 3, self.num_heads, ()))?
.permute((2, 0, 3, 1, 4))?;
let [q, k, v] = projected
.chunk(3, 0)?
.try_into()
.map_err(|_| anyhow!("Chunk size mismatch"))?;
let q = q.squeeze(0)?.contiguous()?;
let k = k.squeeze(0)?.contiguous()?;
let v = v.squeeze(0)?.contiguous()?;
let (q, k) = apply_rotary_pos_emb_roformer(&q, &k, cos, sin)?;
// let (q, k) = self.apply_rope(&q, &k, cos, sin)?;
let attn = eager_attention_forward(&q, &k, &v, None, Some(mask), self.scale)?;
// (b, seq_len, n_head, dim) -> (b, n_head, seq_len, dim)
let attn = attn.transpose(1, 2)?;
let valid_q = Tensor::arange(0f32, max_seqlen as f32, xs.device())?
.reshape((1, 1, max_seqlen, 1))?
.broadcast_lt(
&input_lengths
.reshape((bs, 1, 1, 1))?
.repeat((1, 1, max_seqlen, 1))?,
)?
.broadcast_as(attn.shape())?;
let on_false = attn.zeros_like()?;
let attn = valid_q.where_cond(&attn, &on_false)?;
// (b, n_head, seq_len, dim) -> (b, seq_len, n_head, dim)
let attn = attn.transpose(1, 2)?;
let attn = attn.reshape((bs, max_seqlen, ()))?;
let out = self.out_proj.forward(&attn)?;
Ok(out)
}
}
pub struct MossAudioTokenizerTransformerLayer {
self_attn: MossAudioTokenizerMultiheadAttention,
norm1: LayerNorm,
norm2: LayerNorm,
ffn: TwoLinearMLP,
layer_scale_1: Tensor,
layer_scale_2: Tensor,
}
impl MossAudioTokenizerTransformerLayer {
pub fn new(vb: VarBuilder, config: &MossAudioTokenizerModuleConfig) -> Result<Self> {
let self_attn = MossAudioTokenizerMultiheadAttention::new(
vb.pp("self_attn"),
config.d_model.unwrap(),
config.num_heads.unwrap(),
)?;
let norm1 = get_layer_norm(vb.pp("norm1"), 1e-5, config.d_model.unwrap(), true)?;
let norm2 = get_layer_norm(vb.pp("norm2"), 1e-5, config.d_model.unwrap(), true)?;
let ffn = TwoLinearMLP::new(
vb.pp("ffn"),
config.d_model.unwrap(),
config.dim_feedforward.unwrap(),
config.d_model.unwrap(),
candle_nn::Activation::Gelu,
false,
"0",
"2",
)?;
let layer_scale_1 = vb
.get(config.d_model.unwrap(), "layer_scale_1.scale")?
.unsqueeze(0)?
.unsqueeze(0)?;
let layer_scale_2 = vb
.get(config.d_model.unwrap(), "layer_scale_2.scale")?
.unsqueeze(0)?
.unsqueeze(0)?;
Ok(Self {
self_attn,
norm1,
norm2,
ffn,
layer_scale_1,
layer_scale_2,
})
}
pub fn forward(
&self,
xs: &Tensor,
cos: &Tensor,
sin: &Tensor,
mask: &Tensor,
input_lengths: &Tensor,
) -> Result<Tensor> {
let residual = xs.clone();
let xs = self.norm1.forward(xs)?;
let xs = self.self_attn.forward(&xs, cos, sin, mask, input_lengths)?;
let xs = self.layer_scale_1.broadcast_mul(&xs)?;
let residual = residual.add(&xs)?;
let xs = self.norm2.forward(&residual)?;
let xs = self.ffn.forward(&xs)?;
let xs = self.layer_scale_2.broadcast_mul(&xs)?;
let xs = residual.add(&xs)?;
Ok(xs)
}
}
pub struct MossAudioTokenizerTransformer {
rope: RoPE, // use roformer
context: usize,
layers: Vec<MossAudioTokenizerTransformerLayer>,
}
impl MossAudioTokenizerTransformer {
pub fn new(
vb: VarBuilder,
config: &MossAudioTokenizerModuleConfig,
context: usize,
) -> Result<Self> {
let dim = config.d_model.unwrap() / config.num_heads.unwrap();
let rope = RoPE::new(dim, 10000.0, vb.device())?;
let vb_layers = vb.pp("layers");
let mut layers = vec![];
for i in 0..config.num_layers.unwrap() {
let layer = MossAudioTokenizerTransformerLayer::new(vb_layers.pp(i), config)?;
layers.push(layer);
}
Ok(Self {
rope,
context,
layers,
})
}
pub fn forward(&self, input_embeds: &Tensor, input_lengths: &Tensor) -> Result<Tensor> {
let t = input_embeds.dim(1)?;
let (cos, sin) = self.rope.forward(0, t, input_embeds.device())?;
let mut xs = input_embeds.clone();
let mask = self.build_attn_bias(input_lengths, t)?;
for layer in self.layers.iter() {
xs = layer.forward(&xs, &cos, &sin, &mask, input_lengths)?;
}
Ok(xs)
}
fn build_attn_bias(&self, input_lengths: &Tensor, max_seqlen: usize) -> Result<Tensor> {
let positions = Tensor::arange(0f32, max_seqlen as f32, input_lengths.device())?;
let input_lengths = input_lengths.reshape(((), 1, 1))?;
let valid_k = positions
.reshape((1, 1, max_seqlen))?
.broadcast_lt(&input_lengths)?;
let delta = positions
.reshape((1, max_seqlen, 1))?
.broadcast_sub(&positions.reshape((1, 1, max_seqlen))?)?;
let delta1 = delta.ge(&delta.zeros_like()?)?;
let delta2 = delta
.lt(&Tensor::new(self.context as f32, delta.device())?.broadcast_as(delta.shape())?)?;
let mask = delta1.minimum(&delta2)?;
let mask = mask.broadcast_minimum(&valid_k)?.unsqueeze(1)?;
let on_true = mask.zeros_like()?.to_dtype(candle_core::DType::F32)?;
let on_false = Tensor::new(f32::NEG_INFINITY, mask.device())?.broadcast_as(mask.shape())?;
let mask = mask.where_cond(&on_true, &on_false)?;
Ok(mask)
}
}
pub struct MossAudioTokenizerProjectedTransformer {
input_proj: Linear,
transformer: MossAudioTokenizerTransformer,
output_proj: Linear,
}
impl MossAudioTokenizerProjectedTransformer {
pub fn new(
vb: VarBuilder,
config: &MossAudioTokenizerModuleConfig,
context: usize,
) -> Result<Self> {
let input_proj = linear_no_bias(
config.input_dimension.unwrap(),
config.d_model.unwrap(),
vb.pp("input_proj"),
)?;
let transformer =
MossAudioTokenizerTransformer::new(vb.pp("transformer"), config, context)?;
let output_proj = linear_no_bias(
config.d_model.unwrap(),
config.output_dimension.unwrap(),
vb.pp("output_proj"),
)?;
Ok(Self {
input_proj,
transformer,
output_proj,
})
}
pub fn forward(
&self,
input_embeds: &Tensor,
input_lengths: &Tensor,
) -> Result<(Tensor, Tensor)> {
let xs = self.input_proj.forward(&input_embeds.transpose(1, 2)?)?;
let xs = self.transformer.forward(&xs, input_lengths)?;
let xs = self.output_proj.forward(&xs)?.transpose(1, 2)?;
Ok((xs, input_lengths.clone()))
}
}
pub enum MossAudioTokenizerModule {
PatchedPretransform(MossAudioTokenizerPatchedPretransform),
ProjectedTransformer(MossAudioTokenizerProjectedTransformer),
}
impl MossAudioTokenizerModule {
pub fn forward(
&self,
input_embeds: &Tensor,
input_lengths: &Tensor,
) -> Result<(Tensor, Tensor)> {
match self {
MossAudioTokenizerModule::PatchedPretransform(patch) => {
patch.forward(input_embeds, input_lengths)
}
MossAudioTokenizerModule::ProjectedTransformer(transformer) => {
transformer.forward(input_embeds, input_lengths)
}
}
}
}
pub struct MossAudioTokenizerLFQ {
in_proj: Option<WNConv1d>,
out_proj: Option<WNConv1d>,
codebook: Embedding,
codebook_l2_norm: Tensor,
}
impl MossAudioTokenizerLFQ {
pub fn new(vb: VarBuilder, config: &MossAudioTokenizerQuantizerKwargs) -> Result<Self> {
let in_proj = if config.rvq_dim != config.codebook_dim {
Some(WNConv1d::new(
vb.pp("in_proj"),
config.rvq_dim,
config.codebook_dim,
1,
1,
0,
1,
1,
true,
Some("parametrizations.weight.original0"),
Some("parametrizations.weight.original1"),
)?)
} else {
None
};
let out_proj = if config.rvq_dim != config.codebook_dim {
Some(WNConv1d::new(
vb.pp("out_proj"),
config.codebook_dim,
config.rvq_dim,
1,
1,
0,
1,
1,
true,
Some("parametrizations.weight.original0"),
Some("parametrizations.weight.original1"),
)?)
} else {
None
};
let codebook = embedding(config.codebook_size, config.codebook_dim, vb.pp("codebook"))?;
let codebook_l2_norm = l2_normalize(codebook.embeddings(), 1)?;
Ok(Self {
in_proj,
out_proj,
codebook,
codebook_l2_norm,
})
}
pub fn forward(&self, xs: &Tensor) -> Result<(Tensor, Tensor)> {
let z_e = if let Some(in_proj) = &self.in_proj {
in_proj.forward(xs)?
} else {
xs.clone()
};
let (bs, len, _) = z_e.dims3()?;
let encodings = z_e.transpose(1, 2)?.reshape(((), len))?;
let encodings = l2_normalize(&encodings, 1)?;
let dist1 = encodings.powf(2.0)?.sum_keepdim(1)?;
let dist2 = encodings
.affine(2.0, 0.0)?
.matmul(&self.codebook_l2_norm.t()?)?;
let dist3 = self.codebook_l2_norm.powf(2.0)?.sum_keepdim(1)?.t()?;
let dist = dist1.broadcast_sub(&dist2)?.broadcast_add(&dist3)?;
let indices = dist
.affine(-1.0, 0.0)?
.argmax(1)?
.reshape((bs, ()))?
.to_dtype(candle_core::DType::U32)?;
let z_q = self.codebook.forward(&indices)?.transpose(1, 2)?;
let mut z_q = z_e.add(&z_q.sub(&z_e)?)?;
if let Some(out_proj) = &self.out_proj {
z_q = out_proj.forward(&z_q)?;
}
Ok((z_q, indices))
}
}
pub struct MossAudioTokenizerResidualLFQ {
input_proj: Option<WNConv1d>,
output_proj: Option<WNConv1d>,
quantizers: Vec<MossAudioTokenizerLFQ>,
}
impl MossAudioTokenizerResidualLFQ {
pub fn new(vb: VarBuilder, config: &MossAudioTokenizerQuantizerKwargs) -> Result<Self> {
let input_proj = if config.input_dim != config.rvq_dim {
Some(WNConv1d::new(
vb.pp("input_proj"),
config.input_dim,
config.rvq_dim,
1,
1,
0,
1,
1,
true,
Some("parametrizations.weight.original0"),
Some("parametrizations.weight.original1"),
)?)
} else {
None
};
let output_proj = if config.rvq_dim != config.output_dim {
Some(WNConv1d::new(
vb.pp("output_proj"),
config.rvq_dim,
config.output_dim,
1,
1,
0,
1,
1,
true,
Some("parametrizations.weight.original0"),
Some("parametrizations.weight.original1"),
)?)
} else {
None
};
let vb_quantizers = vb.pp("quantizers");
let mut quantizers = vec![];
for i in 0..config.num_quantizers {
let layer = MossAudioTokenizerLFQ::new(vb_quantizers.pp(i), config)?;
quantizers.push(layer);
}
Ok(Self {
input_proj,
output_proj,
quantizers,
})
}
pub fn forward(&self, input_values: &Tensor, length: &Tensor) -> Result<Tensor> {
let z = if let Some(proj) = &self.input_proj {
proj.forward(input_values)?
} else {
input_values.clone()
};
let max_time = z.dim(2)?;
let mask = Tensor::arange(0f32, max_time as f32, z.device())?
.unsqueeze(0)?
.broadcast_lt(&length.unsqueeze(1)?)?
.unsqueeze(1)?;
// let mut quantized_out = z.zeros_like()?;
let mut residual = z.clone();
let on_false = residual.zeros_like()?;
let mask_reshape = mask.broadcast_as(residual.shape())?;
let mut all_indices = vec![];
for quantizer in self.quantizers.iter() {
let masked_residual = mask_reshape.where_cond(&residual, &on_false)?;
let (z_q_i, indices_i) = quantizer.forward(&masked_residual)?;
all_indices.push(indices_i);
let z_q_i_mask = mask_reshape.where_cond(&z_q_i, &on_false)?;
residual = residual.sub(&z_q_i_mask)?;
}
let all_indices = Tensor::stack(&all_indices, 0)?;
Ok(all_indices)
}
}
pub struct MossAudioTokenizer {
pub sampling_rate: usize,
pub downsample_rate: usize,
pub number_channels: usize,
pub enable_channel_interleave: bool,
encoder: Vec<MossAudioTokenizerModule>,
quantizer: MossAudioTokenizerResidualLFQ,
decoder: Vec<MossAudioTokenizerModule>,
}
impl MossAudioTokenizer {
pub fn new(vb: VarBuilder, config: &MossAudioTokenizerConfig) -> Result<Self> {
let channel_interleave_factor =
if config.enable_channel_interleave && config.number_channels > 1 {
config.number_channels
} else {
1
};
let current_frame_rate = config.sampling_rate * channel_interleave_factor;
let vb_encoder = vb.pp("encoder");
let mut encoder = vec![];
for (layer_id, cfg) in config.encoder_kwargs.iter().enumerate() {
if cfg.module_type == "PatchedPretransform"
&& let Some(patch_size) = cfg.patch_size
{
let layer = MossAudioTokenizerPatchedPretransform::new(patch_size, true);
encoder.push(MossAudioTokenizerModule::PatchedPretransform(layer));
} else if cfg.module_type == "Transformer" {
let context_duration = cfg
.context_duration
.unwrap_or(config.causal_transformer_context_duration);
let context = (current_frame_rate as f64 * context_duration).round() as usize;
let layer = MossAudioTokenizerProjectedTransformer::new(
vb_encoder.pp(layer_id),
cfg,
context,
)?;
encoder.push(MossAudioTokenizerModule::ProjectedTransformer(layer));
} else {
return Err(anyhow!(
"Moss Module only sopport PatchedPretransform and Transformer, but get: {}",
cfg.module_type
));
}
}
let quantizer =
MossAudioTokenizerResidualLFQ::new(vb.pp("quantizer"), &config.quantizer_kwargs)?;
let vb_decoder = vb.pp("decoder");
let mut decoder = vec![];
for (layer_id, cfg) in config.decoder_kwargs.iter().enumerate() {
if cfg.module_type == "PatchedPretransform"
&& let Some(patch_size) = cfg.patch_size
{
let layer = MossAudioTokenizerPatchedPretransform::new(patch_size, true);
decoder.push(MossAudioTokenizerModule::PatchedPretransform(layer));
} else if cfg.module_type == "Transformer" {
let context_duration = cfg
.context_duration
.unwrap_or(config.causal_transformer_context_duration);
let context = (current_frame_rate as f64 * context_duration).round() as usize;
let layer = MossAudioTokenizerProjectedTransformer::new(
vb_decoder.pp(layer_id),
cfg,
context,
)?;
decoder.push(MossAudioTokenizerModule::ProjectedTransformer(layer));
} else {
return Err(anyhow!(
"Moss Module only sopport PatchedPretransform and Transformer, but get: {}",
cfg.module_type
));
}
}
Ok(Self {
sampling_rate: config.sampling_rate,
downsample_rate: config.downsample_rate,
number_channels: config.number_channels,
enable_channel_interleave: config.enable_channel_interleave,
encoder,
quantizer,
decoder,
})
}
fn flatten_channels_for_codec(
&self,
input_values: &Tensor,
length: &Tensor,
) -> Result<(Tensor, Tensor)> {
let (bs, _, audio_len) = input_values.dims3()?;
let input_values = if audio_len % self.downsample_rate != 0 {
let pad_length = self.downsample_rate - (audio_len % self.downsample_rate);
input_values.pad_with_zeros(D::Minus1, 0, pad_length)?
} else {
input_values.clone()
};
if self.number_channels > 1 && self.enable_channel_interleave {
let input_values = input_values
.transpose(1, 2)?
.contiguous()?
.reshape((bs, 1, ()))?;
let length = (length * self.number_channels as f64)?;
Ok((input_values, length))
} else {
Ok((input_values, length.clone()))
}
}
pub fn batch_encode(&self, input_values: &Tensor, length: &Tensor) -> Result<Vec<Tensor>> {
let (mut encoder_hidden_states, mut encoder_hidden_lengths) =
self.flatten_channels_for_codec(input_values, length)?;
for layer in &self.encoder {
(encoder_hidden_states, encoder_hidden_lengths) =
layer.forward(&encoder_hidden_states, &encoder_hidden_lengths)?;
}
let audio_codes = self
.quantizer
.forward(&encoder_hidden_states, &encoder_hidden_lengths)?;
// (dim, bs, len) -> (bs, len, dim)
let audio_codes = audio_codes.permute((1, 2, 0))?;
let mut audio_codes_vec = vec![];
for index in 0..encoder_hidden_lengths.dim(0)? {
let codes_i = audio_codes.i(index)?;
let length = encoder_hidden_lengths.i(index)?.to_scalar::<f32>()? as usize;
let codes_i = codes_i.narrow(0, 0, length)?;
audio_codes_vec.push(codes_i);
}
Ok(audio_codes_vec)
}
pub fn encode_one(&self, wav: &Tensor) -> Result<Tensor> {
// (channel, audio_len) -> (bs=1, channel, audio_len)
let (c, len) = wav.dims2()?;
if c != self.number_channels {
return Err(anyhow!(
"MossAudioTokenizer encode_one need number_channels: {} but the wav channel: {}",
self.number_channels,
c,
));
}
let input_values = wav.unsqueeze(0)?;
let length = Tensor::new(vec![len as f32], wav.device())?;
let audio_vec = self.batch_encode(&input_values, &length)?;
Ok(audio_vec[0].clone())
}
pub fn encode_list(&self, wavs: &Vec<Tensor>) -> Result<Vec<Tensor>> {
if wavs.is_empty() {
return Err(anyhow!(
"MossAudioTokenizer encode_list need wavs len > 0, but the wavs is empty"
));
}
let mut length = vec![];
for wav in wavs.iter() {
let (c, len) = wav.dims2()?;
if c != self.number_channels {
return Err(anyhow!(
"MossAudioTokenizer encode_list need number_channels: {} but the wav channel: {}",
self.number_channels,
c,
));
}
length.push(len as u32);
}
let max_length = *length.iter().max().unwrap_or(&0) as usize;
let mut input_values = vec![];
for wav in wavs.iter() {
let audio_len = wav.dim(1)?;
let wav_ = if audio_len < max_length {
wav.pad_with_zeros(D::Minus1, 0, max_length - audio_len)?
} else {
wav.clone()
};
input_values.push(wav_);
}
let input_values = Tensor::stack(&input_values, 0)?;
let length_tensor = Tensor::new(length.clone(), input_values.device())?
.to_dtype(candle_core::DType::F32)?;
Ok(self.batch_encode(&input_values, &length_tensor)?)
}
}
+159
View File
@@ -0,0 +1,159 @@
use serde::Deserialize;
use crate::models::gpt2::config::GPT2Config;
#[derive(Debug, Deserialize)]
pub struct MossAudioTokenizerConfig {
pub sample_rate: usize,
pub sampling_rate: usize,
pub downsample_rate: usize,
pub causal_transformer_context_duration: f64,
pub number_channels: usize,
pub enable_channel_interleave: bool,
pub compute_dtype: String,
pub dtype: String,
pub code_dim: usize,
pub encoder_kwargs: Vec<MossAudioTokenizerModuleConfig>,
pub decoder_kwargs: Vec<MossAudioTokenizerModuleConfig>,
pub quantizer_type: String,
pub quantizer_kwargs: MossAudioTokenizerQuantizerKwargs,
pub reversed_decoder_kwargs: Vec<MossAudioTokenizerModuleConfig>,
}
#[derive(Debug, Deserialize)]
pub struct MossAudioTokenizerModuleConfig {
pub module_type: String,
pub patch_size: Option<usize>,
pub causal: Option<bool>,
pub context_duration: Option<f64>,
pub conv_layout: Option<bool>,
pub d_model: Option<usize>,
pub dim_feedforward: Option<usize>,
pub gating: Option<String>,
pub input_dimension: Option<usize>,
pub layer_scale: Option<f64>,
pub max_period: Option<usize>,
pub norm: Option<String>,
pub num_heads: Option<usize>,
pub num_layers: Option<usize>,
pub output_dimension: Option<usize>,
pub positional_embedding: Option<String>,
}
#[derive(Debug, Deserialize)]
pub struct MossAudioTokenizerQuantizerKwargs {
pub codebook_dim: usize,
pub codebook_loss_weight: f64,
pub codebook_size: usize,
pub commitment_loss_weight: f64,
pub input_dim: usize,
pub num_quantizers: usize,
pub output_dim: usize,
pub quantizer_dropout: f64,
pub quantizer_type: String,
pub rvq_dim: usize,
}
#[derive(Debug, Deserialize)]
pub struct MossTTSConfig {
pub add_cross_attention: bool,
// Audio Tokenizer Specifics
pub audio_assistant_slot_token_id: u32,
pub audio_codebook_sizes: Vec<usize>,
pub audio_end_token_id: u32,
pub audio_pad_token_id: u32,
pub audio_start_token_id: u32,
pub audio_tokenizer_sample_rate: usize,
pub audio_user_slot_token_id: u32,
pub audio_vocab_size: usize,
// Generation/Model Params (Simplified nullables to Options or defaults if not critical)
pub bad_words_ids: Option<Vec<u32>>,
pub begin_suppress_tokens: Option<Vec<u32>>,
pub bos_token_id: Option<u32>,
pub chunk_size_feed_forward: usize,
pub cross_attention_hidden_size: Option<usize>,
pub decoder_start_token_id: Option<u32>,
pub diversity_penalty: f64,
pub do_sample: bool,
pub dtype: String,
pub early_stopping: bool,
pub encoder_no_repeat_ngram_size: usize,
pub eos_token_id: Option<u32>,
pub exponential_decay_length_penalty: Option<f64>,
pub finetuning_task: Option<String>,
pub forced_bos_token_id: Option<u32>,
pub forced_eos_token_id: Option<u32>,
// GPT2 Backbone Config
pub gpt2_config: GPT2Config,
pub hidden_size: usize,
pub id2label: std::collections::HashMap<usize, String>,
pub im_end_token_id: u32,
pub im_start_token_id: u32,
pub initializer_range: f64,
pub is_decoder: bool,
pub is_encoder_decoder: bool,
pub label2id: std::collections::HashMap<String, usize>,
pub length_penalty: f64,
pub local_transformer_attn_implementation: String,
pub local_transformer_layers: usize,
pub max_length: usize,
pub max_position_embeddings: usize,
pub min_length: usize,
pub model_architecture: String,
pub model_type: String,
pub n_vq: usize,
pub no_repeat_ngram_size: usize,
pub num_beam_groups: usize,
pub num_beams: usize,
pub num_return_sequences: usize,
pub output_attentions: bool,
pub output_hidden_states: bool,
pub output_scores: bool,
pub pad_token_id: u32,
pub prefix: Option<String>,
pub problem_type: Option<String>,
// pub pruned_heads: std::collections::HashMap<String, Vec<usize>>,
pub remove_invalid_values: bool,
pub repetition_penalty: f64,
pub return_dict: bool,
pub return_dict_in_generate: bool,
pub sep_token_id: Option<u32>,
pub suppress_tokens: Option<Vec<u32>>,
pub task_specific_params: Option<serde_json::Value>,
pub temperature: f32,
pub tf_legacy_loss: bool,
pub tie_encoder_decoder: bool,
pub tie_word_embeddings: bool,
pub tokenizer_class: String,
pub tokenizer_use_fast: bool,
pub top_k: usize,
pub top_p: f32,
pub torchscript: bool,
pub typical_p: f64,
pub use_bfloat16: bool,
pub vocab_size: usize,
}
+95
View File
@@ -0,0 +1,95 @@
use std::collections::HashMap;
use crate::{
models::moss::{
audio_tokenizer_nano::MossAudioTokenizer,
config::{MossAudioTokenizerConfig, MossTTSConfig},
processor::MossTTSProcessor,
tts_nano::{MossTTSMode, MossTTSModel},
},
utils::{find_type_files, get_device, get_dtype},
};
use anyhow::{Result, anyhow};
use candle_core::{DType, Device, pickle::read_all_with_key};
use candle_nn::VarBuilder;
use sentencepiece::SentencePieceProcessor;
pub struct MossTTSGenerate {
pub audio_tokenizer: MossAudioTokenizer,
pub text_tokenizer: SentencePieceProcessor,
pub processor: MossTTSProcessor,
pub model: MossTTSModel,
pub device: Device,
}
impl MossTTSGenerate {
pub fn init(
tts_path: &str,
audio_tokenizer_path: &str,
device: Option<&Device>,
dtype: Option<DType>,
) -> Result<Self> {
let audio_tokenizer_config_path = audio_tokenizer_path.to_string() + "/config.json";
let audio_tokenizer_cfg: MossAudioTokenizerConfig =
serde_json::from_slice(&std::fs::read(audio_tokenizer_config_path)?)?;
let model_list = find_type_files(audio_tokenizer_path, "safetensors")?;
let audio_dtype = get_dtype(dtype.clone(), &audio_tokenizer_cfg.dtype);
let device = get_device(device);
let vb = unsafe { VarBuilder::from_mmaped_safetensors(&model_list, audio_dtype, &device)? };
let audio_tokenizer = MossAudioTokenizer::new(vb, &audio_tokenizer_cfg)?;
let text_tokenizer_path = tts_path.to_string() + "/tokenizer.model";
let text_tokenizer = SentencePieceProcessor::open(text_tokenizer_path)
.map_err(|e| anyhow!(format!("load bpe.model file error:{}", e)))?;
let tts_cfg_path = tts_path.to_string() + "/config.json";
let tts_cfg: MossTTSConfig = serde_json::from_slice(&std::fs::read(tts_cfg_path)?)?;
let processor = MossTTSProcessor::new(
&tts_cfg,
audio_tokenizer_cfg.sample_rate,
audio_tokenizer_cfg.number_channels,
&text_tokenizer,
)?;
let model_list = find_type_files(tts_path, "bin")?;
let mut dict_to_hashmap = HashMap::new();
// let cfg_dtype = tts_cfg.dtype.as_str();
let m_dtype = get_dtype(dtype, "bfloat16");
for m in model_list {
let dict = read_all_with_key(m, None)?;
for (k, v) in dict {
// println!("key: {}, tensor shape: {:?}", k, v);
dict_to_hashmap.insert(k, v);
}
}
let vb = VarBuilder::from_tensors(dict_to_hashmap, m_dtype, &device);
let model = MossTTSModel::new(vb, &tts_cfg)?;
Ok(Self {
audio_tokenizer,
text_tokenizer,
processor,
model,
device,
})
}
pub fn generate(
&mut self,
text: &str,
prompt_audio_path: Option<&str>,
prompt_text: Option<&str>,
mode: Option<MossTTSMode>,
) -> Result<()> {
let (mut input_ids, mask) = self.processor.build_inference_input_ids(
text,
prompt_audio_path,
prompt_text,
mode,
&self.audio_tokenizer,
&self.text_tokenizer,
&self.device,
)?;
let _ = self.model.generate(&input_ids, Some(&mask))?;
// println!("input_ids: {}", input_ids);
// println!("mask: {}", mask);
Ok(())
}
}
+5
View File
@@ -0,0 +1,5 @@
pub mod audio_tokenizer_nano;
pub mod config;
pub mod generate;
pub mod processor;
pub mod tts_nano;
+218
View File
@@ -0,0 +1,218 @@
use crate::{
models::moss::{
audio_tokenizer_nano::MossAudioTokenizer, config::MossTTSConfig, tts_nano::MossTTSMode,
},
tokenizer::sentencepiece_encode_vec,
utils::{audio_utils::load_audio_with_resample, prepare_tts_text},
};
use anyhow::{Result, anyhow};
use candle_core::{Device, Tensor};
use sentencepiece::SentencePieceProcessor;
pub struct MossTTSProcessor {
target_sample_rate: usize,
target_channels: usize,
audio_start_token_id: u32,
audio_end_token_id: u32,
audio_user_slot_token_id: u32,
audio_assistant_slot_token_id: u32,
audio_pad_token_id: u32,
n_vq: usize,
prompt_token_ids: Vec<u32>,
user_after_ids: Vec<u32>,
assistant_ids: Vec<u32>,
none_ids: Vec<u32>,
}
impl MossTTSProcessor {
pub fn new(
tts_cfg: &MossTTSConfig,
target_sample_rate: usize,
target_channels: usize,
text_tokenizer: &SentencePieceProcessor,
) -> Result<Self> {
let mut prompt_token_ids = vec![tts_cfg.im_start_token_id];
let user_role_ids = sentencepiece_encode_vec("user\n", text_tokenizer)?;
prompt_token_ids.extend_from_slice(&user_role_ids);
let user_template_pre_ids =
sentencepiece_encode_vec("<user_inst>\n- Reference(s):\n", text_tokenizer)?;
prompt_token_ids.extend_from_slice(&user_template_pre_ids);
let user_after_ids = sentencepiece_encode_vec(
"\n- Instruction:\nNone\n- Tokens:\nNone\n- Quality:\nNone\n- Sound Event:\nNone\n- Ambient Sound:\nNone\n- Language:\nNone\n- Text:\n",
text_tokenizer,
)?;
let mut assistant_ids = vec![];
let user_suffix = sentencepiece_encode_vec("\n</user_inst>", text_tokenizer)?;
assistant_ids.extend_from_slice(&user_suffix);
assistant_ids.push(tts_cfg.im_end_token_id);
let assistant_turn_ids = sentencepiece_encode_vec("\n", text_tokenizer)?;
assistant_ids.extend_from_slice(&assistant_turn_ids);
assistant_ids.push(tts_cfg.im_start_token_id);
let assistant_role_ids = sentencepiece_encode_vec("assistant\n", text_tokenizer)?;
assistant_ids.extend_from_slice(&assistant_role_ids);
let none_ids = sentencepiece_encode_vec("None", text_tokenizer)?;
Ok(Self {
target_sample_rate,
target_channels,
audio_start_token_id: tts_cfg.audio_start_token_id,
audio_end_token_id: tts_cfg.audio_end_token_id,
audio_user_slot_token_id: tts_cfg.audio_user_slot_token_id,
audio_assistant_slot_token_id: tts_cfg.audio_assistant_slot_token_id,
audio_pad_token_id: tts_cfg.audio_pad_token_id,
n_vq: tts_cfg.n_vq,
prompt_token_ids,
user_after_ids,
assistant_ids,
none_ids,
})
}
fn resolved_mode(
&self,
mode: Option<MossTTSMode>,
has_prompt_text: bool,
has_prompt_audio: bool,
) -> Result<MossTTSMode> {
let normalized_mode = mode.unwrap_or(MossTTSMode::VoiceClone);
if normalized_mode == MossTTSMode::VoiceClone {
if !has_prompt_audio {
return Err(anyhow!("voice_clone mode requires prompt_audio_path"));
}
if has_prompt_text {
println!("voice_clone mode does not accept prompt_text");
}
} else {
if has_prompt_text != has_prompt_audio {
return Err(anyhow!(
"continuation mode accepts either target text only, or prompt_text and prompt_audio_path together."
));
}
}
Ok(normalized_mode)
}
pub fn build_inference_input_ids(
&self,
text: &str,
prompt_audio_path: Option<&str>,
prompt_text: Option<&str>,
mode: Option<MossTTSMode>,
audio_tokenizer: &MossAudioTokenizer,
text_tokenizer: &SentencePieceProcessor,
device: &Device,
) -> Result<(Tensor, Tensor)> {
let mode = self.resolved_mode(mode, prompt_text.is_some(), prompt_audio_path.is_some())?;
let audio_code = if let Some(audio_path) = prompt_audio_path {
let audio = load_audio_with_resample(
audio_path,
device,
Some(self.target_sample_rate),
Some(self.target_channels),
)?;
Some(audio_tokenizer.encode_one(&audio)?)
} else {
None
};
let text = &prepare_tts_text(text)?;
let prompt_text = if let Some(prompt_text) = prompt_text {
Some(prepare_tts_text(prompt_text)?)
} else {
None
};
// TODO: 长文本段切分
if mode == MossTTSMode::VoiceClone
&& let Some(prompt_audio_codes) = &audio_code
{
let mut prompt_token_ids = vec![];
prompt_token_ids.extend_from_slice(&self.prompt_token_ids);
prompt_token_ids.push(self.audio_start_token_id);
let prompt_ids_tensor = Self::build_text_raw(
&prompt_token_ids,
self.audio_pad_token_id,
self.n_vq,
device,
)?;
let text_token_ids = sentencepiece_encode_vec(text, text_tokenizer)?;
let mut suffix_token_ids = vec![self.audio_end_token_id];
suffix_token_ids.extend_from_slice(&self.user_after_ids);
suffix_token_ids.extend_from_slice(&text_token_ids);
suffix_token_ids.extend_from_slice(&self.assistant_ids);
suffix_token_ids.push(self.audio_start_token_id);
let audio_prefix_rows = Self::build_audio_prefix_rows(
&prompt_audio_codes,
self.audio_user_slot_token_id,
device,
)?;
let suffix_rows = Self::build_text_raw(
&suffix_token_ids,
self.audio_pad_token_id,
self.n_vq,
device,
)?;
let input_ids =
Tensor::cat(&[&prompt_ids_tensor, &audio_prefix_rows, &suffix_rows], 0)?
.unsqueeze(0)?;
let (bs, len, _) = input_ids.dims3()?;
let mask = Tensor::ones((bs, len), candle_core::DType::U32, device)?;
Ok((input_ids, mask))
} else {
let text = if let Some(prompt_text) = prompt_text {
prompt_text + text
} else {
text.to_string()
};
let text_token_ids = sentencepiece_encode_vec(&text, text_tokenizer)?;
let mut prompt_ids = vec![];
prompt_ids.extend_from_slice(&self.prompt_token_ids);
prompt_ids.extend_from_slice(&self.none_ids);
prompt_ids.extend_from_slice(&self.user_after_ids);
prompt_ids.extend_from_slice(&text_token_ids);
prompt_ids.extend_from_slice(&self.assistant_ids);
prompt_ids.push(self.audio_start_token_id);
let mut input_ids =
Self::build_text_raw(&prompt_ids, self.audio_pad_token_id, self.n_vq, device)?;
if let Some(prompt_audio_codes) = &audio_code {
let audio_prefix_rows = Self::build_audio_prefix_rows(
&prompt_audio_codes,
self.audio_assistant_slot_token_id,
device,
)?;
input_ids = Tensor::cat(&[&input_ids, &audio_prefix_rows], 0)?;
}
input_ids = input_ids.unsqueeze(0)?;
let (bs, len, _) = input_ids.dims3()?;
let mask = Tensor::ones((bs, len), candle_core::DType::U32, device)?;
Ok((input_ids, mask))
}
}
fn build_audio_prefix_rows(
prompt_audio_codes: &Tensor,
slot_token_id: u32,
device: &Device,
) -> Result<Tensor> {
let audio_len = prompt_audio_codes.dim(0)?;
let pad_tensor = Tensor::new(slot_token_id, device)?.broadcast_as((audio_len, 1))?;
let rows = Tensor::cat(&[&pad_tensor, prompt_audio_codes], 1)?;
Ok(rows)
}
fn build_text_raw(
token_ids: &Vec<u32>,
audio_pad_token_id: u32,
n_vq: usize,
device: &Device,
) -> Result<Tensor> {
let id_len = token_ids.len();
//(1, len) -> (len, 1)
let text_tensor = Tensor::from_slice(token_ids, (id_len, 1), device)?;
// (len, n_vq)
let pad_tensor = Tensor::new(audio_pad_token_id, device)?.broadcast_as((id_len, n_vq))?;
let rows = Tensor::cat(&[text_tensor, pad_tensor], 1)?;
Ok(rows)
}
}
+256
View File
@@ -0,0 +1,256 @@
use crate::models::{common::sample::simple_sample, gpt2::GPT2Model, moss::config::MossTTSConfig};
use anyhow::{Result, anyhow};
use candle_core::{D, IndexOp, Tensor};
use candle_nn::{Embedding, Linear, Module, VarBuilder, embedding, linear_no_bias};
#[derive(PartialEq, Debug)]
pub enum MossTTSMode {
Continuation,
VoiceClone,
}
pub struct MossTTSModel {
transformer: GPT2Model,
audio_embeddings: Vec<Embedding>,
text_lm_head: Linear,
audio_lm_heads: Vec<Linear>,
local_transformer: GPT2Model,
audio_assistant_slot_token_id: usize,
audio_end_token_id: usize,
n_vq: usize,
audio_pad_token_id_tensor: Tensor,
audio_codebook_sizes: Vec<usize>,
audio_temperature: f64,
audio_top_k: usize,
audio_top_p: f32,
audio_repetition_penalty: f32,
}
impl MossTTSModel {
pub fn new(vb: VarBuilder, cfg: &MossTTSConfig) -> Result<Self> {
let transformer = GPT2Model::new(
vb.pp("transformer"),
cfg.gpt2_config.n_embd,
cfg.gpt2_config.n_head,
cfg.gpt2_config.n_layer,
cfg.gpt2_config.vocab_size,
// cfg.gpt2_config.n_positions,
)?;
let mut audio_embeddings = vec![];
let audio_embed_vb = vb.pp("audio_embeddings");
for i in 0..cfg.n_vq {
let embed = embedding(
cfg.audio_codebook_sizes[i],
cfg.gpt2_config.n_embd,
audio_embed_vb.pp(i),
)?;
audio_embeddings.push(embed);
}
let text_lm_head = linear_no_bias(
cfg.gpt2_config.n_embd,
cfg.gpt2_config.vocab_size,
vb.pp("text_lm_head"),
)?;
let mut audio_lm_heads = vec![];
let audio_lm_vb = vb.pp("audio_lm_heads");
for i in 0..cfg.n_vq {
let layer = linear_no_bias(
cfg.gpt2_config.n_embd,
cfg.audio_codebook_sizes[i],
audio_lm_vb.pp(i),
)?;
audio_lm_heads.push(layer);
}
let mut local_gpt2_cfg = cfg.gpt2_config.clone();
local_gpt2_cfg.n_layer = cfg.local_transformer_layers;
local_gpt2_cfg.n_positions = cfg.n_vq + 1;
local_gpt2_cfg.n_ctx = cfg.n_vq + 1;
let local_transformer = GPT2Model::new_without_wte(
vb.pp("local_transformer"),
local_gpt2_cfg.n_embd,
local_gpt2_cfg.n_head,
local_gpt2_cfg.n_layer,
local_gpt2_cfg.vocab_size,
// local_gpt2_cfg.n_positions,
)?;
let audio_pad_token_id_tensor = Tensor::new(cfg.audio_pad_token_id, vb.device())?;
// let audio_processor = get_logit_processor(Some(0.8), Some(0.95), Some(25), 34562);
Ok(Self {
transformer,
audio_embeddings,
text_lm_head,
audio_lm_heads,
local_transformer,
audio_assistant_slot_token_id: cfg.audio_assistant_slot_token_id as usize,
audio_end_token_id: cfg.audio_end_token_id as usize,
n_vq: cfg.n_vq,
audio_pad_token_id_tensor,
audio_codebook_sizes: cfg.audio_codebook_sizes.clone(),
audio_temperature: 0.8,
audio_top_k: 25,
audio_top_p: 0.95,
audio_repetition_penalty: 1.2,
})
}
fn build_inputs_embeds(&self, input_ids: &Tensor) -> Result<Tensor> {
let text_ids = input_ids.narrow(D::Minus1, 0, 1)?.squeeze(D::Minus1)?;
let mut inputs_embeds = if let Some(wte) = &self.transformer.wte {
wte.forward(&text_ids)?
} else {
return Err(anyhow!("MossTTS transformer wte can not be none"));
};
for (channel_index, embedding) in self.audio_embeddings.iter().enumerate() {
let channel_ids = input_ids
.narrow(D::Minus1, channel_index + 1, 1)?
.squeeze(D::Minus1)?;
let valid_mask = channel_ids.ne(&self
.audio_pad_token_id_tensor
.broadcast_as(channel_ids.shape())?)?;
let invalid_mask = channel_ids.lt(&channel_ids.zeros_like()?)?;
let embedding_nums = Tensor::new(
self.audio_codebook_sizes[channel_index] as u32,
input_ids.device(),
)?;
let invalid_mask1 =
channel_ids.ge(&embedding_nums.broadcast_as(channel_ids.shape())?)?;
let invalid_mask = valid_mask
.minimum(&invalid_mask.maximum(&invalid_mask1)?)?
.to_dtype(candle_core::DType::U32)?;
if invalid_mask.sum_all()?.to_scalar::<u32>()? > 0 {
return Err(anyhow!("Found out-of-range audio token ids for channel"));
}
let safe_ids = valid_mask.where_cond(&channel_ids, &channel_ids.zeros_like()?)?;
let audio_embeds = embedding.forward(&safe_ids)?;
let audio_embeds = audio_embeds.broadcast_mul(
&valid_mask
.unsqueeze(D::Minus1)?
.to_dtype(audio_embeds.dtype())?,
)?;
inputs_embeds = inputs_embeds.add(&audio_embeds)?;
}
Ok(inputs_embeds)
}
fn sample_next_assistant_text_token(&self, logits: &Tensor) -> Result<usize> {
let logits = logits.squeeze(0)?.squeeze(0)?;
let slot_token_id_logit = logits
.i(self.audio_assistant_slot_token_id)?
.to_dtype(candle_core::DType::F32)?
.to_scalar::<f32>()?;
let end_token_id_logit = logits
.i(self.audio_end_token_id)?
.to_dtype(candle_core::DType::F32)?
.to_scalar::<f32>()?;
println!(
"slot_token_id: {} logit: {slot_token_id_logit}",
self.audio_assistant_slot_token_id
);
println!(
"end_token_id: {} logit: {end_token_id_logit}",
self.audio_end_token_id
);
if slot_token_id_logit > end_token_id_logit {
Ok(self.audio_assistant_slot_token_id)
} else {
Ok(self.audio_end_token_id)
}
}
fn build_generation_row(&self, audio_token_ids: &Tensor) -> Result<Tensor> {
let slot = Tensor::from_slice(
&[self.audio_assistant_slot_token_id as u32],
(1, 1, 1),
audio_token_ids.device(),
)?;
let audio_token_ids = audio_token_ids.unsqueeze(0)?.unsqueeze(0)?;
Ok(Tensor::cat(&[&slot, &audio_token_ids], D::Minus1)?)
}
pub fn generate(&mut self, input_ids: &Tensor, mask: Option<&Tensor>) -> Result<()> {
let sample_len = 2;
let mut seqlen_offset = 0;
let mut seq_len = input_ids.dim(1)?;
let mut generated_frames = vec![];
let mut current_model_input_ids = input_ids.clone();
for step_index in 0..sample_len {
// println!("current_model_input_ids: {:?}", current_model_input_ids);
let inputs_embeds = self.build_inputs_embeds(&current_model_input_ids)?;
let outputs = self.transformer.forward(&inputs_embeds, seqlen_offset)?;
// println!("transformer-----------------------");
let outputs_len = outputs.dim(1)?;
let global_hidden_state = outputs.narrow(1, outputs_len - 1, 1)?;
// println!("global_hidden_state: {}", global_hidden_state);
let mut local_positions = 0usize;
let local_outputs = self
.local_transformer
.forward(&global_hidden_state, local_positions)?;
// println!("local_outputs-----------------------");
let local_len = local_outputs.dim(1)?;
let local_hidden_states = local_outputs.narrow(1, local_len - 1, 1)?;
// println!("local_hidden_states: {}", local_hidden_states);
let text_logits = self.text_lm_head.forward(&local_hidden_states)?;
// println!("text_logits: {}", text_logits.i((0, 0, 0..100))?);
println!("step_index: {}", step_index);
let next_text_token = self.sample_next_assistant_text_token(&text_logits)?;
if next_text_token == self.audio_end_token_id {
self.local_transformer.clear_kv_cache();
break;
}
let mut next_frame_tokens = vec![];
let mut current_local_input = if let Some(wte) = &self.transformer.wte {
wte.forward(&Tensor::from_slice(
&[next_text_token as u32],
(1, 1),
input_ids.device(),
)?)?
} else {
return Err(anyhow!("MossTTS GPT2 wte can not be none"));
};
for channel_index in 0..self.n_vq {
local_positions += 1;
let local_outputs = self
.local_transformer
.forward(&current_local_input, local_positions)?;
let local_len = local_outputs.dim(1)?;
let local_hidden_states = local_outputs.narrow(1, local_len - 1, 1)?;
// println!("local_hidden_states: {local_hidden_states}");
let channel_logits = (&self.audio_lm_heads[channel_index])
.forward(&local_hidden_states)?
.squeeze(0)?
.squeeze(0)?;
// println!("channel_logits: {}", channel_logits.i(0..100)?);
let arg_max = channel_logits.argmax(0)?;
println!("arg_max: {}", arg_max);
let channel_token = simple_sample(
&channel_logits,
true,
Some(self.audio_temperature),
Some(self.audio_top_k),
Some(self.audio_top_p),
Some(&next_frame_tokens),
self.audio_repetition_penalty,
None,
)?;
println!("channel_token: {channel_token}");
next_frame_tokens.push(channel_token);
current_local_input = (&self.audio_embeddings[channel_index]).forward(
&Tensor::from_slice(&[channel_token], (1, 1), input_ids.device())?,
)?;
// println!("current_local_input: {current_local_input}");
}
self.local_transformer.clear_kv_cache();
let next_frame = Tensor::new(next_frame_tokens, input_ids.device())?;
// println!("next_frame: {next_frame}");
current_model_input_ids = self.build_generation_row(&next_frame)?;
seqlen_offset += seq_len;
seq_len = 1;
generated_frames.push(next_frame);
}
let audio_token_ids = Tensor::stack(&generated_frames, 0)?;
println!("audio_token_ids: {audio_token_ids}");
Ok(())
}
}
+1 -1
View File
@@ -1,6 +1,6 @@
use std::time::Instant;
use crate::models::common::generate::get_logit_processor;
use crate::models::common::sample::get_logit_processor;
use crate::params::chat::{
ChatCompletionChunkResponse, ChatCompletionParameters, ChatCompletionResponse,
};
+2 -2
View File
@@ -3,8 +3,8 @@ use std::time::Instant;
use crate::{
models::common::{
MultiModalData,
generate::{GenerationContext, generate_generic_text, get_logit_processor},
modules::{AsrResult, VadFrameResult},
generate::{GenerationContext, generate_generic_text},
modules::{AsrResult, VadFrameResult}, sample::get_logit_processor,
},
params::chat::{ChatCompletionChunkResponse, ChatCompletionParameters, ChatCompletionResponse},
utils::response_utils::{build_chunk_response_with_usage, build_completion_response_with_time},
+1 -1
View File
@@ -82,7 +82,7 @@ impl Qwen3AsrProcessor {
}
pub fn extract_audio_vec(&self, mes: &ChatCompletionParameters) -> Result<Vec<Tensor>> {
let audio_tensors = extract_audios(mes, &self.device, Some(self.sample_rate))?;
let audio_tensors = extract_audios(mes, &self.device, Some(self.sample_rate), Some(1))?;
audio_tensors.iter().map(float_range_normalize).collect()
}
+10 -4
View File
@@ -534,7 +534,8 @@ impl VoxCPMModel {
let audio_start = Tensor::new(vec![self.audio_start_token], &self.device)?;
let text_token = Tensor::cat(&[text_token, audio_start], D::Minus1)?;
let text_length = text_token.dim(0)?;
let mut audio = load_audio_with_resample(&path, &self.device, Some(self.sample_rate))?;
let mut audio =
load_audio_with_resample(&path, &self.device, Some(self.sample_rate), Some(1))?;
let patch_len = self.patch_size * self.chunk_size;
if audio.dim(1)? % patch_len != 0 {
audio =
@@ -574,7 +575,8 @@ impl VoxCPMModel {
let audio_start = Tensor::new(vec![self.audio_start_token], &self.device)?;
let text_token = Tensor::cat(&[text_token, audio_start], D::Minus1)?;
let text_length = text_token.dim(0)?;
let mut audio = load_audio_with_resample(&path, &self.device, Some(self.sample_rate))?;
let mut audio =
load_audio_with_resample(&path, &self.device, Some(self.sample_rate), Some(1))?;
let patch_len = self.patch_size * self.chunk_size;
if audio.dim(1)? % patch_len != 0 {
audio =
@@ -841,8 +843,12 @@ impl VoxCPMModel {
) -> Result<HashMap<String, Tensor>> {
let text_token = self.tokenizer.encode(prompt_text)?;
let text_token = Tensor::from_slice(&text_token, text_token.len(), &self.device)?;
let mut audio =
load_audio_with_resample(&prompt_wav_path, &self.device, Some(self.sample_rate))?;
let mut audio = load_audio_with_resample(
&prompt_wav_path,
&self.device,
Some(self.sample_rate),
Some(1),
)?;
let patch_len = self.patch_size * self.chunk_size;
if audio.dim(1)? % patch_len != 0 {
audio = audio.pad_with_zeros(D::Minus1, 0, patch_len - audio.dim(1)? % patch_len)?;
+8 -3
View File
@@ -38,8 +38,12 @@ impl VoxCPMProcessor {
audio_vae: &AudioVAE,
) -> Result<HashMap<String, Tensor>> {
let (text_token, _) = tokenizer.encode_tensor(prompt_text, &self.device)?;
let mut audio =
load_audio_with_resample(&prompt_wav_path, &self.device, Some(self.sample_rate))?;
let mut audio = load_audio_with_resample(
&prompt_wav_path,
&self.device,
Some(self.sample_rate),
Some(1),
)?;
let patch_len = self.patch_size * self.chunk_size;
if audio.dim(1)? % patch_len != 0 {
audio = audio.pad_with_zeros(D::Minus1, 0, patch_len - audio.dim(1)? % patch_len)?;
@@ -74,7 +78,8 @@ impl VoxCPMProcessor {
let mut text_token = Tensor::cat(&[text_token, audio_start], D::Minus1)?;
let (audio_feat, audio_mask) = if let Some(path) = prompt_wav_path {
let mut audio = load_audio_with_resample(&path, &self.device, Some(self.sample_rate))?;
let mut audio =
load_audio_with_resample(&path, &self.device, Some(self.sample_rate), Some(1))?;
let patch_len = self.patch_size * self.chunk_size;
if audio.dim(1)? % patch_len != 0 {
audio =
+118 -56
View File
@@ -1,4 +1,4 @@
use anyhow::{Result, anyhow};
use anyhow::Result;
use candle_core::{D, DType, Device, IndexOp, Tensor};
use candle_transformers::models::deepseek2::SplitOp;
@@ -21,6 +21,22 @@ pub fn rotate_half(x: &Tensor) -> Result<Tensor> {
Ok(rotate_x)
}
pub fn rotate_half_interleave(x: &Tensor) -> Result<Tensor> {
let x_rank = x.rank();
let x_dim = x.dims();
let half_dim = x_dim[x_rank - 1] / 2;
let mut x_reshape = x_dim[0..x_rank - 1].to_vec();
x_reshape.push(half_dim);
x_reshape.push(2);
let x = x.reshape(x_reshape)?;
let even = x.narrow(D::Minus1, 0, 1)?;
let odd = x.narrow(D::Minus1, 1, 1)?.affine(-1.0, 0.0)?;
let rotate_x = Tensor::cat(&[&odd, &even], D::Minus1)?
.reshape(x_dim)?
.contiguous()?;
Ok(rotate_x)
}
pub fn apply_multimodel_rotary_pos_emb(
q: &Tensor,
k: &Tensor,
@@ -115,6 +131,44 @@ pub fn apply_rotary_pos_emb(
Ok((q_embed, k_embed))
}
pub fn apply_rotary_pos_emb_interleave(
q: &Tensor,
k: &Tensor,
cos: &Tensor,
sin: &Tensor,
tof32: bool,
) -> Result<(Tensor, Tensor)> {
// sin/cos: to (bs, 1, seq_len, head_dim)
// q/k: (bs, n_head, seq_len, head_dim)
let mut cos = cos.clone();
let mut sin = sin.clone();
if cos.rank() == 2 {
// (seq_len, head_dim) -> (1, 1, seq_len, head_dim)
cos = cos.unsqueeze(0)?.unsqueeze(0)?;
sin = sin.unsqueeze(0)?.unsqueeze(0)?;
}
if cos.rank() == 3 {
// (bs, seq_len, head_dim) -> (bs, 1, seq_len, head_dim)
cos = cos.unsqueeze(1)?;
sin = sin.unsqueeze(1)?;
}
let orig_dtype = q.dtype();
let q = if tof32 { &q.to_dtype(DType::F32)? } else { q };
let k = if tof32 { &k.to_dtype(DType::F32)? } else { k };
let cos = cos.to_dtype(q.dtype())?;
let sin = sin.to_dtype(q.dtype())?;
let q_embed = q
.broadcast_mul(&cos)?
.add(&rotate_half_interleave(q)?.broadcast_mul(&sin)?)?
.to_dtype(orig_dtype)?;
let k_embed = k
.broadcast_mul(&cos)?
.add(&rotate_half_interleave(k)?.broadcast_mul(&sin)?)?
.to_dtype(orig_dtype)?;
Ok((q_embed, k_embed))
}
pub fn glm_asr_apply_rotary_pos_emb(
q: &Tensor,
k: &Tensor,
@@ -258,66 +312,46 @@ pub fn glm_ocr_apply_rotary_pos_emb(
Ok((q_embed, k_embed))
}
pub fn roformer_rotate(x: &Tensor) -> Result<Tensor> {
let dims = x.dims();
let last_dim = dims
.last()
.ok_or(anyhow!("Input tensor must have at least one dimension"))?;
if last_dim % 2 != 0 {
return Err(anyhow!(
"Last dimension size must be even, got {}",
last_dim
));
}
let new_dims: Vec<usize> = dims[..dims.len() - 1]
.iter()
.copied()
.chain([last_dim / 2, 2])
.collect();
let x_reshape = x.reshape(new_dims)?;
let x_chunks = x_reshape.chunk(2, D::Minus1)?;
let x1 = &x_chunks[0];
let x2 = &x_chunks[1];
// let x1 = x_reshape.narrow(D::Minus1, 0, 1)?;
// let x2 = x_reshape.narrow(D::Minus1, 1, 1)?;
let x2_neg = x2.affine(-1.0, 0.0)?;
let rotate_x = Tensor::cat(&[&x2_neg, x1], D::Minus1)?;
Ok(rotate_x.flatten(D::Minus2, D::Minus1)?)
}
pub fn apply_rotary_pos_emb_roformer(
q: &Tensor,
k: &Tensor,
cos: &Tensor,
sin: &Tensor,
tof32: bool,
) -> Result<(Tensor, Tensor)> {
let mut cos = cos.clone();
let mut sin = sin.clone();
if cos.rank() == 2 {
// (seq_len, head_dim) -> (1, 1, seq_len, head_dim)
cos = cos.unsqueeze(0)?.unsqueeze(0)?;
sin = sin.unsqueeze(0)?.unsqueeze(0)?;
}
if cos.rank() == 3 {
// (bs, seq_len, head_dim) -> (bs, 1, seq_len, head_dim)
cos = cos.unsqueeze(1)?;
sin = sin.unsqueeze(1)?;
}
let orig_dtype = q.dtype();
let q = if tof32 { &q.to_dtype(DType::F32)? } else { q };
let k = if tof32 { &k.to_dtype(DType::F32)? } else { k };
let cos = cos.to_dtype(q.dtype())?;
let sin = sin.to_dtype(q.dtype())?;
let q_embed = q
.broadcast_mul(&cos)?
.add(&roformer_rotate(q)?.broadcast_mul(&sin)?)?
.to_dtype(orig_dtype)?;
let k_embed = k
.broadcast_mul(&cos)?
.add(&roformer_rotate(k)?.broadcast_mul(&sin)?)?
.to_dtype(orig_dtype)?;
Ok((q_embed, k_embed))
let ori_dtype = q.dtype();
let (bs, n_head, seq_len, dim) = q.dims4()?;
let half_dim = dim / 2;
let rotr = cos
.narrow(D::Minus1, 0, half_dim)?
.to_dtype(candle_core::DType::F32)?;
let roti = sin
.narrow(D::Minus1, 0, half_dim)?
.to_dtype(candle_core::DType::F32)?;
let q = q
.reshape((bs, n_head, seq_len, half_dim, 2))?
.to_dtype(candle_core::DType::F32)?;
let qr = q.narrow(D::Minus1, 0, 1)?.squeeze(D::Minus1)?;
let qi = q.narrow(D::Minus1, 1, 1)?.squeeze(D::Minus1)?;
let k = k
.reshape((bs, n_head, seq_len, half_dim, 2))?
.to_dtype(candle_core::DType::F32)?;
let kr = k.narrow(D::Minus1, 0, 1)?.squeeze(D::Minus1)?;
let ki = k.narrow(D::Minus1, 1, 1)?.squeeze(D::Minus1)?;
let qor = qr.broadcast_mul(&rotr)?.sub(&qi.broadcast_mul(&roti)?)?;
let qoi = qr.broadcast_mul(&roti)?.add(&qi.broadcast_mul(&rotr)?)?;
let kor = kr.broadcast_mul(&rotr)?.sub(&ki.broadcast_mul(&roti)?)?;
let koi = kr.broadcast_mul(&roti)?.add(&ki.broadcast_mul(&rotr)?)?;
let q = Tensor::stack(&[qor, qoi], D::Minus1)?
.reshape((bs, n_head, seq_len, dim))?
.to_dtype(ori_dtype)?;
let k = Tensor::stack(&[kor, koi], D::Minus1)?
.reshape((bs, n_head, seq_len, dim))?
.to_dtype(ori_dtype)?;
Ok((q, k))
}
#[derive(Debug, Clone)]
@@ -554,7 +588,6 @@ impl RoPE {
pub fn new(dim: usize, theta_base: f32, device: &Device) -> Result<Self> {
let inv_freq = compute_default_rope_parameters(dim, theta_base);
let inv_freq = Tensor::from_slice(&inv_freq, (1, inv_freq.len()), device)?;
Ok(Self { inv_freq })
}
pub fn forward(
@@ -577,6 +610,35 @@ impl RoPE {
let sin = emb.sin()?;
Ok((cos, sin))
}
pub fn forward_repeat_interleave(
&self,
seqlen_offset: usize,
seq_len: usize,
device: &Device,
) -> Result<(Tensor, Tensor)> {
let positions = Tensor::arange(
seqlen_offset as f32,
(seqlen_offset + seq_len) as f32,
self.inv_freq.device(),
)?
.reshape((seq_len, 1))?; // (seq_len, 1)
let freqs = positions.matmul(&self.inv_freq)?; // (seq_len, dim / 2)
let cos = freqs
.cos()?
.unsqueeze(D::Minus1)?
.repeat((1, 1, 2))?
.flatten_from(D::Minus2)?
.contiguous()?
.to_device(device)?;
let sin = freqs
.sin()?
.unsqueeze(D::Minus1)?
.repeat((1, 1, 2))?
.flatten_from(D::Minus2)?
.contiguous()?
.to_device(device)?;
Ok((cos, sin))
}
}
pub fn get_xd_cos_sin(
+11 -4
View File
@@ -120,15 +120,22 @@ impl TokenizerModel {
}
}
pub fn sentencepiece_encode_vec(
text: &str,
tokenizer: &SentencePieceProcessor,
) -> Result<Vec<u32>> {
let tokens = tokenizer
.encode(text)
.map_err(|e| anyhow!(format!("tokenizer encode error:{}", e)))?;
Ok(tokens.iter().map(|p| p.id).collect::<Vec<u32>>())
}
pub fn sentencepiece_encode(
text: &str,
tokenizer: &SentencePieceProcessor,
device: &Device,
) -> Result<Tensor> {
let tokens = tokenizer
.encode(text)
.map_err(|e| anyhow!(format!("tokenizer encode error:{}", e)))?;
let token_ids = tokens.iter().map(|p| p.id).collect::<Vec<u32>>();
let token_ids = sentencepiece_encode_vec(text, tokenizer)?;
let tokens_t = Tensor::new(token_ids, device)?.unsqueeze(0)?;
Ok(tokens_t)
}
+35 -12
View File
@@ -470,7 +470,14 @@ pub fn get_audio_format_from_bytes(bytes: &[u8]) -> Result<String> {
}
}
pub fn load_audio_use_symphonia(audio_vec: Vec<u8>, device: &Device) -> Result<(Tensor, usize)> {
/// return
/// audio shape: (channel, audio_len)
/// sample_rate: usize
pub fn load_audio_use_symphonia(
audio_vec: Vec<u8>,
device: &Device,
target_channels: usize,
) -> Result<(Tensor, usize)> {
let extension = get_audio_format_from_bytes(&audio_vec)?;
let content = Cursor::new(audio_vec);
let mss = MediaSourceStream::new(Box::new(content), Default::default());
@@ -558,16 +565,26 @@ pub fn load_audio_use_symphonia(audio_vec: Vec<u8>, device: &Device) -> Result<(
}
}
let mut audio_tensor = Tensor::new(all_samples, device)?;
if channels > 1 {
// 对channel通道求平均, channel维度变为1
audio_tensor = audio_tensor.mean_keepdim(0)?;
if target_channels == channels {
return Ok((audio_tensor, sample_rate as usize));
}
Ok((audio_tensor, sample_rate as usize))
}
pub fn load_audio(path: &str, device: &Device) -> Result<(Tensor, usize)> {
let audio_vec = get_audio_bytes_vec(path)?;
load_audio_use_symphonia(audio_vec, device)
audio_tensor = if channels == 1 {
// Mono to Multi-channel: Repeat
audio_tensor.repeat((target_channels, 1))?
} else if target_channels == 1 {
// Multi-channel to Mono: Mean
audio_tensor.mean_keepdim(0)?
} else {
// Unsupported conversion (e.g., Stereo to 5.1)
return Err(anyhow!(
"target_channels: {}, audio channels: {}, can't change directly",
target_channels,
channels
));
};
Ok((audio_tensor, sample_rate as usize))
}
pub fn resample_audio_from_vec_f32(
@@ -599,12 +616,14 @@ pub fn resample_audio_from_vec_f32(
Ok(audio)
}
/// return shape: (channel, audio_len)
pub fn resample_audio_from_bytes(
audio_vec: Vec<u8>,
device: &Device,
target_sample_rate: Option<usize>,
target_channels: usize,
) -> Result<Tensor> {
let (mut audio, sr) = load_audio_use_symphonia(audio_vec, device)?;
let (mut audio, sr) = load_audio_use_symphonia(audio_vec, device, target_channels)?;
if let Some(target_sample_rate) = target_sample_rate
&& target_sample_rate != sr
{
@@ -613,16 +632,19 @@ pub fn resample_audio_from_bytes(
Ok(audio)
}
/// return shape: (channel, audio_len)
pub fn load_audio_with_resample(
path: &str,
device: &Device,
target_sample_rate: Option<usize>,
target_channels: Option<usize>,
) -> Result<Tensor> {
// hound 只支持wav文件
// let audio_path = get_audio_path(path)?;
// let (mut audio, sr) = load_audio_use_hound(audio_path, device)?;
let target_channels = target_channels.unwrap_or(1);
let audio_vec = get_audio_bytes_vec(path)?;
resample_audio_from_bytes(audio_vec, device, target_sample_rate)
resample_audio_from_bytes(audio_vec, device, target_sample_rate, target_channels)
}
pub fn save_wav(audio: &Tensor, save_path: &str, sample_rate: u32) -> Result<()> {
@@ -692,12 +714,13 @@ pub fn extract_audios(
mes: &ChatCompletionParameters,
device: &Device,
target_sample_rate: Option<usize>,
target_channels: Option<usize>,
) -> Result<Vec<Tensor>> {
let audio_url_vec = extract_audio_url(mes);
// 并行加载音频
audio_url_vec
.par_iter()
.map(|url| load_audio_with_resample(url, device, target_sample_rate))
.map(|url| load_audio_with_resample(url, device, target_sample_rate, target_channels))
.collect()
// #[cfg(not(feature = "ffmpeg"))]
// {
+60
View File
@@ -721,6 +721,66 @@ pub fn bucketize(input: usize, boundaries: &[usize]) -> Result<usize> {
// Ok(index)
}
pub fn contains_cjk(text: &str) -> bool {
for ch in text.chars() {
let c = ch as u32;
if (0x4e00..=0x9fff).contains(&c) // CJK Unified Ideographs
|| (0x3400..=0x4dbf).contains(&c) // CJK Unified Ideographs Extension A
|| (0x3040..=0x30ff).contains(&c) // Hiragana and Katakana
|| (0xac00..=0xd7af).contains(&c) // Hangul Syllables
{
return true;
}
}
false
}
pub fn prepare_tts_text(text: &str) -> Result<String> {
let mut normalized_text = text.trim().to_string();
if normalized_text.eq("") {
return Err(anyhow!("Text cannot be empty."))
}
normalized_text = normalized_text.replace('\n', " ").replace('\r', " ");
while normalized_text.contains(" ") {
normalized_text = normalized_text.replace(" ", " ");
}
if contains_cjk(&normalized_text) {
let cjk_end_punctuations = ['。', '', '', '…', '.', '!', '?'];
if !normalized_text.ends_with(|c: char| cjk_end_punctuations.contains(&c)) {
normalized_text.push('。');
}
return Ok(normalized_text);
}
// Non-CJK (English/Western) logic
// Capitalize first letter if it's lowercase alphabetic
if let Some(first_char) = normalized_text.chars().next() {
if first_char.is_ascii_lowercase() {
let mut chars = normalized_text.chars();
chars.next(); // consume first char
let rest: String = chars.collect();
normalized_text = format!("{}{}", first_char.to_ascii_uppercase(), rest);
}
}
// Add period if ends with alphanumeric
if let Some(last_char) = normalized_text.chars().last() {
if last_char.is_alphanumeric() {
normalized_text.push('.');
}
}
// Add padding if less than 5 words
// Split by whitespace to count words
let word_count = normalized_text.split_whitespace().count();
if word_count < 5 {
normalized_text = format!(" {}", normalized_text); // 8 spaces
}
Ok(normalized_text)
}
#[cfg(test)]
mod tests {
use super::*;
+21 -9
View File
@@ -1,13 +1,5 @@
use aha::models::{
deepseek_ocr::config::DeepseekOCRConfig,
hunyuan_ocr::config::HunYuanVLConfig,
lfm2::config::Lfm2Config,
lfm2vl::config::{Lfm2ProcessorConfig, Lfm2VLConfig},
minicpm4::config::MiniCPM4Config,
paddleocr_vl::config::PaddleOCRVLConfig,
qwen2_5vl::config::Qwen2_5VLConfig,
qwen3vl::config::Qwen3VLConfig,
voxcpm::config::VoxCPMConfig,
deepseek_ocr::config::DeepseekOCRConfig, hunyuan_ocr::config::HunYuanVLConfig, lfm2::config::Lfm2Config, lfm2vl::config::{Lfm2ProcessorConfig, Lfm2VLConfig}, minicpm4::config::MiniCPM4Config, moss::config::{MossAudioTokenizerConfig, MossTTSConfig}, paddleocr_vl::config::PaddleOCRVLConfig, qwen2_5vl::config::Qwen2_5VLConfig, qwen3vl::config::Qwen3VLConfig, voxcpm::config::VoxCPMConfig
};
use anyhow::Result;
@@ -117,3 +109,23 @@ fn lfm2vl_config() -> Result<()> {
println!("{:?}", processor_config);
Ok(())
}
#[test]
fn moss_audio_tokenizer_config() -> Result<()> {
// cargo test -F cuda --test config_tests moss_audio_tokenizer_config -r -- --nocapture
let model_path = "/home/jhq/.aha/openmoss/MOSS-Audio-Tokenizer-Nano/";
let config_path = model_path.to_string() + "/config.json";
let config: MossAudioTokenizerConfig = serde_json::from_slice(&std::fs::read(config_path)?)?;
println!("{:?}", config);
Ok(())
}
#[test]
fn moss_tts_config() -> Result<()> {
// cargo test -F cuda --test config_tests moss_tts_config -r -- --nocapture
let model_path = "/home/jhq/.aha/openmoss/MOSS-TTS-Nano/";
let config_path = model_path.to_string() + "/config.json";
let config: MossTTSConfig = serde_json::from_slice(&std::fs::read(config_path)?)?;
println!("{:?}", config);
Ok(())
}
+20
View File
@@ -0,0 +1,20 @@
use aha::models::moss::generate::MossTTSGenerate;
use anyhow::Result;
#[test]
fn moss_tts() -> Result<()> {
// RUST_BACKTRACE=1 cargo test -F cuda --test test_moss_tts moss_tts -r -- --nocapture
let save_dir =
aha::utils::get_default_save_dir().ok_or(anyhow::anyhow!("Failed to get save dir"))?;
let tts_path = format!("{}/openmoss/MOSS-TTS-Nano/", save_dir);
let audio_tokenizer_path = format!("{}/openmoss/MOSS-Audio-Tokenizer-Nano/", save_dir);
let mut model = MossTTSGenerate::init(&tts_path, &audio_tokenizer_path, None, None)?;
let _ = model.generate(
"您好啊,吃饭了吗,吃的啥啊中午",
Some("file://./assets/audio/jiangjiang.wav"),
Some("哈喽大家好,我是蒋蒋"),
Some(aha::models::moss::tts_nano::MossTTSMode::Continuation),
// None,
)?;
Ok(())
}
+27
View File
@@ -425,3 +425,30 @@ fn silero_vad_weight() -> Result<()> {
Ok(())
}
#[test]
fn moss_tts_nano_weight() -> Result<()> {
// cargo test -F cuda --test weight_test moss_tts_nano_weight -r -- --nocapture
let save_dir =
aha::utils::get_default_save_dir().ok_or(anyhow::anyhow!("Failed to get save dir"))?;
let model_path = format!("{}/openmoss/MOSS-TTS-Nano/pytorch_model.bin", save_dir);
let dict = read_all_with_key(&model_path, None)?;
for (k, v) in dict {
println!("key: {}, tensor shape: {:?}", k, v);
}
Ok(())
}
#[test]
fn moss_audio_tokenizer_nano_weight() -> Result<()> {
// cargo test -F cuda --test weight_test moss_audio_tokenizer_nano_weight -r -- --nocapture
let save_dir =
aha::utils::get_default_save_dir().ok_or(anyhow::anyhow!("Failed to get save dir"))?;
let model_path = format!("{}/openmoss/MOSS-Audio-Tokenizer-Nano/model-00001-of-00001.safetensors", save_dir);
let device = get_device(None);
let weights = safetensors::load(model_path, &device)?;
for (key, tensor) in weights.iter() {
println!("=== {} === {:?}", key, tensor);
}
Ok(())
}