CandleでMambaモデルを試すついでにコンテキストの内容をファイルへ保存して復元できるか試してみました。
はじめに
Mambaは、Self-Attentionで構成された一般的なTransformerとは異なり、選択的状態空間モデル(selective state space model)という構造を採用しているようです。
例えば、CandleにおけるMambaのモデル定義は次のようになっています。
モデルのforward時にコンテキストとなるStateを入力とは別の引数として外部から与えるのが特徴です。1
pub struct State {
pub hs: Vec<Tensor>,
pub prev_xs: Vec<[Tensor; D_CONV]>,
pub pos: usize,
}
#[derive(Clone, Debug)]
pub struct MambaBlock {
in_proj: Linear,
conv1d_bias: Tensor,
conv1d_weights: [Tensor; D_CONV],
x_proj: Linear,
dt_proj: Linear,
a_log: Tensor,
d: Tensor,
out_proj: Linear,
dt_rank: usize,
layer_index: usize,
d_inner: usize,
}
#[derive(Clone, Debug)]
pub struct ResidualBlock {
mixer: MambaBlock,
norm: RmsNorm,
}
#[derive(Clone, Debug)]
pub struct Model {
embedding: candle_nn::Embedding,
layers: Vec<ResidualBlock>,
norm_f: RmsNorm,
lm_head: Linear,
dtype: DType,
}
impl Model {
...
pub fn forward(&self, input_ids: &Tensor, state: &mut State) -> Result<Tensor> {
let _b_size = input_ids.dims1()?;
let mut xs = self.embedding.forward(input_ids)?;
for layer in self.layers.iter() {
xs = layer.forward(&xs, state)?
}
state.pos += 1;
xs.apply(&self.norm_f)?.apply(&self.lm_head)
}
...
}
この仕組みなら、同じコンテキストで複数モデルを同時実行して結果を比較・選定したり等、なかなか面白い事ができそうです。2
また、処理する度に肥大化していく(一般的な)KVキャッシュに比べ、Stateはレイヤー数に応じた固定サイズとなっているのもポイントだと考えます。3
現時点で最新のモデルはMamba 3ですが、Candle v0.11に3のモデル定義は無いので、ここではMamba 2(candle_transformers::models::mamba2)で文章生成を行ってみました。
ついでに、処理後のStateをファイルへ保存しておき、次回実行時に復元してみます。
ちなみに、mamba2::Stateは上記のmamba::Stateとは内容が少し異なります。
pub struct State {
pub hs: Vec<Tensor>,
pub conv_states: Vec<Tensor>,
pub pos: usize,
}
注意点として、Mambaのモデルは基本的にhttps://huggingface.co/state-spacesから取得できますが、Mamba 2系に関してはcandle_transformers::models::mamba2がここのモデルをサポートしておらず、代わりにhttps://huggingface.co/AntonVのmamba2-xxx-hfモデルを使う必要がありました。4
実装
処理の基本的な流れは「CandleでローカルLLMを実行する」と同じですが、入力トークンを1つずつforwardしてStateを更新している点が異なっています。
mamba2::Model には、入力トークンを一括処理するforward_prefillも用意されていますが、ここでは通常のforwardを使いました。
今回使用したMambaモデルは、同じフレーズを繰り返し出力する傾向が強かったため、下記を適用して指定トークン(直近の生成トークン)のスコアを下げる調整を行なっています。
- candle_transformers::utils::apply_repeat_penalty
また、プロンプトの続きから文章を生成するモデルのため、入力プロンプトをそのまま出力(print)しています。
[dependencies]
candle-core = { version = "0.11", features = ["metal"] }
candle-nn = { version = "0.11", features = ["metal"] }
candle-transformers = { version = "0.11", features = ["metal"] }
memmap2 = "0.9.11"
safetensors = "0.8.0"
serde_json = "1.0"
tokenizers = "0.23"
use candle_transformers::models::mamba2::{Config, Model, State};
...
const CONFIG_FILE: &str = "model/config.json";
const TOKENIZER_FILE: &str = "model/tokenizer.json";
const MODEL_FILE: &str = "model/model.safetensors";
const KEY_HS: &str = "hs-";
const KEY_CONV_STATES: &str = "conv-states-";
const KEY_POS: &str = "pos";
const TEMPERATURE: Option<f64> = Some(0.7);
const TOP_P: Option<f64> = Some(0.7);
const REPEAT_PENALTY: f32 = 1.1;
const REPEAT_LAST_N: usize = 32;
const REPEAT_PENALTY_MIN_TOKEN_SIZE: Option<usize> = Some(5);
const MAX_SAMPLE_LEN: usize = 100;
fn main() -> Result<()> {
let device = Device::new_metal(0)?; // Metal適用(macOS環境用)
let dtype = DType::F32;
let mut args = env::args().skip(1);
let prompt = args.next().ok_or("prompt")?;
let seed = args.next().and_then(|x| x.parse().ok()).unwrap_or(12345);
let output_state_file = args.next(); // 出力(save)Stateファイル名
let input_state_file = args.next(); // 入力(load)Stateファイル名
let tokenizer = Tokenizer::from_file(TOKENIZER_FILE)?;
let config: Config = serde_json::from_str(&std::fs::read_to_string(CONFIG_FILE)?)?;
let vb = unsafe { VarBuilder::from_mmaped_safetensors(&[MODEL_FILE], dtype, &device)? };
let model = Model::new(&config, vb.pp("backbone"))?;
let mut logits_proc = LogitsProcessor::new(seed, TEMPERATURE, TOP_P);
// 入力プロンプトの出力
print!("{prompt}");
let tokens = tokenizer.encode(prompt, true)?.get_ids().to_vec();
let eos_token = tokenizer
.token_to_id("<|endoftext|>")
.ok_or("not found endoftext token")?;
let mut state = State::new(1, &config, dtype, &device)?;
// Stateの復元
if let Some(input_file) = input_state_file {
load_state(&mut state, &input_file, &device)?;
}
let mut next_logits = None;
// 入力トークン(プロンプト)の処理
for t in tokens {
let input = Tensor::new(&[t], &device)?;
let logits = model.forward(&input, &mut state)?;
next_logits = Some(logits);
}
let mut output_token_ids = vec![];
// 文章生成
for _ in 0..MAX_SAMPLE_LEN {
let logits = next_logits
.ok_or("no token result")?
.squeeze(0)?
.to_dtype(DType::F32)?;
let logits = apply_repeat_penalty(
&logits,
REPEAT_LAST_N,
REPEAT_PENALTY,
&output_token_ids,
REPEAT_PENALTY_MIN_TOKEN_SIZE,
)?;
let token_id = logits_proc.sample(&logits)?;
if token_id == eos_token {
break;
}
output_token_ids.push(token_id);
let token = tokenizer.decode(&[token_id], true)?;
print!("{token}");
let input = Tensor::new(&[token_id], &device)?;
next_logits = Some(model.forward(&input, &mut state)?);
}
// Stateの保存
if let Some(output_file) = output_state_file {
save_state(&state, &output_file)?;
}
Ok(())
}
fn apply_repeat_penalty(
logits: &Tensor,
repeat_last_n: usize,
repeat_penalty: f32,
token_ids: &Vec<u32>,
min_token_size: Option<usize>,
) -> Result<Tensor> {
if token_ids.len() >= min_token_size.unwrap_or_default().max(1) {
let idx = token_ids.len().saturating_sub(repeat_last_n);
// 指定トークンに対してペナルティを適用(スコアを低下させる)
utils::apply_repeat_penalty(&logits, repeat_penalty, &token_ids[idx..])
.map_err(|e| e.into())
} else {
Ok(logits.to_owned())
}
}
Stateの保存と復元はこのようにしました。
fn save_state(state: &State, file: &str) -> Result<()> {
let device = state.hs.first().map(|x| x.device()).ok_or("no tensor")?;
let mut state_data = vec![];
for (i, t) in state.hs.iter().enumerate() {
state_data.push((format!("{KEY_HS}{i:03}"), t));
}
for (i, t) in state.conv_states.iter().enumerate() {
state_data.push((format!("{KEY_CONV_STATES}{i:03}"), t));
}
let pos = Tensor::from_slice(&[state.pos as u32], 1, device)?;
state_data.push((KEY_POS.into(), &pos));
safetensors::tensor::serialize_to_file(state_data, None, file.as_ref())?;
Ok(())
}
fn load_state(state: &mut State, file: &str, device: &Device) -> Result<()> {
let f = File::open(file)?;
let buffer = unsafe { memmap2::MmapOptions::new().map(&f)? };
let ts = safetensors::SafeTensors::deserialize(&buffer)?;
let mut keys = ts.names();
keys.sort();
let mut hs = vec![];
let mut conv_states = vec![];
let mut pos = 0;
for key in keys {
let t = ts.tensor(key)?.load(device)?;
if key.starts_with(KEY_HS) {
hs.push(t);
} else if key.starts_with(KEY_CONV_STATES) {
conv_states.push(t);
} else if key.eq(KEY_POS) {
pos = t.get(0)?.to_scalar::<u32>()? as usize;
}
}
// Stateへ設定
state.hs = hs;
state.conv_states = conv_states;
state.pos = pos;
Ok(())
}
項目名を付けてStateの内容をsafetensorsファイルへ出力し、ファイルから復元してStateへ反映しているだけです。
復元時の注意点として、serialize_to_fileした際の要素の並びはdeserialize時には再現されず、順序がバラバラになってしまうようなので、項目名をソートして処理しています。
また、試しにhs(隠れ状態)だけを保存して復元してみましたが、復元時の入力プロンプトの内容が上手く反映されない等の不都合が見られました。
動作確認
ここでは、https://huggingface.co/AntonV/mamba2-1.3b-hfから次のファイルを取得して使用しました。
- config.json
- model.safetensors
- tokenizer.json
ただし、config.jsonをserde_json::from_strする際に下記の箇所でエラーとなるため、Infinityを適当な数値へ変更するか、time_step_limitの設定自体を削除する必要がありました。5
{
...
"time_step_limit": [
0.0,
Infinity
],
...
}
実行
前回保存したsafetensorsファイルを使用するようにして、入力プロンプトを適度に変えて3回実施した結果はこのようになりました。
内容の質はともかく6、コンテキストは引き継げているように見えます。
ステップ1
$ cargo run --release "The Object Oriented Programming is" 12345 s1.safetensors > s1.txt
The Object Oriented Programming is a method of programming that has been used for many years in the development and maintenance of software. In object oriented programming, programs are divided into classes or modules which can be called upon by other parts to perform certain tasks. This allows for modularity and flexibility in the program as well as allowing for easy modification of programs without having to rewrite large sections of code that may be used by multiple modules.
The use and development of object oriented programming has been relatively slow in the computer industry as the need for
ステップ2
$ cargo run --release "The merit is" 12345 s2.safetensors s1.safetensors > s2.txt
The merit is that the programmer can easily change the program without having to rewrite large sections of code.
The object oriented programming approach has been used in the development and maintenance software for many years, however it is only recently that this approach have become more widely accepted. The main reason being the ability to easily change or modify a program without having to rewrite large sections of code which may be used by multiple modules.
The use and development of object oriented programming has been relatively slow in the computer industry as it is only
ステップ3
$ cargo run --release "The problem is" 12345 s3.safetensors s2.safetensors > s3.txt
The problem is that the programmer must know what methods and properties are available for a particular object. This is not always easy to determine as many objects may have similar or identical properties, such as having the same method names and signatures but different parameter lists (parameters).
The problem is that a programmer must know what methods are available for an object. This is not always easy to determine as many objects may have similar or identical properties, such as having the same method names and signatures but different parameter lists (parameters).
最後に、モデルによるState保存ファイルサイズの違いはこのようになりました。
| モデル | レイヤー数 | Stateファイルサイズ |
|---|---|---|
| mamba2-1.3b-hf | 48 | 約105MB |
| mamba2-130m-hf | 24 | 約20MB |
-
Transformerにおける一般的なKVキャッシュではなく、このStateでコンテキストを管理するようです。ちなみに、Python用ライブラリのTransformersでは
cache_paramsという引数でStateを与えるようになっていました ↩ -
コンテキストをGitのようにバージョン管理するとか(データサイズ次第ですが)、コンテキスト制御の可能性も考えられそうです ↩
-
Vecの要素数がレイヤー数となっています。ただし、常に一定サイズが必要になる点やコンテキストを凝縮するデメリットはありそうです。 ↩
-
config.jsonの内容が大きく異なっており、互換性がありません。candle-examplesでは、Mamba 1系(便宜上
1としています)のモデルはstate-spacesから取得し、Mamba 2系のモデルはAntonVから取得するようになっています ↩ -
mamba2::Configが使用しない項目なので無くても支障はありません ↩
-
指定のトークン数で処理を切り上げている点も影響しているかもしれません ↩