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}