Skip to main content

zsfm_ttm/
convert.rs

1use std::fs::File;
2use std::io::BufWriter;
3use std::path::Path;
4
5use anyhow::Context;
6use indicatif::{ProgressBar, ProgressStyle};
7use safetensors::SafeTensors;
8use safetensors::Dtype as StDtype;
9
10use zsfm_gguf::{GGMLType, GGUFMetaValue, GGUFWriter};
11use zsfm_hub::ModelFiles;
12
13use crate::config::TtmConfig;
14use crate::tensor_map::map_tensor_name;
15
16pub struct ConvertOptions {
17    pub output_dtype: GGMLType,
18}
19
20pub fn convert(
21    model_id: &str,
22    files: &ModelFiles,
23    config: &TtmConfig,
24    opts: &ConvertOptions,
25    output_path: &Path,
26) -> anyhow::Result<()> {
27    let mut writer = GGUFWriter::new();
28    write_metadata(&mut writer, model_id, config);
29
30    let shard_bytes = load_shard_bytes(&files.safetensors_shards)?;
31    let shard_views: Vec<SafeTensors> = shard_bytes
32        .iter()
33        .map(|b| SafeTensors::deserialize(b).context("deserialize shard"))
34        .collect::<anyhow::Result<_>>()?;
35
36    let total_tensors: usize = shard_views.iter().map(|s| s.len()).sum();
37    println!("Found {} tensors across {} shard(s).", total_tensors, shard_views.len());
38
39    let pb = ProgressBar::new(total_tensors as u64);
40    pb.set_style(
41        ProgressStyle::with_template(
42            "{spinner:.green} [{elapsed_precise}] [{bar:40.cyan/blue}] {pos}/{len} {msg}",
43        )
44        .unwrap()
45        .progress_chars("=>-"),
46    );
47
48    let mut mapped = 0usize;
49    let mut skipped: Vec<String> = Vec::new();
50    let mut fallback_count = 0usize;
51
52    for shard in &shard_views {
53        for (hf_name, tensor_view) in shard.tensors() {
54            pb.set_message(hf_name.to_string());
55
56            let gguf_name = match map_tensor_name(&hf_name) {
57                Some(n) => n,
58                None => {
59                    skipped.push(hf_name.to_string());
60                    pb.inc(1);
61                    continue;
62                }
63            };
64
65            let src_dtype = ggml_type_from_st(tensor_view.dtype())
66                .with_context(|| format!("tensor {hf_name}: unsupported dtype {:?}", tensor_view.dtype()))?;
67
68            let raw_data = tensor_view.data();
69            let py_shape = tensor_view.shape();
70            let n_elems: usize = py_shape.iter().product();
71            let innermost = py_shape.last().copied().unwrap_or(1);
72            let _outermost = py_shape.first().copied().unwrap_or(1);
73
74            let (dst_dtype, gguf_shape, tensor_data) =
75                if opts.output_dtype == GGMLType::Q8_0 && (innermost % 32 != 0 || n_elems % 32 != 0) {
76                    fallback_count += 1;
77                    let data = cast_data(raw_data, src_dtype, GGMLType::F32)
78                        .with_context(|| format!("tensor {hf_name}: cast failed"))?;
79                    let gs = py_shape.iter().rev().map(|&d| d as u64).collect();
80                    (GGMLType::F32, gs, data)
81                } else {
82                    let dst = opts.output_dtype;
83                    let data = cast_data(raw_data, src_dtype, dst)
84                        .with_context(|| format!("tensor {hf_name}: cast failed"))?;
85                    let gs = py_shape.iter().rev().map(|&d| d as u64).collect();
86                    (dst, gs, data)
87                };
88
89            writer.add_tensor(gguf_name, gguf_shape, dst_dtype, tensor_data);
90            mapped += 1;
91            pb.inc(1);
92        }
93    }
94
95    pb.finish_with_message("tensors processed");
96
97    if !skipped.is_empty() {
98        eprintln!("\nWarning: {} tensor(s) skipped (unrecognised names):", skipped.len());
99        for name in &skipped {
100            eprintln!("  {name}");
101        }
102    }
103    if fallback_count > 0 {
104        eprintln!("\nNote: {fallback_count} tensor(s) fell back to F32 (too small for Q8_0 blocks).");
105    }
106
107    println!("Writing {mapped} tensors to {} …", output_path.display());
108    let out_file = File::create(output_path)
109        .with_context(|| format!("create output file {}", output_path.display()))?;
110    let mut buf_writer = BufWriter::new(out_file);
111    writer.write_to(&mut buf_writer)?;
112    println!("Done.");
113
114    Ok(())
115}
116
117fn write_metadata(writer: &mut GGUFWriter, model_id: &str, config: &TtmConfig) {
118    writer.add_metadata("general.architecture",    GGUFMetaValue::String("ttm".into()));
119    writer.add_metadata("general.name",            GGUFMetaValue::String(model_id.into()));
120    writer.add_metadata("ttm.context_length",      GGUFMetaValue::Uint32(config.context_length as u32));
121    writer.add_metadata("ttm.prediction_length",   GGUFMetaValue::Uint32(config.prediction_length as u32));
122    writer.add_metadata("ttm.patch_length",        GGUFMetaValue::Uint32(config.patch_length as u32));
123    writer.add_metadata("ttm.patch_stride",        GGUFMetaValue::Uint32(config.patch_stride as u32));
124    writer.add_metadata("ttm.d_model",             GGUFMetaValue::Uint32(config.d_model as u32));
125    writer.add_metadata("ttm.num_layers",          GGUFMetaValue::Uint32(config.num_layers as u32));
126    writer.add_metadata("ttm.decoder_d_model",     GGUFMetaValue::Uint32(config.decoder_d_model as u32));
127    writer.add_metadata("ttm.decoder_num_layers",  GGUFMetaValue::Uint32(config.decoder_num_layers as u32));
128    writer.add_metadata("ttm.expansion_factor",    GGUFMetaValue::Uint32(config.expansion_factor as u32));
129    writer.add_metadata("ttm.adaptive_patching_levels", GGUFMetaValue::Uint32(config.adaptive_patching_levels as u32));
130    writer.add_metadata("ttm.num_patches",         GGUFMetaValue::Uint32(config.num_patches as u32));
131    writer.add_metadata("ttm.scaling",             GGUFMetaValue::String(config.scaling.clone()));
132    writer.add_metadata("ttm.gated_attn",          GGUFMetaValue::Bool(config.gated_attn));
133    writer.add_metadata("ttm.norm_eps",            GGUFMetaValue::Float64(config.norm_eps));
134}
135
136fn ggml_type_from_st(dtype: StDtype) -> anyhow::Result<GGMLType> {
137    match dtype {
138        StDtype::F32  => Ok(GGMLType::F32),
139        StDtype::F16  => Ok(GGMLType::F16),
140        StDtype::BF16 => Ok(GGMLType::BF16),
141        other => anyhow::bail!("unsupported safetensors dtype: {other:?}"),
142    }
143}
144
145fn load_shard_bytes(shards: &[std::path::PathBuf]) -> anyhow::Result<Vec<Vec<u8>>> {
146    shards
147        .iter()
148        .map(|p| std::fs::read(p).with_context(|| format!("read shard {}", p.display())))
149        .collect()
150}
151
152fn cast_data(data: &[u8], src: GGMLType, dst: GGMLType) -> anyhow::Result<Vec<u8>> {
153    if src == dst {
154        return Ok(data.to_vec());
155    }
156    if dst == GGMLType::Q8_0 {
157        let f32_values = decode_to_f32(data, src)?;
158        return quantize_q8_0(&f32_values);
159    }
160    match (src, dst) {
161        (GGMLType::F32, GGMLType::F16) => {
162            let f32_values = parse_f32_le(data)?;
163            let mut out = Vec::with_capacity(f32_values.len() * 2);
164            for v in f32_values {
165                let bits = f32_to_f16_bits(v);
166                out.extend_from_slice(&bits.to_le_bytes());
167            }
168            Ok(out)
169        }
170        (GGMLType::F32, GGMLType::BF16) => {
171            let f32_values = parse_f32_le(data)?;
172            let mut out = Vec::with_capacity(f32_values.len() * 2);
173            for v in f32_values {
174                let bits = (v.to_bits() >> 16) as u16;
175                out.extend_from_slice(&bits.to_le_bytes());
176            }
177            Ok(out)
178        }
179        (GGMLType::F16, GGMLType::BF16) => {
180            let mut out = Vec::with_capacity(data.len());
181            for chunk in data.chunks_exact(2) {
182                let f16_bits = u16::from_le_bytes([chunk[0], chunk[1]]);
183                let f32_val = f16_to_f32(f16_bits);
184                let bits = (f32_val.to_bits() >> 16) as u16;
185                out.extend_from_slice(&bits.to_le_bytes());
186            }
187            Ok(out)
188        }
189        (GGMLType::BF16, GGMLType::F32) => {
190            let mut out = Vec::with_capacity(data.len() * 2);
191            for chunk in data.chunks_exact(2) {
192                let bf16_bits = u16::from_le_bytes([chunk[0], chunk[1]]);
193                let f32_bits = (bf16_bits as u32) << 16;
194                out.extend_from_slice(&f32_bits.to_le_bytes());
195            }
196            Ok(out)
197        }
198        (GGMLType::BF16, GGMLType::F16) => {
199            let mut out = Vec::with_capacity(data.len());
200            for chunk in data.chunks_exact(2) {
201                let bf16_bits = u16::from_le_bytes([chunk[0], chunk[1]]);
202                let f32_bits = (bf16_bits as u32) << 16;
203                let f32_val = f32::from_bits(f32_bits);
204                let f16_bits = f32_to_f16_bits(f32_val);
205                out.extend_from_slice(&f16_bits.to_le_bytes());
206            }
207            Ok(out)
208        }
209        (GGMLType::F16, GGMLType::F32) => {
210            let mut out = Vec::with_capacity(data.len() * 2);
211            for chunk in data.chunks_exact(2) {
212                let f16_bits = u16::from_le_bytes([chunk[0], chunk[1]]);
213                let f32_val = f16_to_f32(f16_bits);
214                out.extend_from_slice(&f32_val.to_bits().to_le_bytes());
215            }
216            Ok(out)
217        }
218        _ => anyhow::bail!("unsupported cast: {src:?} → {dst:?}"),
219    }
220}
221
222fn decode_to_f32(data: &[u8], src: GGMLType) -> anyhow::Result<Vec<f32>> {
223    match src {
224        GGMLType::F32  => parse_f32_le(data),
225        GGMLType::F16  => data
226            .chunks_exact(2)
227            .map(|c| Ok(f16_to_f32(u16::from_le_bytes([c[0], c[1]]))))
228            .collect(),
229        GGMLType::BF16 => data
230            .chunks_exact(2)
231            .map(|c| {
232                let bf16_bits = u16::from_le_bytes([c[0], c[1]]);
233                Ok(f32::from_bits((bf16_bits as u32) << 16))
234            })
235            .collect(),
236        GGMLType::Q8_0 => anyhow::bail!("Q8_0 → re-quantization not supported as source"),
237    }
238}
239
240fn quantize_q8_0(values: &[f32]) -> anyhow::Result<Vec<u8>> {
241    const BLOCK: usize = 32;
242    if values.len() % BLOCK != 0 {
243        anyhow::bail!("Q8_0 requires element count divisible by {BLOCK}, got {}", values.len());
244    }
245    let n_blocks = values.len() / BLOCK;
246    let mut out = vec![0u8; n_blocks * 34];
247    for b in 0..n_blocks {
248        let blk = &values[b * BLOCK..(b + 1) * BLOCK];
249        let amax = blk.iter().copied().map(f32::abs).fold(0.0f32, f32::max);
250        let d = if amax == 0.0 { 0.0f32 } else { amax / 127.0 };
251        let d_inv = if d == 0.0 { 0.0f32 } else { 1.0 / d };
252        let base = b * 34;
253        let d_f16 = f32_to_f16_bits(d);
254        out[base..base + 2].copy_from_slice(&d_f16.to_le_bytes());
255        for i in 0..BLOCK {
256            let q = (blk[i] * d_inv).round().clamp(-127.0, 127.0) as i8;
257            out[base + 2 + i] = q as u8;
258        }
259    }
260    Ok(out)
261}
262
263fn parse_f32_le(data: &[u8]) -> anyhow::Result<Vec<f32>> {
264    if data.len() % 4 != 0 {
265        anyhow::bail!("f32 data length not divisible by 4");
266    }
267    Ok(data
268        .chunks_exact(4)
269        .map(|c| f32::from_le_bytes([c[0], c[1], c[2], c[3]]))
270        .collect())
271}
272
273fn f32_to_f16_bits(v: f32) -> u16 {
274    let bits = v.to_bits();
275    let sign = ((bits >> 16) & 0x8000) as u16;
276    let exp = ((bits >> 23) & 0xFF) as i32;
277    let mantissa = bits & 0x007F_FFFF;
278    if exp == 0xFF {
279        return sign | 0x7C00 | if mantissa != 0 { 0x0200 } else { 0 };
280    }
281    let new_exp = exp - 127 + 15;
282    if new_exp >= 31 { return sign | 0x7C00; }
283    if new_exp <= 0 {
284        if new_exp < -10 { return sign; }
285        let m = (mantissa | 0x0080_0000) >> (1 - new_exp);
286        return sign | (m >> 13) as u16;
287    }
288    sign | ((new_exp as u16) << 10) | (mantissa >> 13) as u16
289}
290
291fn f16_to_f32(bits: u16) -> f32 {
292    let sign = ((bits & 0x8000) as u32) << 16;
293    let exp = ((bits >> 10) & 0x1F) as i32;
294    let mantissa = (bits & 0x03FF) as u32;
295    let f32_bits = if exp == 0 {
296        if mantissa == 0 { sign }
297        else {
298            let mut m = mantissa;
299            let mut e = 0i32;
300            while m & 0x0400 == 0 { m <<= 1; e += 1; }
301            sign | ((127 - 15 - e + 1) as u32) << 23 | (m & 0x03FF) << 13
302        }
303    } else if exp == 31 {
304        sign | 0x7F80_0000 | (mantissa << 13)
305    } else {
306        sign | ((exp + 127 - 15) as u32) << 23 | (mantissa << 13)
307    };
308    f32::from_bits(f32_bits)
309}