CoolFace
Apppublic

rahul7star/gemma4-e4b

sourceHugging Faceupdated 5mo agoView on Hugging Face
0likes
app.py431 linesDownload Raw Back to root
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()