1use crate::Engine;
45use cudarc::driver::CudaSlice;
46use memra_gguf::dequant::bf16_to_f32;
47use memra_gguf::safetensors::StModel;
48use std::path::Path;
49
50pub const SV_HIDDEN: usize = 1536;
51pub const SV_HEADS: usize = 16;
52pub const SV_HEAD_DIM: usize = SV_HIDDEN / SV_HEADS; pub const SV_INTER: usize = 8960; pub const SV_DEPTH: usize = 47;
55pub const SV_PATCH: usize = 14;
56pub const SV_POS_GRID: usize = 52; pub const SV_PATCH_IN: usize = 3 * SV_PATCH * SV_PATCH; pub const SV_IMAGE_SIZE: usize = 728;
60pub const SV_TILE_SIZE: usize = 504;
61pub const SV_GRID_MAIN: usize = SV_IMAGE_SIZE / SV_PATCH; pub const SV_GRID_TILE: usize = SV_TILE_SIZE / SV_PATCH; pub const SV_MAIN_ROWS: usize = 169;
65pub const SV_TILE_ROWS: usize = 81;
66pub const SV_MAX_IMAGE_SIZE: usize = 3024;
68const LN_EPS: f32 = 1e-5;
69const ROPE_THETA: f32 = 10000.0;
70const MEAN: [f32; 3] = [0.481_454_66, 0.457_827_5, 0.408_210_73];
72const STD: [f32; 3] = [0.268_629_54, 0.261_302_58, 0.275_777_11];
73
74struct Lin {
75 w: CudaSlice<f32>,
76 b: Option<CudaSlice<f32>>,
77 in_f: usize,
78 out_f: usize,
79}
80
81struct SBlock {
82 ln1_w: CudaSlice<f32>,
83 ln1_b: CudaSlice<f32>,
84 ln2_w: CudaSlice<f32>,
85 ln2_b: CudaSlice<f32>,
86 ls1: Vec<f32>,
87 ls2: Vec<f32>,
88 qkv: Lin,
89 proj: Lin,
90 fc: Lin,
91 cproj: Lin,
92}
93
94struct Conv3x3s2 {
97 w: CudaSlice<f32>, b: CudaSlice<f32>,
99 c_in: usize,
100 c_out: usize,
101}
102
103pub struct StepVisionTower {
104 patch: Lin, pos: Vec<f32>,
107 ln_pre_w: CudaSlice<f32>,
108 ln_pre_b: CudaSlice<f32>,
109 blocks: Vec<SBlock>,
110 down1: Conv3x3s2, down2: Conv3x3s2, proj: Lin, }
114
115fn read_f32(m: &StModel, name: &str) -> Result<Vec<f32>, Box<dyn std::error::Error>> {
116 let (info, raw) = m
117 .raw(name)
118 .ok_or_else(|| format!("step vision tensor missing: {name}"))?;
119 match info.dtype.as_str() {
120 "BF16" => Ok(raw
121 .chunks_exact(2)
122 .map(|c| bf16_to_f32(u16::from_le_bytes([c[0], c[1]])))
123 .collect()),
124 "F32" => Ok(raw
125 .chunks_exact(4)
126 .map(|c| f32::from_le_bytes(c.try_into().unwrap()))
127 .collect()),
128 other => Err(format!("step vision tensor {name}: unsupported dtype {other}").into()),
129 }
130}
131
132fn load_lin(
133 e: &Engine,
134 m: &StModel,
135 stem: &str,
136 in_f: usize,
137 out_f: usize,
138 bias: bool,
139) -> Result<Lin, Box<dyn std::error::Error>> {
140 let w = read_f32(m, &format!("{stem}.weight"))?;
141 assert_eq!(w.len(), in_f * out_f, "{stem}.weight shape");
142 let b = if bias {
143 let b = read_f32(m, &format!("{stem}.bias"))?;
144 assert_eq!(b.len(), out_f, "{stem}.bias shape");
145 Some(e.htod(&b)?)
146 } else {
147 None
148 };
149 Ok(Lin {
150 w: e.htod(&w)?,
151 b,
152 in_f,
153 out_f,
154 })
155}
156
157fn window_size(long: usize, short: usize) -> usize {
161 if long <= SV_IMAGE_SIZE {
162 if long as f64 / short as f64 > 1.5 {
163 short
164 } else {
165 0
166 }
167 } else if long as f64 / short as f64 > 4.0 {
168 short.min(SV_TILE_SIZE)
169 } else {
170 SV_TILE_SIZE
171 }
172}
173
174fn pad_rule(w: usize, h: usize) -> (usize, usize) {
176 let ratio = w as f64 / h as f64;
177 if w.min(h) < 32 && (ratio > 4.0 || ratio < 0.25) {
178 let s = w.max(h);
179 (s, s)
180 } else {
181 (w, h)
182 }
183}
184
185fn cap_rule(w: usize, h: usize) -> (usize, usize) {
187 if w.max(h) > SV_MAX_IMAGE_SIZE {
188 let s = SV_MAX_IMAGE_SIZE as f64 / w.max(h) as f64;
189 ((w as f64 * s) as usize, (h as f64 * s) as usize)
190 } else {
191 (w, h)
192 }
193}
194
195fn crop_snap(side: usize, win: usize) -> usize {
197 let ratio = side as f64 / win as f64;
198 if ratio < 1.0 {
199 return side;
200 }
201 let whole = side / win;
202 let n = if ratio - whole as f64 > 0.2 {
203 whole + 1
204 } else {
205 whole
206 };
207 win * n
208}
209
210pub struct StepImagePlan {
215 pub n_tiles: usize,
216 pub newline_mask: Vec<bool>,
217}
218
219impl StepImagePlan {
220 pub fn n_rows(&self) -> usize {
222 self.n_tiles * SV_TILE_ROWS + SV_MAIN_ROWS
223 }
224 pub fn n_prompt_tokens(&self) -> usize {
226 let newlines = self.newline_mask.iter().filter(|&&b| b).count();
227 self.n_tiles * (SV_TILE_ROWS + 2) + newlines + SV_MAIN_ROWS + 2
228 }
229}
230
231fn plan_for_dims(w0: usize, h0: usize) -> StepImagePlan {
232 let (w, h) = pad_rule(w0, h0);
233 let (w, h) = cap_rule(w, h);
234 let win = window_size(w.max(h), w.min(h));
235 if win == 0 {
236 return StepImagePlan {
237 n_tiles: 0,
238 newline_mask: Vec::new(),
239 };
240 }
241 let (cw, ch) = (crop_snap(w, win), crop_snap(h, win));
242 let x_num = (cw / win).max(1);
245 let y_num = (ch / win).max(1);
246 let n = x_num * y_num;
247 let mut mask = vec![false; n];
248 let mut newlines: Vec<usize> = (0..n).filter(|i| (i + 1) % x_num == 0).collect();
249 if newlines.last() == Some(&(n - 1)) {
250 newlines.pop(); }
252 for i in newlines {
253 mask[i] = true;
254 }
255 StepImagePlan {
256 n_tiles: n,
257 newline_mask: mask,
258 }
259}
260
261pub fn step_plan_image(bytes: &[u8]) -> Result<StepImagePlan, String> {
264 let (w, h) = crate::vision_pre::image_header_dims(bytes)?;
265 if w.saturating_mul(h) > crate::vision_pre::IMG_MAX_DECODE_PIXELS {
266 return Err(format!(
267 "image {w}x{h} exceeds the decode budget ({} px) — refused before decode",
268 crate::vision_pre::IMG_MAX_DECODE_PIXELS
269 ));
270 }
271 if w < 2 || h < 2 {
272 return Err(format!("image too small: {w}x{h}"));
273 }
274 Ok(plan_for_dims(w, h))
275}
276
277pub struct StepVisionUnit {
281 pub main: Vec<f32>,
283 pub tiles: Vec<Vec<f32>>,
285 pub newline_mask: Vec<bool>,
286}
287
288impl StepVisionUnit {
289 pub fn n_rows(&self) -> usize {
290 self.tiles.len() * SV_TILE_ROWS + SV_MAIN_ROWS
291 }
292}
293
294fn patchify(img: &image::RgbImage, g: usize) -> Vec<f32> {
296 let mut rows = vec![0f32; g * g * SV_PATCH_IN];
297 for py in 0..g {
298 for px in 0..g {
299 let dst = &mut rows[(py * g + px) * SV_PATCH_IN..(py * g + px + 1) * SV_PATCH_IN];
300 for c in 0..3 {
301 for ky in 0..SV_PATCH {
302 for kx in 0..SV_PATCH {
303 let p =
304 img.get_pixel((px * SV_PATCH + kx) as u32, (py * SV_PATCH + ky) as u32);
305 dst[(c * SV_PATCH + ky) * SV_PATCH + kx] =
306 ((p[c] as f32) / 255.0 - MEAN[c]) / STD[c];
307 }
308 }
309 }
310 }
311 }
312 rows
313}
314
315pub fn step_prep_image(bytes: &[u8]) -> Result<StepVisionUnit, Box<dyn std::error::Error>> {
320 step_plan_image(bytes)?;
321 let (hw, hh) = crate::vision_pre::image_header_dims(bytes)?;
322 let mut reader = image::ImageReader::new(std::io::Cursor::new(bytes)).with_guessed_format()?;
323 let mut limits = image::Limits::default();
324 limits.max_image_width = Some(hw as u32);
325 limits.max_image_height = Some(hh as u32);
326 reader.limits(limits);
327 let mut img = reader.decode()?.to_rgb8();
328 let (w0, h0) = (img.width() as usize, img.height() as usize);
329 let (pw, ph) = pad_rule(w0, h0);
331 if (pw, ph) != (w0, h0) {
332 let mut padded = image::RgbImage::new(pw as u32, ph as u32);
333 image::imageops::replace(&mut padded, &img, 0, 0);
334 img = padded;
335 }
336 let (cw, ch) = cap_rule(img.width() as usize, img.height() as usize);
338 if (cw, ch) != (img.width() as usize, img.height() as usize) {
339 img = image::imageops::resize(
340 &img,
341 cw as u32,
342 ch as u32,
343 image::imageops::FilterType::Triangle,
344 );
345 }
346 let (w, h) = (img.width() as usize, img.height() as usize);
347 let main_img = image::imageops::resize(
349 &img,
350 SV_IMAGE_SIZE as u32,
351 SV_IMAGE_SIZE as u32,
352 image::imageops::FilterType::Triangle,
353 );
354 let main = patchify(&main_img, SV_GRID_MAIN);
355 let win = window_size(w.max(h), w.min(h));
357 let (mut tiles, mut newline_mask) = (Vec::new(), Vec::new());
358 if win > 0 {
359 let (sw, sh) = (crop_snap(w, win), crop_snap(h, win));
360 let snapped = if (sw, sh) != (w, h) {
361 image::imageops::resize(
362 &img,
363 sw as u32,
364 sh as u32,
365 image::imageops::FilterType::Triangle,
366 )
367 } else {
368 img
369 };
370 let x_num = (sw / win).max(1);
371 let y_num = (sh / win).max(1);
372 let n = x_num * y_num;
373 for ty in 0..y_num {
374 for tx in 0..x_num {
375 let crop = image::imageops::crop_imm(
376 &snapped,
377 (tx * win) as u32,
378 (ty * win) as u32,
379 win as u32,
380 win as u32,
381 )
382 .to_image();
383 let tile = image::imageops::resize(
384 &crop,
385 SV_TILE_SIZE as u32,
386 SV_TILE_SIZE as u32,
387 image::imageops::FilterType::Triangle,
388 );
389 tiles.push(patchify(&tile, SV_GRID_TILE));
390 }
391 }
392 let mut newlines: Vec<usize> = (0..n).filter(|i| (i + 1) % x_num == 0).collect();
393 if newlines.last() == Some(&(n - 1)) {
394 newlines.pop();
395 }
396 newline_mask = vec![false; n];
397 for i in newlines {
398 newline_mask[i] = true;
399 }
400 }
401 Ok(StepVisionUnit {
402 main,
403 tiles,
404 newline_mask,
405 })
406}
407
408impl StepVisionTower {
411 pub fn load(e: &Engine, dir: &Path) -> Result<Self, Box<dyn std::error::Error>> {
416 let m = StModel::open(dir)?;
417 let p = "model.vision_model";
418 let patch = {
419 let w = read_f32(&m, &format!("{p}.conv1.weight"))?;
422 assert_eq!(w.len(), SV_HIDDEN * SV_PATCH_IN, "conv1.weight shape");
423 Lin {
424 w: e.htod(&w)?,
425 b: None,
426 in_f: SV_PATCH_IN,
427 out_f: SV_HIDDEN,
428 }
429 };
430 let pos = read_f32(&m, &format!("{p}.positional_embedding"))?;
431 assert_eq!(
432 pos.len(),
433 SV_POS_GRID * SV_POS_GRID * SV_HIDDEN,
434 "positional_embedding shape"
435 );
436 let ln_pre_w = e.htod(&read_f32(&m, &format!("{p}.ln_pre.weight"))?)?;
437 let ln_pre_b = e.htod(&read_f32(&m, &format!("{p}.ln_pre.bias"))?)?;
438 let mut blocks = Vec::with_capacity(SV_DEPTH);
439 for il in 0..SV_DEPTH {
440 let bp = format!("{p}.transformer.resblocks.{il}");
441 let ls1 = read_f32(&m, &format!("{bp}.ls_1.gamma"))?;
442 let ls2 = read_f32(&m, &format!("{bp}.ls_2.gamma"))?;
443 assert_eq!(ls1.len(), SV_HIDDEN, "ls_1.gamma shape");
444 assert_eq!(ls2.len(), SV_HIDDEN, "ls_2.gamma shape");
445 blocks.push(SBlock {
446 ln1_w: e.htod(&read_f32(&m, &format!("{bp}.ln_1.weight"))?)?,
447 ln1_b: e.htod(&read_f32(&m, &format!("{bp}.ln_1.bias"))?)?,
448 ln2_w: e.htod(&read_f32(&m, &format!("{bp}.ln_2.weight"))?)?,
449 ln2_b: e.htod(&read_f32(&m, &format!("{bp}.ln_2.bias"))?)?,
450 ls1,
451 ls2,
452 qkv: {
453 let w = read_f32(&m, &format!("{bp}.attn.in_proj_weight"))?;
455 let b = read_f32(&m, &format!("{bp}.attn.in_proj_bias"))?;
456 assert_eq!(w.len(), 3 * SV_HIDDEN * SV_HIDDEN, "in_proj_weight shape");
457 assert_eq!(b.len(), 3 * SV_HIDDEN, "in_proj_bias shape");
458 Lin {
459 w: e.htod(&w)?,
460 b: Some(e.htod(&b)?),
461 in_f: SV_HIDDEN,
462 out_f: 3 * SV_HIDDEN,
463 }
464 },
465 proj: load_lin(
466 e,
467 &m,
468 &format!("{bp}.attn.out_proj"),
469 SV_HIDDEN,
470 SV_HIDDEN,
471 true,
472 )?,
473 fc: load_lin(e, &m, &format!("{bp}.mlp.c_fc"), SV_HIDDEN, SV_INTER, true)?,
474 cproj: load_lin(
475 e,
476 &m,
477 &format!("{bp}.mlp.c_proj"),
478 SV_INTER,
479 SV_HIDDEN,
480 true,
481 )?,
482 });
483 }
484 let load_conv = |stem: &str,
485 c_in: usize,
486 c_out: usize|
487 -> Result<Conv3x3s2, Box<dyn std::error::Error>> {
488 let w = read_f32(&m, &format!("{stem}.weight"))?;
489 let b = read_f32(&m, &format!("{stem}.bias"))?;
490 assert_eq!(w.len(), c_out * c_in * 9, "{stem}.weight shape");
491 assert_eq!(b.len(), c_out, "{stem}.bias shape");
492 Ok(Conv3x3s2 {
493 w: e.htod(&w)?,
494 b: e.htod(&b)?,
495 c_in,
496 c_out,
497 })
498 };
499 let down1 = load_conv(&format!("{p}.vit_downsampler1"), SV_HIDDEN, 2 * SV_HIDDEN)?;
500 let down2 = load_conv(
501 &format!("{p}.vit_downsampler2"),
502 2 * SV_HIDDEN,
503 4 * SV_HIDDEN,
504 )?;
505 let proj = {
506 let w = read_f32(&m, "model.vit_large_projector.weight")?;
510 assert_eq!(w.len() % (4 * SV_HIDDEN), 0, "vit_large_projector shape");
511 let out_f = w.len() / (4 * SV_HIDDEN);
512 Lin {
513 w: e.htod(&w)?,
514 b: None,
515 in_f: 4 * SV_HIDDEN,
516 out_f,
517 }
518 };
519 eprintln!(
520 "[step-vision] tower loaded from {} ({SV_DEPTH} blocks, out_width {}, f32-resident)",
521 dir.display(),
522 proj.out_f
523 );
524 Ok(Self {
525 patch,
526 pos,
527 ln_pre_w,
528 ln_pre_b,
529 blocks,
530 down1,
531 down2,
532 proj,
533 })
534 }
535
536 pub fn out_width(&self) -> usize {
539 self.proj.out_f
540 }
541
542 fn linear_bias(
543 &self,
544 e: &Engine,
545 x: &CudaSlice<f32>,
546 l: &Lin,
547 m: usize,
548 ) -> Result<CudaSlice<f32>, Box<dyn std::error::Error>> {
549 let mut y = e.linear(x, &l.w, m, l.in_f, l.out_f)?;
550 if let Some(b) = &l.b {
551 for r in 0..m {
552 e.add_row_inplace(&mut y, b, l.out_f, r * l.out_f)?;
553 }
554 }
555 Ok(y)
556 }
557
558 fn pos_for_grid(&self, g: usize) -> Vec<f32> {
562 if g == SV_POS_GRID {
563 return self.pos.clone();
564 }
565 let scale = SV_POS_GRID as f32 / g as f32;
566 let mut out = vec![0f32; g * g * SV_HIDDEN];
567 for y in 0..g {
568 for x in 0..g {
569 let sy = ((y as f32 + 0.5) * scale - 0.5).clamp(0.0, (SV_POS_GRID - 1) as f32);
570 let sx = ((x as f32 + 0.5) * scale - 0.5).clamp(0.0, (SV_POS_GRID - 1) as f32);
571 let (y0, x0) = (sy.floor() as usize, sx.floor() as usize);
572 let (y1, x1) = ((y0 + 1).min(SV_POS_GRID - 1), (x0 + 1).min(SV_POS_GRID - 1));
573 let (fy, fx) = (sy - y0 as f32, sx - x0 as f32);
574 let dst = &mut out[(y * g + x) * SV_HIDDEN..(y * g + x + 1) * SV_HIDDEN];
575 for c in 0..SV_HIDDEN {
576 let p00 = self.pos[(y0 * SV_POS_GRID + x0) * SV_HIDDEN + c];
577 let p01 = self.pos[(y0 * SV_POS_GRID + x1) * SV_HIDDEN + c];
578 let p10 = self.pos[(y1 * SV_POS_GRID + x0) * SV_HIDDEN + c];
579 let p11 = self.pos[(y1 * SV_POS_GRID + x1) * SV_HIDDEN + c];
580 dst[c] = p00 * (1.0 - fy) * (1.0 - fx)
581 + p01 * (1.0 - fy) * fx
582 + p10 * fy * (1.0 - fx)
583 + p11 * fy * fx;
584 }
585 }
586 }
587 out
588 }
589
590 fn im2col(x: &[f32], g: usize, c_in: usize) -> (Vec<f32>, usize) {
594 let og = (g - 1) / 2 + 1;
595 let mut out = vec![0f32; og * og * c_in * 9];
596 for oy in 0..og {
597 for ox in 0..og {
598 let dst = &mut out[(oy * og + ox) * c_in * 9..(oy * og + ox + 1) * c_in * 9];
599 for ky in 0..3usize {
600 for kx in 0..3usize {
601 let iy = (2 * oy + ky) as isize - 1;
602 let ix = (2 * ox + kx) as isize - 1;
603 if iy < 0 || ix < 0 || iy >= g as isize || ix >= g as isize {
604 continue; }
606 let src = &x[((iy as usize) * g + ix as usize) * c_in..];
607 for c in 0..c_in {
608 dst[c * 9 + ky * 3 + kx] = src[c];
609 }
610 }
611 }
612 }
613 }
614 (out, og)
615 }
616
617 pub fn forward(
621 &self,
622 e: &Engine,
623 patches: &[f32],
624 g: usize,
625 ) -> Result<CudaSlice<f32>, Box<dyn std::error::Error>> {
626 let n = g * g;
627 assert_eq!(patches.len(), n * SV_PATCH_IN, "patch buffer shape");
628 if n > 12288 {
629 return Err(format!(
630 "step vision segment {n} patches exceeds the sdpa shared-memory ceiling (12288)"
631 )
632 .into());
633 }
634 let dbg = std::env::var("MEMRA_VISION_DEBUG").ok();
635 let dump = |tag: &str, buf: &[f32]| {
636 if let Some(dir) = dbg.as_deref() {
637 let raw: Vec<u8> = buf.iter().flat_map(|v| v.to_le_bytes()).collect();
638 let _ = std::fs::write(format!("{dir}/rust_{tag}.bin"), raw);
639 }
640 };
641 let xd = e.htod(patches)?;
643 let embedded = self.linear_bias(e, &xd, &self.patch, n)?;
644 let pos = self.pos_for_grid(g);
645 let pos_d = e.htod(&pos)?;
646 let mut summed = e.zeros(n * SV_HIDDEN)?;
647 e.add(&embedded, &pos_d, &mut summed, n * SV_HIDDEN)?;
648 let mut x = e.zeros(n * SV_HIDDEN)?;
649 e.layer_norm_bias(
650 &summed,
651 &self.ln_pre_w,
652 &self.ln_pre_b,
653 &mut x,
654 SV_HIDDEN,
655 n,
656 LN_EPS,
657 )?;
658 if dbg.is_some() {
659 dump("pre_blocks", &e.dtoh(&x)?);
660 }
661 let half = SV_HEAD_DIM / 2; let quarter = half / 2; let inv_freq: Vec<f32> = (0..quarter)
668 .map(|i| ROPE_THETA.powf(-2.0 * (i as f32) / half as f32))
669 .collect();
670 let mut cos_t = vec![0f32; n * half]; let mut sin_t = vec![0f32; n * half];
672 for t in 0..n {
673 let (row, col) = (t / g, t % g);
674 for i in 0..quarter {
675 let ac = col as f32 * inv_freq[i];
676 let ar = row as f32 * inv_freq[i];
677 cos_t[t * half + i] = ac.cos();
678 sin_t[t * half + i] = ac.sin();
679 cos_t[t * half + quarter + i] = ar.cos();
680 sin_t[t * half + quarter + i] = ar.sin();
681 }
682 }
683 let scale = 1.0 / (SV_HEAD_DIM as f32).sqrt();
684 for (ib, blk) in self.blocks.iter().enumerate() {
685 let mut h = e.zeros(n * SV_HIDDEN)?;
688 e.layer_norm_bias(&x, &blk.ln1_w, &blk.ln1_b, &mut h, SV_HIDDEN, n, LN_EPS)?;
689 let qkv = self.linear_bias(e, &h, &blk.qkv, n)?;
690 let qkv_h = e.dtoh(&qkv)?;
691 let mut qh = vec![0f32; n * SV_HIDDEN];
692 let mut kh = vec![0f32; n * SV_HIDDEN];
693 let mut vh = vec![0f32; n * SV_HIDDEN];
694 for t in 0..n {
695 let row = &qkv_h[t * 3 * SV_HIDDEN..(t + 1) * 3 * SV_HIDDEN];
696 let dst = t * SV_HIDDEN;
697 vh[dst..dst + SV_HIDDEN].copy_from_slice(&row[2 * SV_HIDDEN..3 * SV_HIDDEN]);
698 for hd in 0..SV_HEADS {
699 let o = hd * SV_HEAD_DIM;
700 for hf in 0..2usize {
704 for i in 0..quarter {
705 let (c, s) = (
706 cos_t[t * half + hf * quarter + i],
707 sin_t[t * half + hf * quarter + i],
708 );
709 let d = hf * half + 2 * i;
710 let (qa, qb) = (row[o + d], row[o + d + 1]);
711 qh[dst + o + d] = qa * c - qb * s;
712 qh[dst + o + d + 1] = qb * c + qa * s;
713 let (ka, kb) = (row[SV_HIDDEN + o + d], row[SV_HIDDEN + o + d + 1]);
714 kh[dst + o + d] = ka * c - kb * s;
715 kh[dst + o + d + 1] = kb * c + ka * s;
716 }
717 }
718 }
719 }
720 let (qd, kd, vd) = (e.htod(&qh)?, e.htod(&kh)?, e.htod(&vh)?);
721 let mut od = e.zeros(n * SV_HIDDEN)?;
722 e.sdpa_naive(
723 &qd,
724 &kd,
725 &vd,
726 &mut od,
727 SV_HEAD_DIM,
728 SV_HEADS,
729 SV_HEADS,
730 n,
731 n,
732 scale,
733 false,
734 )?;
735 let attn = self.linear_bias(e, &od, &blk.proj, n)?;
736 let mut ah = e.dtoh(&attn)?;
738 for t in 0..n {
739 for c in 0..SV_HIDDEN {
740 ah[t * SV_HIDDEN + c] *= blk.ls1[c];
741 }
742 }
743 let ad = e.htod(&ah)?;
744 let mut xr = e.zeros(n * SV_HIDDEN)?;
745 e.add(&x, &ad, &mut xr, n * SV_HIDDEN)?;
746 let mut h2 = e.zeros(n * SV_HIDDEN)?;
748 e.layer_norm_bias(&xr, &blk.ln2_w, &blk.ln2_b, &mut h2, SV_HIDDEN, n, LN_EPS)?;
749 let f1 = self.linear_bias(e, &h2, &blk.fc, n)?;
750 let mut fh = e.dtoh(&f1)?;
751 for v in fh.iter_mut() {
752 *v = *v / (1.0 + (-1.702 * *v).exp());
754 }
755 let fd = e.htod(&fh)?;
756 let f2 = self.linear_bias(e, &fd, &blk.cproj, n)?;
757 let mut mh = e.dtoh(&f2)?;
758 for t in 0..n {
759 for c in 0..SV_HIDDEN {
760 mh[t * SV_HIDDEN + c] *= blk.ls2[c];
761 }
762 }
763 let md = e.htod(&mh)?;
764 let mut xn = e.zeros(n * SV_HIDDEN)?;
765 e.add(&xr, &md, &mut xn, n * SV_HIDDEN)?;
766 x = xn;
767 if dbg.is_some() && ib == 0 {
768 dump("blk0", &e.dtoh(&x)?);
769 }
770 }
771 if dbg.is_some() {
773 dump("post_blocks", &e.dtoh(&x)?);
774 }
775 let xh = e.dtoh(&x)?;
778 let (col1, g1) = Self::im2col(&xh, g, SV_HIDDEN);
779 let c1 = e.htod(&col1)?;
780 let mut y1 = e.linear(
781 &c1,
782 &self.down1.w,
783 g1 * g1,
784 self.down1.c_in * 9,
785 self.down1.c_out,
786 )?;
787 for r in 0..g1 * g1 {
788 e.add_row_inplace(
789 &mut y1,
790 &self.down1.b,
791 self.down1.c_out,
792 r * self.down1.c_out,
793 )?;
794 }
795 let y1h = e.dtoh(&y1)?;
796 let (col2, g2) = Self::im2col(&y1h, g1, self.down2.c_in);
797 let c2 = e.htod(&col2)?;
798 let mut y2 = e.linear(
799 &c2,
800 &self.down2.w,
801 g2 * g2,
802 self.down2.c_in * 9,
803 self.down2.c_out,
804 )?;
805 for r in 0..g2 * g2 {
806 e.add_row_inplace(
807 &mut y2,
808 &self.down2.b,
809 self.down2.c_out,
810 r * self.down2.c_out,
811 )?;
812 }
813 if dbg.is_some() {
814 dump("downsampled", &e.dtoh(&y2)?);
815 }
816 let out = self.linear_bias(e, &y2, &self.proj, g2 * g2)?;
817 if dbg.is_some() {
818 dump("projected", &e.dtoh(&out)?);
819 }
820 Ok(out)
821 }
822
823 pub fn forward_unit(
826 &self,
827 e: &Engine,
828 unit: &StepVisionUnit,
829 rows: &mut CudaSlice<f32>,
830 row_off: usize,
831 ) -> Result<usize, Box<dyn std::error::Error>> {
832 let w = self.out_width();
833 let mut off = row_off;
834 for tile in &unit.tiles {
835 let emb = self.forward(e, tile, SV_GRID_TILE)?;
836 e.dtod_copy_into(&emb, rows, off * w)?;
837 off += SV_TILE_ROWS;
838 }
839 let emb = self.forward(e, &unit.main, SV_GRID_MAIN)?;
840 e.dtod_copy_into(&emb, rows, off * w)?;
841 off += SV_MAIN_ROWS;
842 Ok(off - row_off)
843 }
844}
845
846#[cfg(test)]
847mod tests {
848 use super::*;
849
850 #[test]
853 fn tiling_plan_cells() {
854 let p = plan_for_dims(600, 400);
856 assert_eq!((p.n_tiles, p.n_prompt_tokens()), (0, 171));
857 let p = plan_for_dims(728, 728);
858 assert_eq!(p.n_tiles, 0);
859 let p = plan_for_dims(700, 300);
863 assert_eq!(p.n_tiles, 3);
864 assert_eq!(p.newline_mask, vec![false, false, false]);
865 let p = plan_for_dims(1600, 900);
869 assert_eq!(p.n_tiles, 6);
870 assert_eq!(
871 p.newline_mask,
872 vec![false, false, true, false, false, false]
873 );
874 assert_eq!(p.n_prompt_tokens(), 670);
876 let p = plan_for_dims(200, 20);
878 assert_eq!(p.n_tiles, 0);
879 let p = plan_for_dims(4000, 1000);
882 assert_eq!(p.n_tiles, 12);
883 assert_eq!(p.n_rows(), 12 * 81 + 169);
884 }
885
886 #[test]
888 fn downsampler_geometry() {
889 let x = vec![0f32; 52 * 52 * 4];
890 let (_, og) = StepVisionTower::im2col(&x, 52, 4);
891 assert_eq!(og, 26);
892 let x = vec![0f32; 26 * 26 * 4];
893 let (_, og) = StepVisionTower::im2col(&x, 26, 4);
894 assert_eq!(og, 13);
895 let x = vec![0f32; 36 * 36 * 4];
896 let (_, og) = StepVisionTower::im2col(&x, 36, 4);
897 assert_eq!(og, 18);
898 let x = vec![0f32; 18 * 18 * 4];
899 let (_, og) = StepVisionTower::im2col(&x, 18, 4);
900 assert_eq!(og, 9);
901 }
902
903 #[test]
905 fn im2col_values() {
906 let x: Vec<f32> = (1..=9).map(|v| v as f32).collect();
909 let (col, og) = StepVisionTower::im2col(&x, 3, 1);
910 assert_eq!(og, 2);
911 assert_eq!(&col[0..9], &[0., 0., 0., 0., 1., 2., 0., 4., 5.]);
912 assert_eq!(&col[27..36], &[5., 6., 0., 8., 9., 0., 0., 0., 0.]);
914 }
915
916 fn png_bytes(w: u32, h: u32) -> Vec<u8> {
917 let img = image::RgbImage::from_fn(w, h, |x, y| {
918 image::Rgb([(x % 251) as u8, (y % 241) as u8, ((x + y) % 253) as u8])
919 });
920 let mut buf = std::io::Cursor::new(Vec::new());
921 img.write_to(&mut buf, image::ImageFormat::Png).unwrap();
922 buf.into_inner()
923 }
924
925 #[test]
928 fn prep_matches_plan() {
929 for (w, h) in [(64u32, 64u32), (1600, 900), (700, 300), (900, 3000)] {
930 let bytes = png_bytes(w, h);
931 let plan = step_plan_image(&bytes).unwrap();
932 let prep = step_prep_image(&bytes).unwrap();
933 assert_eq!(prep.tiles.len(), plan.n_tiles, "{w}x{h} tile count");
934 assert_eq!(prep.newline_mask, plan.newline_mask, "{w}x{h} newline mask");
935 assert_eq!(prep.main.len(), SV_GRID_MAIN * SV_GRID_MAIN * SV_PATCH_IN);
936 for t in &prep.tiles {
937 assert_eq!(t.len(), SV_GRID_TILE * SV_GRID_TILE * SV_PATCH_IN);
938 }
939 assert_eq!(prep.n_rows(), plan.n_rows());
940 }
941 }
942
943 #[test]
946 fn patchify_normalization() {
947 let img = image::RgbImage::from_pixel(
948 SV_IMAGE_SIZE as u32,
949 SV_IMAGE_SIZE as u32,
950 image::Rgb([128, 128, 128]),
951 );
952 let rows = patchify(&img, SV_GRID_MAIN);
953 let want: Vec<f32> = (0..3).map(|c| (128.0 / 255.0 - MEAN[c]) / STD[c]).collect();
954 let r0 = &rows[..SV_PATCH_IN];
955 for c in 0..3 {
956 for i in 0..SV_PATCH * SV_PATCH {
957 assert!((r0[c * SV_PATCH * SV_PATCH + i] - want[c]).abs() < 1e-6);
958 }
959 }
960 }
961}