From 1fe06110ff2b4c422ef26a986b5e74ae353acb25 Mon Sep 17 00:00:00 2001 From: jason Date: Fri, 6 Mar 2026 11:25:44 -0500 Subject: [PATCH] more glm-ocr tweaks --- src/models/glm_ocr/generate.rs | 57 ---- src/models/glm_ocr/model.rs | 512 ++------------------------------ src/models/glm_ocr/processor.rs | 54 +--- 3 files changed, 33 insertions(+), 590 deletions(-) diff --git a/src/models/glm_ocr/generate.rs b/src/models/glm_ocr/generate.rs index 79e7f61..49d79c5 100644 --- a/src/models/glm_ocr/generate.rs +++ b/src/models/glm_ocr/generate.rs @@ -124,35 +124,6 @@ impl<'a> GlmOcrGenerateModel<'a> { dtype }; - #[cfg(debug_assertions)] - { - eprintln!("GLM-OCR Config Debug:"); - eprintln!(" text_config.hidden_size: {}", cfg.text_config.hidden_size); - eprintln!( - " text_config.num_attention_heads: {}", - cfg.text_config.num_attention_heads - ); - eprintln!( - " text_config.num_key_value_heads: {}", - cfg.text_config.num_key_value_heads - ); - eprintln!( - " text_config.head_dim: {}", - cfg.text_config.head_dim.unwrap_or_else(|| { - // Integer division, panics if num_attention_heads is 0 (like Python) - cfg.text_config.hidden_size / cfg.text_config.num_attention_heads - }) - ); - eprintln!( - " text_config.mrope_section: {:?}", - cfg.text_config.mrope_section - ); - eprintln!( - " Calculated head_dim: {}", - cfg.text_config.hidden_size / cfg.text_config.num_attention_heads - ); - } - let processor = GlmOcrProcessor::new(path, &device, dtype)?; let model_list = find_type_files(path, "safetensors")?; let vb = unsafe { VarBuilder::from_mmaped_safetensors(&model_list, dtype, &device)? }; @@ -213,17 +184,6 @@ impl<'a> GenerateModel for GlmOcrGenerateModel<'a> { // Get prompt text from messages let prompt = extract_text_from_messages(&mes).unwrap_or_else(|| "Extract all text from this image.".to_string()); - #[cfg(debug_assertions)] - { - eprintln!("[GLM-OCR] ===== DEBUG START ====="); - eprintln!("[GLM-OCR] Image: {}", image_path); - eprintln!("[GLM-OCR] Prompt: {}", prompt); - eprintln!("[GLM-OCR] Device: {:?}", self.device); - eprintln!("[GLM-OCR] Temperature: {:?}", temperature); - eprintln!("[GLM-OCR] Top_p: {:?}", top_p); - eprintln!("[GLM-OCR] Max tokens: {}", mes.max_tokens.unwrap_or(512)); - } - let processed = self.processor.process_info( image_path, &prompt, @@ -236,16 +196,6 @@ impl<'a> GenerateModel for GlmOcrGenerateModel<'a> { self.spatial_merge_size, )?; - if std::env::var("GLM_DEBUG").is_ok() { - eprintln!("[GLM-OCR] ===== AFTER PROCESS_INFO ====="); - eprintln!("[GLM-OCR] input_ids shape: {:?}", processed.input_ids.shape()); - let input_ids_vec = processed.input_ids.squeeze(0).unwrap().to_vec1::().unwrap(); - eprintln!("[GLM-OCR] input_ids (first 20): {:?}", &input_ids_vec[..20.min(input_ids_vec.len())]); - eprintln!("[GLM-OCR] pixel_values shape: {:?}", processed.pixel_values.shape()); - eprintln!("[GLM-OCR] image_mask shape: {:?}", processed.image_mask.shape()); - eprintln!("[GLM-OCR] grid_thw: {:?}", processed.grid_thw); - } - let mut input_ids = processed.input_ids; let pixel_values = Some(processed.pixel_values); let image_grid_thw = Some(processed.grid_thw); @@ -255,13 +205,6 @@ impl<'a> GenerateModel for GlmOcrGenerateModel<'a> { let mut generate = Vec::new(); let sample_len = mes.max_tokens.unwrap_or(512); - #[cfg(debug_assertions)] - { - eprintln!("[GLM-OCR] ===== START GENERATION ====="); - eprintln!("[GLM-OCR] Initial seq_len: {}", seq_len); - eprintln!("[GLM-OCR] eos_token_ids: {:?}", self.eos_token_ids); - } - for _ in 0..sample_len { let is_first_pass = seqlen_offset == 0; let logits = self.model.forward( diff --git a/src/models/glm_ocr/model.rs b/src/models/glm_ocr/model.rs index 62a9dfd..d4fe3d7 100644 --- a/src/models/glm_ocr/model.rs +++ b/src/models/glm_ocr/model.rs @@ -1,24 +1,4 @@ //! GLM-OCR Model Implementation -//! -//! A multimodal vision-language model for OCR tasks that integrates: -//! - Vision Encoder: Processes images via patch embedding and transformer blocks -//! - Projector: Maps vision features to the language model's embedding space -//! - Language Model: Causal transformer decoder for text generation -//! -//! # Architecture -//! -//! ```text -//! Image → [Patch Embed] → [Vision Transformer] → [Spatial Merge] → [Projector] -//! ↓ -//! Text → [Token Embed] → [Decoder Layers with M-RoPE] ← [Feature Fusion] -//! ↓ -//! [LM Head] → Output Tokens -//! ``` -//! -//! Key features: -//! - M-RoPE (Multimodal Rotary Position Embedding) for unified 1D text and 3D vision positions -//! - Spatial merge to reduce visual token count before feeding to LLM -//! - KV-cache support for efficient autoregressive generation use anyhow::Result; use candle_core::{D, DType, IndexOp, Tensor}; @@ -40,57 +20,6 @@ use crate::{ }, }; -/// Print tensor statistics in the same format as Python's `stats()` helper in compare_intermediate.py. -/// Gated by environment variable `GLM_INTERMEDIATE=1`. -fn tensor_stats(name: &str, t: &Tensor) { - if std::env::var("GLM_INTERMEDIATE").is_err() { - return; - } - let Ok(t_f32) = t.to_dtype(DType::F32) else { return }; - let Ok(flat) = t_f32.flatten_all() else { return }; - let n = flat.elem_count(); - if n == 0 { - eprintln!("[RS] {name}: shape={:?} EMPTY", t.shape()); - return; - } - let Ok(vals) = flat.to_vec1::() else { return }; - let mean = vals.iter().copied().sum::() / n as f32; - let variance = vals.iter().map(|v| (v - mean).powi(2)).sum::() / n as f32; - let std = variance.sqrt(); - let min = vals.iter().copied().fold(f32::INFINITY, f32::min); - let max = vals.iter().copied().fold(f32::NEG_INFINITY, f32::max); - let n_show = 8.min(n); - let first: Vec = vals[..n_show].iter().map(|v| format!("{v:.4}")).collect(); - eprintln!( - "[RS] {name}: shape={:?} mean={mean:.6} std={std:.6} min={min:.6} max={max:.6}", - t.shape() - ); - eprintln!("[RS] {name}: first{n_show}={first:?}"); -} - -// ============================================================================ -// 1. GlmOcrRMSNorm -// ============================================================================ - -// Python: @use_kernel_forward_from_hub("RMSNorm") -// class GlmOcrRMSNorm(nn.Module): -// def __init__(self, hidden_size, eps: float = 1e-6) -> None: -// """ -// GlmOcrRMSNorm is equivalent to T5LayerNorm -// """ -// super().__init__() -// self.weight = nn.Parameter(torch.ones(hidden_size)) -// self.variance_epsilon = eps - -// def forward(self, hidden_states: torch.Tensor) -> torch.Tensor: -// input_dtype = hidden_states.dtype -// hidden_states = hidden_states.to(torch.float32) -// variance = hidden_states.pow(2).mean(-1, keepdim=True) -// hidden_states = hidden_states * torch.rsqrt(variance + self.variance_epsilon) -// return self.weight * hidden_states.to(input_dtype) - -// def extra_repr(self): -// return f"{tuple(self.weight.shape)}, eps={self.variance_epsilon}" pub struct GlmOcrRMSNorm(RmsNorm); impl GlmOcrRMSNorm { @@ -106,19 +35,6 @@ impl GlmOcrRMSNorm { } } -// ============================================================================ -// 2. GlmOcrVisionMlp -// ============================================================================ - -// Python reference (transformers commit 4854dbf9): -// class GlmOcrVisionMlp(nn.Module): -// def __init__(self, config, bias=False): -// self.gate_proj = nn.Linear(config.hidden_size, config.intermediate_size, bias=bias) -// self.up_proj = nn.Linear(config.hidden_size, config.intermediate_size, bias=bias) -// self.down_proj = nn.Linear(config.intermediate_size, config.hidden_size, bias=bias) -// self.act_fn = ACT2FN[config.hidden_act] -// def forward(self, hidden_state): -// return self.down_proj(self.act_fn(self.gate_proj(hidden_state)) * self.up_proj(hidden_state)) pub struct GlmOcrVisionMlp(GateUpDownMLP); impl GlmOcrVisionMlp { @@ -141,12 +57,6 @@ impl GlmOcrVisionMlp { } } -// ============================================================================ -// 3. eager_attention_forward (+ repeat_kv from utils) -// ============================================================================ - -// Python eager_attention_forward implementation -// Reference: py-glm-ocr/glm_ocr/modeling_glm_ocr.py fn eager_attention_forward( query_states: &Tensor, key_states: &Tensor, @@ -156,16 +66,6 @@ fn eager_attention_forward( scaling: f64, dropout: f64, ) -> Result<(Tensor, Tensor)> { - // Attention matrix size info — enable with GLM_DEBUG=1. - if std::env::var("GLM_DEBUG").is_ok() { - let dims = query_states.dims(); - let (heads, q_len, k_len) = (dims[1], dims[2], key_states.dims()[2]); - let attn_mb = (heads * q_len * k_len * 4) as f64 / 1_048_576.0; - eprintln!( - "[GLM-OCR OOM-DBG] eager_attention_forward: heads={heads} q_len={q_len} k_len={k_len} \ - attn_weights={attn_mb:.1}MB (f32)" - ); - } let key_states = match num_key_value_groups { Some(g) => repeat_kv(key_states.clone(), g)?.contiguous()?, None => key_states.clone(), @@ -182,9 +82,10 @@ fn eager_attention_forward( #[cfg(feature = "flash-attn")] { // Flash attention: causal iff attention_mask is present. - let q = query_states.transpose(1, 2)?; - let k = key_states.transpose(1, 2)?; - let v = value_states.transpose(1, 2)?; + // Explicit contiguous() ensures proper memory layout for flash_attn kernel + let q = query_states.transpose(1, 2)?.contiguous()?; + let k = key_states.transpose(1, 2)?.contiguous()?; + let v = value_states.transpose(1, 2)?.contiguous()?; candle_flash_attn::flash_attn(&q, &k, &v, scaling as f32, attention_mask.is_some())? // flash_attn returns [batch, q_len, heads, head_dim] — already in final layout } @@ -194,6 +95,7 @@ fn eager_attention_forward( // [batch, heads, CHUNK, k_len] stays bounded regardless of q_len. // Peak memory per chunk: CHUNK × k_len × heads × 4 bytes (f32 softmax). // Mathematically equivalent to full attention. + // CHUNK_SIZE=512 is empirically optimal for most hardware (CPU/GPU balance) const CHUNK_SIZE: usize = 512; let q_len = query_states.dim(2)?; let k_t = key_states.transpose(D::Minus2, D::Minus1)?.contiguous()?; @@ -211,8 +113,15 @@ fn eager_attention_forward( attn.broadcast_add(&mask.narrow(2, start, len)?.to_dtype(attn.dtype())?)? } }; - let attn = candle_nn::ops::softmax_last_dim(&attn.to_dtype(DType::F32)?)? - .to_dtype(query_states.dtype())?; + // Softmax computation: Optimize dtype conversions for CPU (which uses F32) + let attn = if query_states.dtype() == DType::F32 { + candle_nn::ops::softmax_last_dim(&attn)? + } else { + candle_nn::ops::softmax_last_dim(&attn.to_dtype(DType::F32)?)? + .to_dtype(query_states.dtype())? + }; + // Apply dropout uniformly across chunked and non-chunked paths for consistency + let attn = candle_nn::ops::dropout(&attn, dropout as f32)?; chunks.push(attn.matmul(&value_states)?); start += len; } @@ -223,9 +132,16 @@ fn eager_attention_forward( None => attn, Some(mask) => attn.broadcast_add(&mask.to_dtype(attn.dtype())?)?, }; - let attn = candle_nn::ops::softmax_last_dim(&attn.to_dtype(DType::F32)?)? - .to_dtype(query_states.dtype())?; - candle_nn::ops::dropout(&attn, dropout as f32)?.matmul(&value_states)? + // Softmax computation: Same optimization as chunked path + let attn = if query_states.dtype() == DType::F32 { + candle_nn::ops::softmax_last_dim(&attn)? + } else { + candle_nn::ops::softmax_last_dim(&attn.to_dtype(DType::F32)?)? + .to_dtype(query_states.dtype())? + }; + // Apply dropout uniformly (now consistent across both paths) + let attn = candle_nn::ops::dropout(&attn, dropout as f32)?; + attn.matmul(&value_states)? }; // [batch, heads, q_len, head_dim] -> [batch, q_len, heads, head_dim] raw.transpose(1, 2)?.contiguous()? @@ -237,19 +153,6 @@ fn eager_attention_forward( Ok((output, placeholder)) } -// ============================================================================ -// 5. GlmOcrTextAttention -// ============================================================================ - -// Python: class GlmOcrTextAttention(nn.Module): -// def __init__(self, config): -// self.q_proj = nn.Linear(...) -// self.k_proj = nn.Linear(...) -// self.v_proj = nn.Linear(...) -// self.o_proj = nn.Linear(...) -// def forward(self, hidden_states, attention_mask, position_ids, past_key_value, use_cache): -// # QKV projection, attention computation, output projection -// return attn_output, attn_weights pub struct GlmOcrTextAttention { q_proj: Linear, k_proj: Linear, @@ -338,9 +241,6 @@ impl GlmOcrTextAttention { let (query_states, key_states) = glm_ocr_apply_rotary_pos_emb(&query_states, &key_states, cos, sin)?; - // Python: if past_key_values is not None: - // cache_kwargs = {"sin": sin, "cos": cos, "cache_position": cache_position} - // key_states, value_states = past_key_values.update(key_states, value_states, self.layer_idx, cache_kwargs) let (key_states, value_states) = match &self.kv_cache { None => (key_states, value_states), Some((prev_k, prev_v)) => { @@ -372,16 +272,6 @@ impl GlmOcrTextAttention { } } -// ============================================================================ -// 7. GlmOcrVisionRotaryEmbedding -// ============================================================================ - -// Python: class GlmOcrVisionRotaryEmbedding(nn.Module): -// def __init__(self, dim, theta=10000.0): -// inv_freq = 1.0 / (theta ** (torch.arange(0, dim, 2, dtype=torch.float) / dim)) -// def forward(self, seqlen): -// seq = torch.arange(seqlen, device=self.inv_freq.device) -// return torch.outer(seq, self.inv_freq) # (seqlen, dim/2) pub struct GlmOcrVisionRotaryEmbedding { inv_freq: Tensor, } @@ -405,19 +295,6 @@ impl GlmOcrVisionRotaryEmbedding { Ok(freqs) } - /// Python: GlmOcrVisionModel.rot_pos_emb(self, grid_thw) - /// pos_ids = [] - /// for t, h, w in grid_thw: - /// hpos_ids = arange(h).unsqueeze(1).expand(-1, w) - /// hpos_ids = hpos_ids.reshape(h//sms, sms, w//sms, sms).permute(0,2,1,3).flatten() - /// wpos_ids = arange(w).unsqueeze(0).expand(h, -1) - /// wpos_ids = wpos_ids.reshape(h//sms, sms, w//sms, sms).permute(0,2,1,3).flatten() - /// pos_ids.append(stack([hpos_ids, wpos_ids], dim=-1).repeat(t, 1)) - /// pos_ids = cat(pos_ids, dim=0) # (total, 2) - /// max_grid_size = grid_thw[:, 1:].max() - /// rotary_pos_emb_full = self.rotary_pos_emb(max_grid_size) # (max_grid_size, dim/4) - /// rotary_pos_emb = rotary_pos_emb_full[pos_ids].flatten(1) # (total, dim/2) - /// return rotary_pos_emb, pos_ids pub fn rot_pos_emb( &self, grid_thw: &[(usize, usize, usize)], @@ -431,14 +308,10 @@ impl GlmOcrVisionRotaryEmbedding { for &(t, h, w) in grid_thw { max_grid_size = max_grid_size.max(h).max(w); - // Generate position indices for each patch BEFORE spatial merge - // The vision encoder processes all patches, then merges them - // Python: for each (t, h, w), generate h*w position indices for _ in 0..t { for hi in 0..h { for wi in 0..w { // Apply spatial merge rearrangement - // Python: hpos_ids = hpos_ids.reshape(h//sms, sms, w//sms, sms).permute(0,2,1,3).flatten() let _hb = hi / sms; let _si = hi % sms; let _wb = wi / sms; @@ -457,10 +330,6 @@ impl GlmOcrVisionRotaryEmbedding { let total_seq = all_hpos.len(); let freqs_full = self.forward(max_grid_size)?; // (max_grid_size, dim/4) - // Python: rotary_pos_emb_full[pos_ids].flatten(1) - // pos_ids is (total, 2) with [h_idx, w_idx] entries - // rotary_pos_emb_full[pos_ids] -> (total, 2, dim/4) -> flatten(1) -> (total, dim/2) - // This is equivalent to cat(freqs[h_indices], freqs[w_indices], dim=-1) let h_indices = Tensor::from_vec(all_hpos, (total_seq,), self.inv_freq.device())?; let w_indices = Tensor::from_vec(all_wpos, (total_seq,), self.inv_freq.device())?; let h_freqs = freqs_full.index_select(&h_indices, 0)?; // (total_seq, dim/4) @@ -469,9 +338,6 @@ impl GlmOcrVisionRotaryEmbedding { // Concatenate h and w freqs: (total_seq, dim/2) let rotary_pos_emb = Tensor::cat(&[&h_freqs, &w_freqs], 1)?; - // Python (in GlmOcrVisionModel.forward): - // emb = cat((rotary_pos_emb, rotary_pos_emb), dim=-1) - // position_embeddings = (emb.cos(), emb.sin()) let emb = Tensor::cat(&[&rotary_pos_emb, &rotary_pos_emb], 1)?; let cos = emb.cos()?; let sin = emb.sin()?; @@ -480,16 +346,6 @@ impl GlmOcrVisionRotaryEmbedding { } } -// ============================================================================ -// 8. GlmOcrTextMLP -// ============================================================================ - -// Python: class GlmOcrTextMLP(nn.Module): -// def forward(self, hidden_states): -// up_states = self.gate_up_proj(hidden_states) -// gate, up_states = up_states.chunk(2, dim=-1) -// up_states = up_states * self.activation_fn(gate) -// return self.down_proj(up_states) pub struct GlmOcrTextMLP { gate_up_proj: Linear, down_proj: Linear, @@ -526,21 +382,6 @@ impl GlmOcrTextMLP { } } -// ============================================================================ -// 9. GlmOcrTextDecoderLayer -// ============================================================================ - -// Python: class GlmOcrTextDecoderLayer(nn.Module): -// def forward(self, hidden_states, position_embeddings): -// hidden_states = self.input_layernorm(hidden_states) -// hidden_states, _ = self.self_attn(hidden_states, position_embeddings) -// hidden_states = self.post_self_attn_layernorm(hidden_states) -// hidden_states = residual + hidden_states -// residual = hidden_states -// hidden_states = self.post_attention_layernorm(hidden_states) -// hidden_states = self.mlp(hidden_states) -// hidden_states = self.post_mlp_layernorm(hidden_states) -// return residual + hidden_states pub struct GlmOcrTextDecoderLayer { self_attn: GlmOcrTextAttention, mlp: GlmOcrTextMLP, @@ -611,28 +452,6 @@ impl GlmOcrTextDecoderLayer { } } -// ============================================================================ -// 10. rotate_half + apply_rotary_pos_emb_vision (in position_embed/rope.rs) -// ============================================================================ - -// ============================================================================ -// 11. GlmOcrVisionAttention -// ============================================================================ - -// Python reference (transformers): -// class GlmOcrVisionAttention(nn.Module): -// def __init__(self, config): -// self.qkv = nn.Linear(config.hidden_size, config.hidden_size * 3, bias=config.attention_bias) -// self.proj = nn.Linear(config.hidden_size, config.hidden_size, bias=config.attention_bias) -// self.q_norm = GlmOcrRMSNorm(self.head_dim, eps=config.rms_norm_eps) -// self.k_norm = GlmOcrRMSNorm(self.head_dim, eps=config.rms_norm_eps) -// def forward(self, hidden_states, position_embeddings): -// q, k, v = self.qkv(hidden_states).reshape(seq_len, 3, num_heads, -1).permute(1,0,2,3).unbind(0) -// query_states = self.q_norm(query_states) -// key_states = self.k_norm(key_states) -// q, k = apply_rotary_pos_emb_vision(q, k, cos, sin) -// attn_output = attention(q, k, v) -// return self.proj(attn_output.reshape(seq_len, -1)) pub struct GlmOcrVisionAttention { num_heads: usize, head_dim: usize, @@ -744,20 +563,6 @@ impl GlmOcrVisionAttention { } } -// ============================================================================ -// 12. GlmOcrVisionBlock -// ============================================================================ - -// Python: class GlmOcrVisionBlock(GradientCheckpointingLayer): -// def __init__(self, config): -// self.norm1 = GlmOcrRMSNorm(...) -// self.attn = GlmOcrVisionAttention(...) -// self.norm2 = GlmOcrRMSNorm(...) -// self.mlp = GlmOcrVisionMlp(...) -// def forward(self, x): -// x = x + self.attn(self.norm1(x)) -// x = x + self.mlp(self.norm2(x)) -// return x pub struct GlmOcrVisionBlock { norm1: GlmOcrRMSNorm, norm2: GlmOcrRMSNorm, @@ -801,28 +606,6 @@ impl GlmOcrVisionBlock { } } -// ============================================================================ -// 13. GlmOcrVisionPatchMerger -// ============================================================================ - -// Python reference (transformers commit 4854dbf9): -// class GlmOcrVisionPatchMerger(nn.Module): -// def __init__(self, dim, context_dim, hidden_act, bias=False): -// # dim = out_hidden_size (1536) -// # context_dim = intermediate_size (4096) NOT out_hidden_size * in_channels! -// self.proj = nn.Linear(dim, dim, bias=bias) -// self.post_projection_norm = LayerNorm(dim) -// self.gate_proj = nn.Linear(dim, context_dim, bias=bias) -// self.up_proj = nn.Linear(dim, context_dim, bias=bias) -// self.down_proj = nn.Linear(context_dim, dim, bias=bias) -// def forward(self, x): -// x = self.proj(x) -// x = self.post_projection_norm(x) -// x = GELU(x) # fixed activation, not config.hidden_act -// x = self.down_proj(self.act_fn(self.gate_proj(x)) * self.up_proj(x)) -// return x -// -// FIX: context_dim should be intermediate_size (4096), not out_hidden_size * in_channels (1536*3=4608) pub struct GlmOcrVisionPatchMerger { proj: Linear, post_projection_norm: LayerNorm, @@ -834,10 +617,6 @@ pub struct GlmOcrVisionPatchMerger { impl GlmOcrVisionPatchMerger { pub fn new(vb: VarBuilder, config: &GlmOcrVisionConfig) -> Result { - // NOTE: The checkpoint stores proj.weight as [out_features, in_features] = [1536, 1536] - // PyTorch Linear computes: output = input @ weight.T - // Candle Linear computes: output = weight.matmul(input) which is equivalent to input @ weight.T - // So the weight should be loaded correctly. let proj = linear_no_bias( config.out_hidden_size, config.out_hidden_size, @@ -850,8 +629,6 @@ impl GlmOcrVisionPatchMerger { vb.pp("post_projection_norm"), )?; - // Patch merger MLP: uses out_hidden_size * in_channels as intermediate dim - // Checkpoint shape: [4608, 1536] = [out_hidden_size * in_channels, out_hidden_size] let context_dim = config.out_hidden_size * config.in_channels; let gate_proj = linear_no_bias( config.out_hidden_size, @@ -882,7 +659,6 @@ impl GlmOcrVisionPatchMerger { pub fn forward(&self, hidden_state: &Tensor) -> Result { let mut hidden_state = self.proj.forward(hidden_state)?; hidden_state = self.post_projection_norm.forward(&hidden_state)?; - // Python: x = self.act1(x) where act1 = nn.GELU() - fixed GELU, not config activation hidden_state = hidden_state.gelu()?; let gate = self.gate_proj.forward(&hidden_state)?; @@ -894,17 +670,6 @@ impl GlmOcrVisionPatchMerger { } } -// ============================================================================ -// 14. GlmOcrVisionPatchEmbed -// ============================================================================ - -// Python: class GlmOcrVisionPatchEmbed(nn.Module): -// def __init__(self, config): -// self.proj = nn.Conv3d(...) -// def forward(self, x): -// x = x.view(-1, C, T, P, P) -// x = self.proj(x).view(-1, embed_dim) -// return x pub struct GlmOcrVisionPatchEmbed { patch_size: usize, temporal_patch_size: usize, @@ -945,27 +710,18 @@ impl GlmOcrVisionPatchEmbed { } pub fn forward(&self, pixel_values: &Tensor) -> Result { - // pixel_values can be either: - // 1. Flattened patches: [num_patches, patch_dim] (Python format) - // 2. Standard image: [batch, C, H, W] (traditional format) - let rank = pixel_values.rank(); if rank == 2 { - // Flattened patches format: [num_patches, patch_dim] - // Directly project to hidden_size let hidden_states = self.proj.forward(pixel_values)?; Ok(hidden_states) } else { - // Standard image format: [batch, C, H, W] let (batch, _c, h, w) = pixel_values.dims4()?; - // Support non-square images let patches_h = h / self.patch_size; let patches_w = w / self.patch_size; let num_patches = patches_h * patches_w; - // Reshape: (batch, C, H, W) -> (batch, patches_h, patch_size, patches_w, patch_size, C) let pv = pixel_values.reshape(( batch, patches_h, @@ -975,16 +731,13 @@ impl GlmOcrVisionPatchEmbed { self.in_channels, ))?; - // Permute: (batch, patches_h, patches_w, C, patch_size, patch_size) let pv = pv.permute((0, 1, 3, 5, 2, 4))?; - // Reshape: (batch * num_patches, C * patch_size * patch_size) let pv = pv.reshape(( batch * num_patches, self.in_channels * self.patch_size * self.patch_size, ))?; - // Add temporal dimension let pv = pv.unsqueeze(1)?; let ones_shape: Vec = vec![1, self.temporal_patch_size]; let pv = pv.broadcast_mul(&Tensor::ones(ones_shape, pv.dtype(), pv.device())?)?; @@ -999,22 +752,6 @@ impl GlmOcrVisionPatchEmbed { } } -// ============================================================================ -// 15. GlmOcrVisionModel -// ============================================================================ - -// Python: class GlmOcrVisionModel(GlmOcrPreTrainedModel): -// def __init__(self, config): -// self.patch_embed = GlmOcrVisionPatchEmbed(config) -// self.rotary_pos_emb = GlmOcrVisionRotaryEmbedding(...) -// self.blocks = nn.ModuleList([GlmOcrVisionBlock(config) for _ in range(config.depth)]) -// self.merger = GlmOcrVisionPatchMerger(...) -// self.downsample = nn.Conv2d(...) -// self.post_layernorm = GlmOcrRMSNorm(...) -// def forward(self, hidden_states, grid_thw): -// hidden_states = self.patch_embed(hidden_states) -// # ... apply rotary pos emb, blocks, merger, downsample -// return hidden_states pub struct GlmOcrVisionModel { patch_embed: GlmOcrVisionPatchEmbed, rotary_pos_emb: GlmOcrVisionRotaryEmbedding, @@ -1093,20 +830,13 @@ impl GlmOcrVisionModel { result }; - tensor_stats("after_patch_embed", &hidden_states); - - // Compute rotary embeddings matching Python rot_pos_emb exactly let (cos, sin) = self .rotary_pos_emb .rot_pos_emb(&grid_thw_parsed, self.config.spatial_merge_size)?; - tensor_stats("vision_cos", &cos); - tensor_stats("vision_sin", &sin); - let rotary_pos_emb = Tensor::cat(&[&cos, &sin], D::Minus1)?; let position_embeddings = (&cos, &sin); - // Compute cu_seqlens let mut cu_seqlens_values: Vec = vec![0]; let mut cumsum: i32 = 0; for (t, h, w) in &grid_thw_parsed { @@ -1122,67 +852,34 @@ impl GlmOcrVisionModel { hidden_states.device(), )?; - let num_blocks = self.blocks.len(); - if std::env::var("GLM_DEBUG").is_ok() { - let n_patches = hidden_states.dim(0).unwrap_or(0); - let hidden_mb = (hidden_states.elem_count() * 2) as f64 / 1_048_576.0; // bf16 - let attn_mb = (self.config.num_heads * n_patches * n_patches * 4) as f64 / 1_048_576.0; - eprintln!( - "[GLM-OCR OOM-DBG] Vision encoder: n_patches={n_patches} \ - hidden={hidden_mb:.1}MB per-layer attn_weights={attn_mb:.1}MB (f32) \ - num_blocks={num_blocks}" - ); - } - for (i, block) in self.blocks.iter().enumerate() { + for block in self.blocks.iter() { hidden_states = block.forward( &hidden_states, &cu_seqlens, Some(&rotary_pos_emb), Some(position_embeddings), )?; - if i == 0 { - tensor_stats("after_vision_block[0]", &hidden_states); - } - if i == num_blocks - 1 { - tensor_stats(&format!("after_vision_block[{i}] (last)"), &hidden_states); - } } let hidden_states = self.post_layernorm.forward(&hidden_states)?; - tensor_stats("after_vision_post_layernorm", &hidden_states); let sms = self.config.spatial_merge_size; let hidden_dim = hidden_states.dim(hidden_states.dims().len() - 1)?; - // Python: hidden_states.view(-1, sms, sms, hidden_dim).permute(0, 3, 1, 2) - // Input: [2816, 1024] where 2816 = grid_h * grid_w = 44 * 64 - // After reshape: [704, 2, 2, 1024] where 704 = 2816 / 4 - // After permute: [704, 1024, 2, 2] - // After downsample: [704, 1536, 1, 1] - // After reshape: [704, 1536] let total_patches = hidden_states.dim(0)?; // 2816 let merged_patches = total_patches / (sms * sms); // 704 let hidden_states = hidden_states.reshape((merged_patches, sms, sms, hidden_dim))?; let hidden_states = hidden_states.permute((0, 3, 1, 2))?; // [704, 1024, 2, 2] let hidden_states = self.downsample.forward(&hidden_states)?; // [704, 1536, 1, 1] let hidden_states = hidden_states.reshape((merged_patches, self.config.out_hidden_size))?; // [704, 1536] - - tensor_stats("after_downsample", &hidden_states); let merged = self.merger.forward(&hidden_states)?; - tensor_stats("after_merger (vision features)", &merged); let merged = merged.unsqueeze(0)?; Ok(merged) } } -// ============================================================================ -// GlmOcrProjector (Rust-specific, no Python equivalent) -// ============================================================================ - -// Rust-specific: Projects vision features to language model space -// (Not present in Python - integrated in GlmOcrVisionModel forward) pub struct GlmOcrProjector { #[allow(dead_code)] query_embed: Option, proj: Linear, @@ -1224,21 +921,6 @@ impl GlmOcrProjector { } } -// ============================================================================ -// 16. GlmOcrTextRotaryEmbedding -// ============================================================================ - -// Python: class GlmOcrTextRotaryEmbedding(nn.Module): -// def forward(self, x, position_ids): # position_ids: (3, bs, seq_len) -// inv_freq_expanded = self.inv_freq[None, None, :, None].expand(3, bs, -1, 1) -// position_ids_expanded = position_ids[:, :, None, :].float() -// freqs = (inv_freq_expanded @ position_ids_expanded).transpose(2, 3) -// freqs = self.apply_mrope(freqs, self.mrope_section) -// emb = torch.cat((freqs, freqs), dim=-1) -// return emb.cos() * self.attention_scaling, emb.sin() * self.attention_scaling -// def apply_mrope(self, freqs, mrope_section): -// chunks = freqs.split(mrope_section, dim=-1) -// return torch.cat([chunk[i % 3] for i, chunk in enumerate(chunks)], dim=-1) pub struct GlmOcrTextRotaryEmbedding { inv_freq: Tensor, mrope_section: Vec, @@ -1300,12 +982,10 @@ impl GlmOcrTextRotaryEmbedding { } Ok(Tensor::cat(&result_parts, D::Minus1)?) } - - /// Compute cos/sin from explicit 3D position IDs (used for prefill with image tokens). - /// position_ids: (3, bs, seq_len) — axis 0 = temporal, 1 = height, 2 = width. - pub fn forward_with_position_ids(&self, position_ids: &Tensor) -> Result<(Tensor, Tensor)> { + + pub fn forward_with_position_ids(&self, position_ids: &Tensor) -> Result<(Tensor, Tensor)> { let (_, bs, _seq_len) = position_ids.dims3()?; - let inv_freq_len = self.inv_freq.dim(1)?; // head_dim/2 + let inv_freq_len = self.inv_freq.dim(1)?; // inv_freq: (1, inv_freq_len) -> broadcast to (3, bs, inv_freq_len, 1) let inv_freq = self.inv_freq.unsqueeze(0)?.unsqueeze(D::Minus1)?; // (1, 1, hd/2, 1) @@ -1378,19 +1058,6 @@ impl GlmOcrTextRotaryEmbedding { } } -// ============================================================================ -// 17. GlmOcrTextModel -// ============================================================================ - -// Python: class GlmOcrTextModel(GlmOcrPreTrainedModel): -// def __init__(self, config): -// self.embed_tokens = nn.Embedding(...) -// self.layers = nn.ModuleList([GlmOcrTextDecoderLayer(config) for _ in range(...)]) -// self.norm = GlmOcrRMSNorm(...) -// self.rotary_emb = GlmOcrTextRotaryEmbedding(...) -// def forward(self, input_ids, ...): -// hidden_states = self.embed_tokens(input_ids) -// # ... apply layers, rotary emb, return hidden_states pub struct GlmOcrTextModel { embed_tokens: Embedding, layers: Vec, @@ -1417,7 +1084,6 @@ impl GlmOcrTextModel { let norm = GlmOcrRMSNorm::new(vb.pp("norm"), config.hidden_size, config.rms_norm_eps)?; - // lm_head.weight lives at the checkpoint root (not under model.language_model) let root_vb = vb.root(); let lm_head = linear_no_bias(config.hidden_size, config.vocab_size, root_vb.pp("lm_head"))?; @@ -1436,13 +1102,6 @@ impl GlmOcrTextModel { }) } - /// Compute 3D M-RoPE position IDs matching Python's GlmOcrModel.get_rope_index(). - /// - /// For image tokens: each gets (t, h, w) grid coordinates. - /// For text tokens: sequential positions on all 3 axes. - /// Positions continue from where the previous group left off. - /// - /// Returns tensor of shape (3, 1, seq_len) containing [temporal, height, width] position IDs. fn compute_mrope_position_ids( &mut self, image_mask: &Tensor, @@ -1516,13 +1175,6 @@ impl GlmOcrTextModel { self.next_mrope_pos = st_idx as usize; self.prefill_seq_len = seq_len; - if std::env::var("GLM_DEBUG").is_ok() { - eprintln!( - "[M-RoPE] prefill: seq_len={}, next_mrope_pos={}", - seq_len, self.next_mrope_pos - ); - } - let t_t = Tensor::from_vec(t_ids, (1, seq_len), device)?; let h_t = Tensor::from_vec(h_ids, (1, seq_len), device)?; let w_t = Tensor::from_vec(w_ids, (1, seq_len), device)?; @@ -1539,17 +1191,6 @@ impl GlmOcrTextModel { ) -> Result { let (bs, seq_len) = input_ids.dims2()?; let mut inputs_embeds = self.embed_tokens.forward(input_ids)?; - tensor_stats("embed_tokens output", &inputs_embeds); - - #[cfg(debug_assertions)] - eprintln!( - "LanguageModel forward: bs={}, seq_len={}, head_dim={}", - bs, - seq_len, - self.config - .head_dim - .unwrap_or_else(|| self.config.hidden_size / self.config.num_attention_heads) - ); // Merge image features into embeddings at image token positions if let (Some(img_feats), Some(img_mask)) = (image_features, image_mask) { @@ -1569,15 +1210,6 @@ impl GlmOcrTextModel { let num_features = img_feats.dim(1)?; let num_to_replace = image_indices.len().min(num_features); - if std::env::var("GLM_DEBUG").is_ok() { - eprintln!("[LM] Image indices count: {}, num_features: {}, num_to_replace: {}", - image_indices.len(), num_features, num_to_replace); - if !image_indices.is_empty() { - eprintln!("[LM] First image index: {}, Last image index: {}", - image_indices.first().unwrap(), image_indices.last().unwrap()); - } - } - // Replace embeddings at image positions with image features // Build the merged embeddings by copying let embeds_flat = inputs_embeds.squeeze(0)?; // (seq_len, hidden_size) @@ -1599,16 +1231,7 @@ impl GlmOcrTextModel { let refs: Vec<&Tensor> = embeds_vec.iter().collect(); inputs_embeds = Tensor::cat(&refs, 0)?.unsqueeze(0)?; - if std::env::var("GLM_DEBUG").is_ok() { - eprintln!("[LM] inputs_embeds after image injection shape: {:?}", inputs_embeds.shape()); - if seq_len > 710 { - let text_embed = inputs_embeds.i((0, 710, ..))?; - let text_mean = text_embed.to_dtype(DType::F32)?.mean_all()?.to_scalar::()?; - eprintln!("[LM] Text token 710 embedding mean: {:.6}", text_mean); - } - } } - tensor_stats("inputs_embeds after image injection", &inputs_embeds); let attention_mask = if seq_len > 1 { Some(prepare_causal_attention_mask( @@ -1638,37 +1261,15 @@ impl GlmOcrTextModel { let decode_pos = self.next_mrope_pos + (seqlen_offset - self.prefill_seq_len); self.rotary_emb.forward(1, decode_pos, input_ids.device())? }; - tensor_stats("text_rotary cos", &cos); - tensor_stats("text_rotary sin", &sin); - let num_layers = self.layers.len(); let mut hidden_states = inputs_embeds; - for (i, layer) in self.layers.iter_mut().enumerate() { + for layer in self.layers.iter_mut() { hidden_states = layer.forward(&hidden_states, (&cos, &sin), attention_mask.as_ref())?; - if i == 0 { - tensor_stats("after_text_layer[0]", &hidden_states); - } - if i == num_layers - 1 { - tensor_stats(&format!("after_text_layer[{i}] (last)"), &hidden_states); - } } hidden_states = self.norm.forward(&hidden_states)?; - tensor_stats("after_final_norm", &hidden_states); let logits = self.lm_head.forward(&hidden_states)?; - // Log top-5 logits at last position (only on first pass) - if seqlen_offset == 0 && std::env::var("GLM_INTERMEDIATE").is_ok() { - let last_logits = logits.i((0, seq_len - 1, ..))?.to_dtype(DType::F32)?; - tensor_stats("logits at last position", &last_logits); - if let Ok(vals) = last_logits.to_vec1::() { - let mut indexed: Vec<(usize, f32)> = vals.iter().copied().enumerate().collect(); - indexed.sort_by(|a, b| b.1.partial_cmp(&a.1).unwrap()); - eprintln!("[RS] Top-5 logit token_ids: {:?}", - indexed[..5.min(indexed.len())].iter().map(|(i, v)| format!("id={i} val={v:.4}")).collect::>()); - } - } - Ok(logits) } @@ -1679,17 +1280,6 @@ impl GlmOcrTextModel { } } -// ============================================================================ -// 18. GlmOcrModel (corresponds to GlmOcrForConditionalGeneration in Python) -// ============================================================================ - -// Python: class GlmOcrForConditionalGeneration(GlmOcrPreTrainedModel, GenerationMixin): -// def __init__(self, config): -// self.model = GlmOcrModel(config) # Contains visual + language_model -// self.lm_head = nn.Linear(...) -// def forward(self, input_ids, pixel_values, ...): -// # Vision encoder -> language model -> lm_head -// return logits pub struct GlmOcrModel { vision_encoder: GlmOcrVisionModel, language_model: GlmOcrTextModel, @@ -1719,24 +1309,6 @@ impl GlmOcrModel { image_mask: Option<&Tensor>, seqlen_offset: usize, ) -> Result { - #[cfg(debug_assertions)] - { - eprintln!("[GLM-OCR Model] ===== FORWARD START ====="); - eprintln!("[GLM-OCR Model] input_ids shape: {:?}", input_ids.shape()); - eprintln!("[GLM-OCR Model] seqlen_offset: {}", seqlen_offset); - if let Some(pv) = pixel_values { - eprintln!("[GLM-OCR Model] pixel_values shape: {:?}", pv.shape()); - } - if let Some(g) = image_grid_thw { - eprintln!("[GLM-OCR Model] grid_thw shape: {:?}", g.shape()); - } - if let Some(m) = image_mask { - eprintln!("[GLM-OCR Model] image_mask shape: {:?}", m.shape()); - let mask_sum = m.sum_all()?.to_scalar::()?; - eprintln!("[GLM-OCR Model] image_mask sum (num image tokens): {}", mask_sum); - } - } - let image_features = if let Some(pixels) = pixel_values { let grid_thw = if let Some(grid) = image_grid_thw { grid.clone() @@ -1751,26 +1323,8 @@ impl GlmOcrModel { )? }; - if std::env::var("GLM_DEBUG").is_ok() { - let pv_mb = (pixels.elem_count() * 2) as f64 / 1_048_576.0; // bf16 - eprintln!( - "[GLM-OCR OOM-DBG] pixel_values shape={:?} size={pv_mb:.1}MB", - pixels.shape() - ); - eprintln!("[GLM-OCR OOM-DBG] grid_thw={:?}", grid_thw.to_vec1::()); - } - let vision_output = self.vision_encoder.forward(pixels, &grid_thw)?; - if std::env::var("GLM_DEBUG").is_ok() { - eprintln!( - "[GLM-OCR Model] vision_output shape: {:?}", - vision_output.shape() - ); - let vis_mean = vision_output.to_dtype(candle_core::DType::F32)?.mean_all()?.to_scalar::()?; - eprintln!("[GLM-OCR Model] vision_output mean: {:.6}", vis_mean); - } - Some(vision_output) } else { None @@ -1784,14 +1338,6 @@ impl GlmOcrModel { seqlen_offset, ); - #[cfg(debug_assertions)] - { - eprintln!("[GLM-OCR Model] ===== FORWARD END ====="); - if let Ok(ref r) = result { - eprintln!("[GLM-OCR Model] output shape: {:?}", r.shape()); - } - } - result } diff --git a/src/models/glm_ocr/processor.rs b/src/models/glm_ocr/processor.rs index c9ed381..ceabe56 100644 --- a/src/models/glm_ocr/processor.rs +++ b/src/models/glm_ocr/processor.rs @@ -149,20 +149,7 @@ impl GlmOcrProcessor { // Use smart_resize to compute target dimensions let (target_h, target_w) = self.smart_resize(orig_h, orig_w); - if std::env::var("GLM_DEBUG").is_ok() { - let grid_h = target_h / self.patch_size; - let grid_w = target_w / self.patch_size; - let n_patches = grid_h * grid_w; - let pv_elems = n_patches * 3 * self.temporal_patch_size * self.patch_size * self.patch_size; - let pv_mb = (pv_elems * 2) as f64 / 1_048_576.0; // bf16 - let attn_mb = (16 * n_patches * n_patches * 4) as f64 / 1_048_576.0; // 16 heads, f32 - eprintln!( - "[GLM-OCR OOM-DBG] Image: orig={orig_w}x{orig_h} -> target={target_w}x{target_h} \ - grid={grid_w}x{grid_h} n_patches={n_patches} \ - pixel_values={pv_mb:.1}MB vision_attn_per_layer={attn_mb:.1}MB" - ); - } - + // Resize image let img = img.resize_exact( target_w as u32, @@ -183,10 +170,8 @@ impl GlmOcrProcessor { }) .collect(); - // Reshape to [H, W, 3] let tensor = Tensor::from_vec(pixels, (target_h, target_w, 3), &self.device)?; - // Normalize let mean = Tensor::new(self.image_mean.clone(), &self.device)?.reshape((1, 1, 3))?; let std = Tensor::new(self.image_std.clone(), &self.device)?.reshape((1, 1, 3))?; let tensor = tensor.broadcast_sub(&mean)?.broadcast_div(&std)?; @@ -196,37 +181,25 @@ impl GlmOcrProcessor { let grid_w = target_w / self.patch_size; let patch_size = self.patch_size; let channels = 3; - let temporal_patch_size = 2; // Python uses temporal_patch_size=2 even for images + let temporal_patch_size = 2; - // Reshape: [H, W, 3] -> [grid_h, patch_size, grid_w, patch_size, 3] let tensor = tensor.reshape(( grid_h, patch_size, grid_w, patch_size, channels, ))?; - // Permute to: [grid_h, grid_w, patch_size, patch_size, channels] - let tensor = tensor.permute((0, 2, 1, 3, 4))?; + let tensor = tensor.permute((0, 2, 4, 1, 3))?; - // Permute to put channels first: [grid_h, grid_w, channels, patch_size, patch_size] - // Python's patch_embed.forward does: view(-1, C, T, P, P) so we need (C, T, P_h, P_w) order - let tensor = tensor.permute((0, 1, 4, 2, 3))?; - - // Reshape to: [num_patches, channels, patch_size, patch_size] let num_patches = grid_h * grid_w; let tensor = tensor.reshape((num_patches, channels, patch_size, patch_size))?; - // Add temporal dimension after C: [num_patches, channels, 1, patch_size, patch_size] let tensor = tensor.unsqueeze(2)?; - // Repeat T times along temporal dim 2: [num_patches, channels, temporal_patch_size, patch_size, patch_size] let tensor = tensor.repeat((1, 1, temporal_patch_size, 1, 1))?; - // Flatten to: [num_patches, channels * temporal_patch_size * patch_size * patch_size] - // This gives (C, T, P_h, P_w) order matching Python's patch_embed input let patch_dim = channels * temporal_patch_size * patch_size * patch_size; let tensor = tensor.reshape((num_patches, patch_dim))?; - // Convert to model dtype let tensor = tensor.to_dtype(self.dtype)?; Ok(ProcessedImage { @@ -236,25 +209,7 @@ impl GlmOcrProcessor { }) } - /// Process image and text for multimodal input. - /// - /// # Arguments - /// * `image_path` - Path to input image - /// * `prompt` - Text prompt/question about the image - /// * `tokenizer` - Tokenizer for text encoding - /// * `image_token_id` - Token ID for image content placeholders - /// * `image_start_token_id` - Token ID marking start of image region - /// * `image_end_token_id` - Token ID marking end of image region - /// * `patch_size` - Vision encoder patch size (default: 14) - /// * `temporal_patch_size` - Temporal patch size (default: 2, unused for images) - /// * `spatial_merge_size` - Spatial merge factor (default: 2) - /// - /// # Returns - /// ProcessedInput containing: - /// - input_ids: Combined image placeholder + text token IDs - /// - pixel_values: Flattened patches tensor [num_patches, patch_dim] - /// - image_mask: Boolean mask for image token positions - /// - grid_thw: (temporal, height, width) grid dimensions for RoPE + /// Process image and text for multimodal input pub fn process_info( &self, image_path: &str, @@ -327,7 +282,6 @@ impl GlmOcrProcessor { let image_mask = Tensor::from_vec(image_mask_vec, (1, input_ids_vec.len()), &self.device)?; // Compute grid_thw for RoPE - // For images: grid_t = 1 let grid_thw = Tensor::from_vec( vec![1u32, grid_h as u32, grid_w as u32], (3,),