#[derive(Debug, Clone, Copy, PartialEq, Eq)]
pub enum GemmaQuantBits {
Two,
Four,
Eight,
}
impl GemmaQuantBits {
pub fn from_num_bits(n: u32) -> Option<Self> {
match n {
2 => Some(Self::Two),
4 => Some(Self::Four),
8 => Some(Self::Eight),
_ => None,
}
}
pub fn values_per_byte(self) -> usize {
match self {
Self::Two => 4,
Self::Four => 2,
Self::Eight => 1,
}
}
fn field_bits(self) -> u32 {
match self {
Self::Two => 2,
Self::Four => 4,
Self::Eight => 8,
}
}
fn zero_point(self) -> i32 {
match self {
Self::Two => 2,
Self::Four => 8,
Self::Eight => 0,
}
}
pub fn packed_len(self, n: usize) -> usize {
n.div_ceil(self.values_per_byte())
}
}
pub fn unpack_row(packed: &[u8], n: usize, bits: GemmaQuantBits) -> Vec<i32> {
let mut out = Vec::with_capacity(n);
match bits {
GemmaQuantBits::Eight => {
for &b in packed.iter().take(n) {
out.push(b as i8 as i32);
}
}
_ => {
let vpb = bits.values_per_byte();
let shift = bits.field_bits();
let mask = (1i32 << shift) - 1;
let zp = bits.zero_point();
'outer: for &byte in packed {
let mut acc = byte as i32;
for _ in 0..vpb {
if out.len() == n {
break 'outer;
}
out.push((acc & mask) - zp);
acc >>= shift;
}
}
}
}
out
}
pub fn dequantize_matrix(
packed: &[u8],
scale: &[f32],
out: usize,
inn: usize,
bits: GemmaQuantBits,
) -> anyhow::Result<Vec<f32>> {
if scale.len() != out {
anyhow::bail!(
"gemma-qat: weight_scale len {} != out rows {out}",
scale.len()
);
}
let bytes_per_row = bits.packed_len(inn);
let expected = bytes_per_row * out;
if packed.len() != expected {
anyhow::bail!(
"gemma-qat: packed len {} != out*bytes_per_row ({out}*{bytes_per_row}={expected}) \
for {bits:?} inn={inn}",
packed.len()
);
}
let mut w = vec![0f32; out * inn];
for o in 0..out {
let row = &packed[o * bytes_per_row..(o + 1) * bytes_per_row];
let q = unpack_row(row, inn, bits);
let s = scale[o];
let dst = &mut w[o * inn..(o + 1) * inn];
for (d, qv) in dst.iter_mut().zip(q.iter()) {
*d = *qv as f32 * s;
}
}
Ok(w)
}
#[derive(Debug, Clone)]
enum QuantPat {
Exact(String),
EndsWith(String),
LayerMlpUpTo14,
AnyLayerMlp,
AnyLayerSelfAttn,
AudioNotLconvStart,
AudioLconvStart,
Contains(String),
}
impl QuantPat {
fn parse(pat: &str) -> Self {
let unesc = |s: &str| s.replace("\\.", ".").replace("\\d", "");
match pat {
"^lm_head$" => QuantPat::Exact("lm_head".into()),
r"language_model\.embed_tokens$" => {
QuantPat::EndsWith("language_model.embed_tokens".into())
}
r"language_model\.embed_tokens_per_layer$" => {
QuantPat::EndsWith("language_model.embed_tokens_per_layer".into())
}
r"language_model\.layers\.(\d|1[0-4])\.mlp\." => QuantPat::LayerMlpUpTo14,
r"language_model\.layers\.\d+\.mlp\." => QuantPat::AnyLayerMlp,
r"language_model\.layers\.\d+\.self_attn\." => QuantPat::AnyLayerSelfAttn,
r"language_model\.layers\.\d+\.per_layer_input_gate$" => {
QuantPat::EndsWith("per_layer_input_gate".into())
}
r"language_model\.layers\.\d+\.per_layer_projection$" => {
QuantPat::EndsWith("per_layer_projection".into())
}
r"audio_tower(?!.*lconv1d\.linear_start)" => QuantPat::AudioNotLconvStart,
r"audio_tower\.layers\.\d+\.lconv1d\.linear_start\." => QuantPat::AudioLconvStart,
other if other.starts_with('^') && other.ends_with('$') => {
QuantPat::Exact(unesc(&other[1..other.len() - 1]))
}
other if other.ends_with('$') => QuantPat::EndsWith(unesc(&other[..other.len() - 1])),
other => QuantPat::Contains(unesc(other)),
}
}
fn rank(&self) -> u8 {
match self {
QuantPat::Exact(_) => 0,
QuantPat::EndsWith(_) => 1,
QuantPat::LayerMlpUpTo14 => 2,
QuantPat::AudioLconvStart => 2,
QuantPat::AnyLayerSelfAttn => 3,
QuantPat::AnyLayerMlp => 4,
QuantPat::AudioNotLconvStart => 5,
QuantPat::Contains(_) => 6,
}
}
fn matches(&self, name: &str) -> bool {
let layer_idx = |seg: &str| -> Option<usize> {
let i = name.find(seg)? + seg.len();
let rest = &name[i..];
let end = rest
.find(|c: char| !c.is_ascii_digit())
.unwrap_or(rest.len());
rest[..end].parse().ok()
};
match self {
QuantPat::Exact(s) => name == s,
QuantPat::EndsWith(s) => name.ends_with(s),
QuantPat::Contains(s) => name.contains(s),
QuantPat::LayerMlpUpTo14 => {
name.contains(".mlp.")
&& layer_idx("language_model.layers.").is_some_and(|l| l <= 14)
}
QuantPat::AnyLayerMlp => {
name.contains(".mlp.") && name.contains("language_model.layers.")
}
QuantPat::AnyLayerSelfAttn => {
name.contains(".self_attn.") && name.contains("language_model.layers.")
}
QuantPat::AudioNotLconvStart => {
name.contains("audio_tower") && !name.contains("lconv1d.linear_start")
}
QuantPat::AudioLconvStart => {
name.contains("audio_tower") && name.contains("lconv1d.linear_start")
}
}
}
}
#[derive(Debug, Clone)]
pub struct GemmaQuantPlan {
default_bits: u32,
quantize_embeddings: bool,
not_convert: Vec<String>,
rules: Vec<(QuantPat, u32)>,
}
impl GemmaQuantPlan {
pub fn from_json(quant_cfg: &serde_json::Value) -> Self {
let default_bits = quant_cfg
.get("num_bits")
.and_then(|v| v.as_u64())
.unwrap_or(4) as u32;
let quantize_embeddings = quant_cfg
.get("quantize_embeddings")
.and_then(|v| v.as_bool())
.unwrap_or(false);
let not_convert = quant_cfg
.get("modules_to_not_convert")
.and_then(|v| v.as_array())
.map(|a| {
a.iter()
.filter_map(|s| s.as_str().map(str::to_string))
.collect()
})
.unwrap_or_default();
let mut rules: Vec<(QuantPat, u32)> = quant_cfg
.get("module_quant_configs")
.and_then(|v| v.as_object())
.map(|obj| {
obj.iter()
.filter_map(|(pat, v)| {
v.get("num_bits")
.and_then(|b| b.as_u64())
.map(|b| (QuantPat::parse(pat), b as u32))
})
.collect()
})
.unwrap_or_default();
rules.sort_by_key(|(pat, _)| pat.rank());
Self {
default_bits,
quantize_embeddings,
not_convert,
rules,
}
}
pub fn quantize_embeddings(&self) -> bool {
self.quantize_embeddings
}
pub fn resolve_bits(&self, name: &str) -> Option<GemmaQuantBits> {
if self.not_convert.iter().any(|m| name.contains(m.as_str())) {
return None;
}
let bits = self
.rules
.iter()
.find(|(pat, _)| pat.matches(name))
.map(|(_, b)| *b)
.unwrap_or(self.default_bits);
GemmaQuantBits::from_num_bits(bits)
}
}
#[cfg(test)]
mod tests {
use super::*;
#[test]
fn widths_map_from_num_bits() {
assert_eq!(GemmaQuantBits::from_num_bits(2), Some(GemmaQuantBits::Two));
assert_eq!(GemmaQuantBits::from_num_bits(4), Some(GemmaQuantBits::Four));
assert_eq!(
GemmaQuantBits::from_num_bits(8),
Some(GemmaQuantBits::Eight)
);
assert_eq!(GemmaQuantBits::from_num_bits(3), None);
}
#[test]
fn unpack_4bit_low_nibble_first_signed() {
let bytes = [0x10u8, 0xF8u8];
let v = unpack_row(&bytes, 4, GemmaQuantBits::Four);
assert_eq!(v, vec![-8, -7, 0, 7]);
}
#[test]
fn unpack_2bit_lsb_first_signed() {
let bytes = [0xE4u8];
let v = unpack_row(&bytes, 4, GemmaQuantBits::Two);
assert_eq!(v, vec![-2, -1, 0, 1]);
}
#[test]
fn unpack_8bit_is_native_int8() {
let bytes = [0x00u8, 0x7F, 0x80, 0xFF];
let v = unpack_row(&bytes, 4, GemmaQuantBits::Eight);
assert_eq!(v, vec![0, 127, -128, -1]);
}
#[test]
fn unpack_respects_partial_trailing_field() {
let bytes = [0xF8u8, 0x00u8];
let v = unpack_row(&bytes, 3, GemmaQuantBits::Four);
assert_eq!(v, vec![0, 7, -8]);
}
#[test]
fn dequantize_applies_per_row_scale() {
let packed = [0x10u8, 0xF8, 0x98, 0x21];
let scale = [0.5f32, 2.0];
let w = dequantize_matrix(&packed, &scale, 2, 4, GemmaQuantBits::Four).unwrap();
assert_eq!(
w,
vec![
-4.0, -3.5, 0.0, 3.5, 0.0, 2.0, -14.0, -12.0, ]
);
}
#[test]
fn dequantize_rejects_size_mismatch() {
let packed = [0u8; 2];
let scale = [1.0f32; 2];
assert!(dequantize_matrix(&packed, &scale, 2, 4, GemmaQuantBits::Four).is_err());
}
#[test]
fn packed_len_rounds_up() {
assert_eq!(GemmaQuantBits::Two.packed_len(10), 3); assert_eq!(GemmaQuantBits::Four.packed_len(7), 4); assert_eq!(GemmaQuantBits::Eight.packed_len(5), 5);
}
const E2B_QUANT_CFG: &str = r#"{
"module_quant_configs": {
"^lm_head$": {"num_bits": 2},
"audio_tower(?!.*lconv1d\\.linear_start)": {"num_bits": 2},
"audio_tower\\.layers\\.\\d+\\.lconv1d\\.linear_start\\.": {"num_bits": 4},
"language_model\\.embed_tokens$": {"num_bits": 2},
"language_model\\.embed_tokens_per_layer$": {"num_bits": 4},
"language_model\\.layers\\.(\\d|1[0-4])\\.mlp\\.": {"num_bits": 4},
"language_model\\.layers\\.\\d+\\.mlp\\.": {"num_bits": 2},
"language_model\\.layers\\.\\d+\\.per_layer_input_gate$": {"num_bits": 8},
"language_model\\.layers\\.\\d+\\.per_layer_projection$": {"num_bits": 8},
"language_model\\.layers\\.\\d+\\.self_attn\\.": {"num_bits": 4},
"vision_tower": {"num_bits": 8}
},
"modules_to_not_convert": [
"model.vision_tower.patch_embedder",
"model.audio_tower.subsample_conv_projection",
"model.audio_tower.output_proj",
"relative_k_proj",
"model.embed_audio",
"model.embed_vision",
"per_layer_model_projection"
],
"num_bits": 4,
"quantize_embeddings": true
}"#;
fn plan() -> GemmaQuantPlan {
GemmaQuantPlan::from_json(&serde_json::from_str(E2B_QUANT_CFG).unwrap())
}
#[test]
fn resolve_bits_matches_real_config() {
let p = plan();
let b = |n: &str| p.resolve_bits(n);
use GemmaQuantBits::*;
assert_eq!(
b("model.language_model.layers.0.self_attn.q_proj.weight"),
Some(Four)
);
assert_eq!(
b("model.language_model.layers.34.self_attn.v_proj.weight"),
Some(Four)
);
assert_eq!(
b("model.language_model.layers.0.mlp.gate_proj.weight"),
Some(Four)
);
assert_eq!(
b("model.language_model.layers.14.mlp.down_proj.weight"),
Some(Four)
);
assert_eq!(
b("model.language_model.layers.15.mlp.gate_proj.weight"),
Some(Two)
);
assert_eq!(
b("model.language_model.layers.20.mlp.gate_proj.weight"),
Some(Two)
);
assert_eq!(
b("model.language_model.layers.3.per_layer_input_gate"),
Some(Eight)
);
assert_eq!(
b("model.language_model.layers.3.per_layer_projection"),
Some(Eight)
);
assert_eq!(b("model.language_model.embed_tokens"), Some(Two));
assert_eq!(b("model.language_model.embed_tokens_per_layer"), Some(Four));
assert_eq!(b("lm_head"), Some(Two));
assert_eq!(
b("model.vision_tower.encoder.layers.0.mlp.gate_proj.linear.weight"),
Some(Eight)
);
assert_eq!(
b("model.audio_tower.layers.0.feed_forward1.ffw_layer_1.linear.weight"),
Some(Two)
);
assert_eq!(
b("model.audio_tower.layers.0.lconv1d.linear_start.weight"),
Some(Four)
);
}
#[test]
fn modules_to_not_convert_stay_fp() {
let p = plan();
assert_eq!(
p.resolve_bits("model.language_model.per_layer_model_projection.weight"),
None
);
assert_eq!(
p.resolve_bits("model.vision_tower.patch_embedder.input_proj.weight"),
None
);
assert_eq!(
p.resolve_bits("model.embed_vision.embedding_projection.weight"),
None
);
assert_eq!(p.resolve_bits("model.audio_tower.output_proj.weight"), None);
assert!(p.quantize_embeddings());
}
}