rahul7star/gemma4-e4b
0
1import gradio as gr2import torch3import time4import traceback5from threading import Thread6from transformers import AutoModelForCausalLM, AutoTokenizer, TextIteratorStreamer7 8model_id = "rahul7star/gemma-4-finetune"9 10 11def log(msg):12 print(f"[DEBUG] {msg}", flush=True)13 14 15log("Starting Gemma 4 debug app")16log(f"Model ID: {model_id}")17log(f"Torch version: {torch.__version__}")18log(f"CUDA available: {torch.cuda.is_available()}")19 20if torch.cuda.is_available():21 log(f"CUDA device count: {torch.cuda.device_count()}")22 log(f"CUDA device name: {torch.cuda.get_device_name(0)}")23 24 25# ============================================================26# Load Tokenizer27# ============================================================28 29log("Loading tokenizer...")30 31tokenizer = AutoTokenizer.from_pretrained(32 model_id,33 trust_remote_code=True,34)35 36log("Tokenizer loaded")37log(f"Tokenizer class: {tokenizer.__class__.__name__}")38log(f"Vocab size: {len(tokenizer)}")39log(f"EOS token: {tokenizer.eos_token} / {tokenizer.eos_token_id}")40log(f"PAD token: {tokenizer.pad_token} / {tokenizer.pad_token_id}")41log(f"Chat template exists: {tokenizer.chat_template is not None}")42 43if tokenizer.pad_token_id is None:44 tokenizer.pad_token = tokenizer.eos_token45 log("PAD token was missing, set PAD token = EOS token")46 47 48# ============================================================49# Load Model50# ============================================================51 52log("Loading model...")53 54model = AutoModelForCausalLM.from_pretrained(55 model_id,56 device_map="cpu",57 low_cpu_mem_usage=True,58 torch_dtype=torch.bfloat16,59 trust_remote_code=True,60)61 62model.eval()63 64log("Model loaded")65log(f"Model class: {model.__class__.__name__}")66log(f"Model device: {model.device}")67log(f"Model dtype: {next(model.parameters()).dtype}")68 69 70# ============================================================71# Config Logs72# ============================================================73 74cfg = model.config75text_cfg = getattr(cfg, "text_config", None)76vision_cfg = getattr(cfg, "vision_config", None)77 78log("========== MAIN MODEL CONFIG ==========")79log(f"model_type: {getattr(cfg, 'model_type', None)}")80log(f"architectures: {getattr(cfg, 'architectures', None)}")81log(f"is_encoder_decoder: {getattr(cfg, 'is_encoder_decoder', None)}")82log(f"text_config exists: {text_cfg is not None}")83log(f"vision_config exists: {vision_cfg is not None}")84log("=======================================")85 86if text_cfg is not None:87 log("========== TEXT CONFIG ==========")88 log(f"model_type: {getattr(text_cfg, 'model_type', None)}")89 log(f"hidden_size: {getattr(text_cfg, 'hidden_size', None)}")90 log(f"intermediate_size: {getattr(text_cfg, 'intermediate_size', None)}")91 log(f"num_hidden_layers: {getattr(text_cfg, 'num_hidden_layers', None)}")92 log(f"num_attention_heads: {getattr(text_cfg, 'num_attention_heads', None)}")93 log(f"num_key_value_heads: {getattr(text_cfg, 'num_key_value_heads', None)}")94 log(f"head_dim: {getattr(text_cfg, 'head_dim', None)}")95 log(f"vocab_size: {getattr(text_cfg, 'vocab_size', None)}")96 log(f"max_position_embeddings: {getattr(text_cfg, 'max_position_embeddings', None)}")97 log(f"rope_theta: {getattr(text_cfg, 'rope_theta', None)}")98 log(f"rms_norm_eps: {getattr(text_cfg, 'rms_norm_eps', None)}")99 log(f"attention_bias: {getattr(text_cfg, 'attention_bias', None)}")100 log(f"use_cache: {getattr(text_cfg, 'use_cache', None)}")101 log(f"sliding_window: {getattr(text_cfg, 'sliding_window', None)}")102 log(f"query_pre_attn_scalar: {getattr(text_cfg, 'query_pre_attn_scalar', None)}")103 log(f"final_logit_softcapping: {getattr(text_cfg, 'final_logit_softcapping', None)}")104 log(f"attn_logit_softcapping: {getattr(text_cfg, 'attn_logit_softcapping', None)}")105 log("=================================")106 107if vision_cfg is not None:108 log("========== VISION CONFIG ==========")109 log(f"model_type: {getattr(vision_cfg, 'model_type', None)}")110 log(f"hidden_size: {getattr(vision_cfg, 'hidden_size', None)}")111 log(f"intermediate_size: {getattr(vision_cfg, 'intermediate_size', None)}")112 log(f"num_hidden_layers: {getattr(vision_cfg, 'num_hidden_layers', None)}")113 log(f"num_attention_heads: {getattr(vision_cfg, 'num_attention_heads', None)}")114 log(f"image_size: {getattr(vision_cfg, 'image_size', None)}")115 log(f"patch_size: {getattr(vision_cfg, 'patch_size', None)}")116 log("===================================")117 118 119# ============================================================120# Parameter Logs121# ============================================================122 123total_params = sum(p.numel() for p in model.parameters())124trainable_params = sum(p.numel() for p in model.parameters() if p.requires_grad)125 126log(f"Total parameters: {total_params:,}")127log(f"Trainable parameters: {trainable_params:,}")128 129 130# ============================================================131# Module Listing132# ============================================================133 134log("========== TEXT MODEL MODULES ==========")135 136text_keywords = [137 "language_model",138 "text_model",139 "model.layers",140 "self_attn",141 "mlp",142 "input_layernorm",143 "post_attention_layernorm",144 "q_proj",145 "k_proj",146 "v_proj",147 "o_proj",148 "gate_proj",149 "up_proj",150 "down_proj",151 "rotary",152 "embed_tokens",153 "lm_head",154]155 156count = 0157 158for name, module in model.named_modules():159 lower = name.lower()160 161 if "vision_tower" in lower:162 continue163 164 if any(k in lower for k in text_keywords):165 log(f"{name} => {module.__class__.__name__}")166 count += 1167 168 if count >= 200:169 log("Stopped text module logging after 200 entries")170 break171 172log(f"Text modules logged: {count}")173log("========================================")174 175 176# ============================================================177# Deep Forward Hooks178# ============================================================179 180DEBUG_HOOKS = True181HOOK_EVERY_N_CALLS = 20182_hook_calls = {}183 184 185def tensor_stats(x):186 if not torch.is_tensor(x):187 return str(type(x))188 189 with torch.no_grad():190 xf = x.detach().float()191 return (192 f"shape={tuple(x.shape)}, "193 f"dtype={x.dtype}, "194 f"device={x.device}, "195 f"mean={xf.mean().item():.5f}, "196 f"std={xf.std().item():.5f}, "197 f"min={xf.min().item():.5f}, "198 f"max={xf.max().item():.5f}"199 )200 201 202def get_first_tensor(obj):203 if torch.is_tensor(obj):204 return obj205 206 if isinstance(obj, (list, tuple)):207 for item in obj:208 t = get_first_tensor(item)209 if t is not None:210 return t211 212 if isinstance(obj, dict):213 for item in obj.values():214 t = get_first_tensor(item)215 if t is not None:216 return t217 218 return None219 220 221def make_hook(name):222 def hook(module, inputs, output):223 if not DEBUG_HOOKS:224 return225 226 _hook_calls[name] = _hook_calls.get(name, 0) + 1227 228 if _hook_calls[name] % HOOK_EVERY_N_CALLS != 1:229 return230 231 inp = get_first_tensor(inputs)232 out = get_first_tensor(output)233 234 log(f"HOOK: {name}")235 236 if inp is not None:237 log(f" input -> {tensor_stats(inp)}")238 239 if out is not None:240 log(f" output -> {tensor_stats(out)}")241 242 return hook243 244 245def attach_debug_hooks():246 wanted = [247 "self_attn",248 "q_proj",249 "k_proj",250 "v_proj",251 "o_proj",252 "mlp",253 "gate_proj",254 "up_proj",255 "down_proj",256 "input_layernorm",257 "post_attention_layernorm",258 "rotary_emb",259 "lm_head",260 ]261 262 attached = 0263 264 for name, module in model.named_modules():265 lower = name.lower()266 267 if "vision_tower" in lower:268 continue269 270 if any(w in lower for w in wanted):271 module.register_forward_hook(make_hook(name))272 attached += 1273 274 log(f"Attached debug hooks: {attached}")275 276 277#attach_debug_hooks()278 279 280# ============================================================281# Generation Function282# ============================================================283 284def generate_response(message, history):285 start_time = time.time()286 287 log("========== NEW GENERATION ==========")288 log(f"User message: {message}")289 log(f"History turns: {len(history)}")290 291 messages = []292 293 for item in history:294 try:295 user_msg, bot_msg = item296 messages.append({"role": "user", "content": user_msg})297 messages.append({"role": "assistant", "content": bot_msg})298 except Exception as e:299 log(f"History parse warning: {e}")300 log(f"Bad history item: {item}")301 302 messages.append({"role": "user", "content": message})303 304 log(f"Total chat messages: {len(messages)}")305 306 try:307 inputs = tokenizer.apply_chat_template(308 messages,309 return_tensors="pt",310 return_dict=True,311 add_generation_prompt=True,312 ).to(model.device)313 314 input_token_count = inputs["input_ids"].shape[-1]315 316 log(f"Input tensor shape: {inputs['input_ids'].shape}")317 log(f"Input tokens: {input_token_count}")318 log(f"Input device: {inputs['input_ids'].device}")319 320 log("========== TOKEN DEBUG ==========")321 ids = inputs["input_ids"][0].tolist()322 log(f"First 20 token ids: {ids[:20]}")323 log(f"Last 20 token ids: {ids[-20:]}")324 log(f"Decoded prompt preview: {tokenizer.decode(ids[-200:], skip_special_tokens=False)}")325 log("=================================")326 327 except Exception as e:328 log("Chat template/tokenization failed")329 log(traceback.format_exc())330 yield f"Tokenization error: {e}"331 return332 333 streamer = TextIteratorStreamer(334 tokenizer,335 timeout=420.0,336 skip_prompt=True,337 skip_special_tokens=True,338 )339 340 341 generate_kwargs = dict(342 **inputs,343 streamer=streamer,344 max_new_tokens=1024,345 temperature=0.7,346 do_sample=False,347 top_p=0.9,348 pad_token_id=tokenizer.pad_token_id,349 eos_token_id=tokenizer.eos_token_id,350 )351 352 log("Generation kwargs:")353 log("max_new_tokens=1024")354 log("temperature=0.7")355 log("do_sample=True")356 log("top_p=0.9")357 358 def run_generation():359 try:360 log("Generation thread started")361 362 gen_start = time.time()363 364 with torch.no_grad():365 model.generate(**generate_kwargs)366 367 gen_time = time.time() - gen_start368 log(f"Generation thread finished in {gen_time:.2f}s")369 370 except Exception as e:371 log("Generation Error")372 log(traceback.format_exc())373 374 streamer.text_queue.put(375 f"\n[Generation thread crashed. Reason: {e}]"376 )377 streamer.end()378 379 t = Thread(target=run_generation)380 t.start()381 382 partial_text = ""383 token_chunks = 0384 385 try:386 for new_text in streamer:387 token_chunks += 1388 partial_text += new_text389 390 if token_chunks % 20 == 0:391 elapsed = time.time() - start_time392 log(393 f"Streaming chunks: {token_chunks}, "394 f"chars: {len(partial_text)}, "395 f"elapsed: {elapsed:.2f}s"396 )397 398 yield partial_text399 400 except Exception as e:401 log("Streaming Error")402 log(traceback.format_exc())403 yield partial_text + f"\n\n[Streaming error: {e}]"404 405 finally:406 elapsed = time.time() - start_time407 log("========== GENERATION DONE ==========")408 log(f"Output chars: {len(partial_text)}")409 log(f"Streaming chunks: {token_chunks}")410 log(f"Elapsed seconds: {elapsed:.2f}")411 log("=====================================")412 413 414# ============================================================415# Gradio UI416# ============================================================417 418demo = gr.ChatInterface(419 fn=generate_response,420 title="Gemma 4 E4B - Debug",421 examples=[422 "Explain quantum entanglement simply.",423 "Write a Python function to add two numbers.",424 "Explain how RoPE works in transformer attention.",425 ],426 cache_examples=False,427)428 429if __name__ == "__main__":430 log("Launching Gradio app...")431 demo.launch()