tiny-random/model-weight-inspector
0
1import json2import tempfile3import os4import glob5import shutil6import io7import time8import threading9import sys10 11import gradio as gr12import torch13from huggingface_hub import hf_hub_download, scan_cache_dir, whoami14from safetensors import safe_open15 16# Default token from HF_TOKEN environment variable (for HuggingFace Spaces)17DEFAULT_HF_TOKEN = os.environ.get("HF_TOKEN")18 19 20def hf_login(token: str, session_token: str):21 """Login to Hugging Face with provided token (per-user session)."""22 if not token:23 return "❌ Please provide a token", "Not logged in", session_token24 25 try:26 user_info = whoami(token=token)27 username = user_info.get('name', 'Unknown')28 return f"✅ Successfully logged in as: {username}", f"✅ Logged in as {username}", token29 except Exception as e:30 return f"❌ Login failed: {str(e)}", "❌ Not logged in", session_token31 32 33def hf_logout(session_token: str):34 """Logout from Hugging Face (clear session token)."""35 return "✅ Successfully logged out", "Not logged in", None36 37 38def check_hf_status(session_token: str):39 """Check current HF login status for this session."""40 # Check session token first, then fall back to default token41 token = session_token or DEFAULT_HF_TOKEN42 43 if not token:44 return "ℹ️ Not logged in", "Not logged in", session_token45 46 try:47 user_info = whoami(token=token)48 username = user_info.get('name', 'Unknown')49 source = "(session)" if session_token else "(default HF_TOKEN)"50 return f"✅ Currently logged in as: {username} {source}", f"✅ Logged in as {username}", session_token51 except Exception:52 return "ℹ️ Not logged in", "Not logged in", session_token53 54 55def get_param(model_id: str, param_key: str, log_buffer: io.StringIO, progress: gr.Progress, token: str = None):56 """57 Download and return a specific parameter tensor from a Hugging Face model.58 """59 # Use session token or fall back to default token60 auth_token = token or DEFAULT_HF_TOKEN61 62 # Redirect stderr to log buffer for real-time tqdm updates63 original_stderr = sys.stderr64 sys.stderr = log_buffer65 66 try:67 # Try to download the index file (for sharded models)68 try:69 log_buffer.write(f"📥 Downloading index file for {model_id}...\n")70 progress(0.1, desc="Downloading index...")71 72 index_path = hf_hub_download(73 model_id, "model.safetensors.index.json", token=auth_token)74 75 log_buffer.write(f"✓ Index file found: {index_path}\n")76 77 with open(index_path, "r", encoding="utf-8") as f:78 index = json.load(f)79 weight_map = index["weight_map"]80 if param_key not in weight_map:81 raise KeyError(82 f"Parameter '{param_key}' not found in model. Available keys: {list(weight_map.keys())[:10]}..."83 )84 shard_file = weight_map[param_key]85 log_buffer.write(f"✓ Parameter found in shard: {shard_file}\n")86 except Exception as e:87 if "404" in str(e) or "not found" in str(e).lower():88 log_buffer.write("ℹ️ No index file, trying single model file...\n")89 shard_file = "model.safetensors"90 else:91 raise92 93 log_buffer.write(f"📥 Downloading shard: {shard_file}...\n")94 progress(0.3, desc=f"Downloading {shard_file}...")95 96 shard_path = hf_hub_download(model_id, shard_file, token=auth_token)97 98 log_buffer.write(f"\n✓ Shard downloaded: {shard_path}\n")99 progress(0.7, desc="Loading tensor...")100 101 log_buffer.write(f"🔍 Loading tensor '{param_key}'...\n")102 with safe_open(shard_path, framework="pt") as f:103 tensor = f.get_tensor(param_key)104 log_buffer.write(f"✓ Tensor loaded successfully\n")105 progress(0.9, desc="Finalizing...")106 107 return tensor108 finally:109 # Restore original stderr110 sys.stderr = original_stderr111 112 113def get_available_keys(model_id: str, token: str = None):114 """Get all available parameter keys from a model."""115 # Use session token or fall back to default token116 auth_token = token or DEFAULT_HF_TOKEN117 118 try:119 index_path = hf_hub_download(model_id, "model.safetensors.index.json", token=auth_token)120 with open(index_path, "r", encoding="utf-8") as f:121 index = json.load(f)122 return sorted(index["weight_map"].keys())123 except Exception:124 # Try single file125 try:126 shard_path = hf_hub_download(model_id, "model.safetensors", token=auth_token)127 with safe_open(shard_path, framework="pt") as f:128 return sorted(f.keys())129 except Exception as e:130 return []131 132 133def format_tensor_info(tensor: torch.Tensor) -> str:134 """Format tensor information for display."""135 info = []136 info.append(f"**Shape:** {list(tensor.shape)}")137 info.append(f"**Dtype:** {tensor.dtype}")138 info.append(f"**Device:** {tensor.device}")139 info.append(f"**Numel:** {tensor.numel():,}")140 141 # Handle special dtypes that don't support statistical operations142 try:143 # Convert FP8 and other special dtypes to float32 for stats144 if str(tensor.dtype) in ['torch.float8_e4m3fn', 'torch.float8_e5m2']:145 stats_tensor = tensor.to(torch.float32)146 else:147 stats_tensor = tensor148 149 info.append(f"**Min:** {stats_tensor.min().item():.6f}")150 info.append(f"**Max:** {stats_tensor.max().item():.6f}")151 info.append(f"**Mean:** {stats_tensor.float().mean().item():.6f}")152 info.append(f"**Std:** {stats_tensor.float().std().item():.6f}")153 except Exception as e:154 info.append(f"**Stats:** Unable to compute (dtype not supported)")155 156 return "<br>".join(info)157 158 159def fetch_param(model_id: str, param_key: str, session_token: str, progress=gr.Progress()):160 """Fetch parameter and return formatted info and tensor preview."""161 log_buffer = io.StringIO()162 last_log_value = ""163 164 if not model_id or not param_key:165 yield "Please provide both model ID and parameter key.", "", None, "❌ Missing required inputs"166 return167 168 try:169 log_buffer.write(f"🚀 Starting download for {model_id}\n")170 log_buffer.write(f"🎯 Target parameter: {param_key}\n\n")171 progress(0, desc="Initializing...")172 yield "", "", None, log_buffer.getvalue()173 time.sleep(0.5)174 175 # Start download in background thread176 download_complete = threading.Event()177 download_error = [None] # Use list to store exception from thread178 result_tensor = [None] # Use list to store result from thread179 180 def download_thread():181 try:182 result_tensor[0] = get_param(model_id, param_key, log_buffer, progress, session_token)183 except Exception as e:184 download_error[0] = e185 finally:186 download_complete.set()187 188 thread = threading.Thread(target=download_thread, daemon=True)189 thread.start()190 191 # Poll log buffer every 1 second while download is running192 while not download_complete.is_set():193 current_log = log_buffer.getvalue()194 if current_log != last_log_value:195 yield "", "", None, current_log196 last_log_value = current_log197 time.sleep(1)198 199 # Final log update after download completes200 current_log = log_buffer.getvalue()201 if current_log != last_log_value:202 yield "", "", None, current_log203 last_log_value = current_log204 205 # Check for errors206 if download_error[0]:207 raise download_error[0]208 209 tensor = result_tensor[0]210 info = format_tensor_info(tensor)211 212 # Create tensor preview (first few elements)213 log_buffer.write(f"\n📊 Creating preview...\n")214 yield "", "", None, log_buffer.getvalue()215 216 flat = tensor.flatten()217 preview_size = min(100, flat.numel())218 219 # Convert to float32 for FP8 types for display220 if str(tensor.dtype) in ['torch.float8_e4m3fn', 'torch.float8_e5m2']:221 preview = flat[:preview_size].to(torch.float32).tolist()222 else:223 preview = flat[:preview_size].tolist()224 225 # Format preview in multiple lines (10 values per line)226 # Adapt to different data types227 preview_lines = []228 for i in range(0, len(preview), 10):229 line_values = preview[i:i+10]230 if tensor.dtype in [torch.float32, torch.float64, torch.float16, torch.bfloat16] or str(tensor.dtype) in ['torch.float8_e4m3fn', 'torch.float8_e5m2']:231 preview_lines.append(", ".join(f"{v:.6f}" for v in line_values))232 elif tensor.dtype in [torch.int8, torch.int16, torch.int32, torch.int64, torch.uint8]:233 preview_lines.append(", ".join(f"{v}" for v in line_values))234 elif tensor.dtype == torch.bool:235 preview_lines.append(", ".join(f"{v}" for v in line_values))236 else:237 preview_lines.append(", ".join(str(v) for v in line_values))238 239 preview_str = f"**First {preview_size} values:**\n```\n" + \240 "\n".join(preview_lines) + "\n```"241 242 # if flat.numel() > preview_size:243 # preview_str += f"\n\n... and {flat.numel() - preview_size:,} more values"244 245 # Save tensor for download246 log_buffer.write(f"💾 Saving tensor for download...\n")247 yield info, preview_str, None, log_buffer.getvalue()248 249 temp_dir = tempfile.gettempdir()250 safe_param_key = param_key.replace("/", "_").replace(".", "_")251 download_path = os.path.join(temp_dir, f"{safe_param_key}.pt")252 torch.save(tensor, download_path)253 log_buffer.write(f"✓ Saved to: {download_path}\n")254 255 progress(1.0, desc="Complete!")256 log_buffer.write(f"\n✅ All operations completed successfully!\n")257 yield info, preview_str, download_path, log_buffer.getvalue()258 except Exception as e:259 log_buffer.write(f"\n❌ Error: {str(e)}\n")260 yield f"**Error:** {str(e)}", "", None, log_buffer.getvalue()261 262 263def list_keys(model_id: str, session_token: str):264 """List all available keys for a model."""265 if not model_id:266 return "Please provide a model ID."267 268 try:269 keys = get_available_keys(model_id, session_token)270 if not keys:271 return "No keys found or failed to load model."272 return "\n".join(keys)273 except Exception as e:274 return f"**Error:** {str(e)}"275 276 277def clear_temp_files():278 """Clear all .pt files from temp directory."""279 try:280 temp_dir = tempfile.gettempdir()281 pt_files = glob.glob(os.path.join(temp_dir, "*.pt"))282 count = len(pt_files)283 deleted_files = []284 for file in pt_files:285 try:286 os.remove(file)287 deleted_files.append(os.path.basename(file))288 except Exception:289 pass290 291 if deleted_files:292 files_list = "\n".join(deleted_files)293 return f"✅ Cleared {count} temporary file(s):\n\n{files_list}"294 else:295 return "✅ No temporary files to clear"296 except Exception as e:297 return f"❌ Error: {str(e)}"298 299 300def clear_hf_cache():301 """Clear Hugging Face cache directory."""302 try:303 cache_info = scan_cache_dir()304 total_size = cache_info.size_on_disk305 total_repos = len(cache_info.repos)306 307 if total_repos == 0:308 return "✅ Hugging Face cache is already empty"309 310 # Get cache directory and clear it311 cache_dir = os.path.expanduser("~/.cache/huggingface/hub")312 if os.path.exists(cache_dir):313 shutil.rmtree(cache_dir)314 os.makedirs(cache_dir)315 size_mb = total_size / (1024 * 1024)316 return f"✅ Cleared Hugging Face cache: {total_repos} repo(s), {size_mb:.2f} MB freed"317 else:318 return "✅ Hugging Face cache directory not found"319 except Exception as e:320 return f"❌ Error: {str(e)}"321 322 323def get_cache_info():324 """Get size information about caches."""325 try:326 # Temp files327 temp_dir = tempfile.gettempdir()328 pt_files = glob.glob(os.path.join(temp_dir, "*.pt"))329 temp_size = sum(os.path.getsize(f)330 for f in pt_files if os.path.exists(f))331 temp_size_mb = temp_size / (1024 * 1024)332 333 info = f"📊 Cache Info:\n\n"334 info += f"═══ Temp .pt files: {len(pt_files)} file(s), {temp_size_mb:.2f} MB ═══\n"335 336 if pt_files:337 for file in pt_files:338 size = os.path.getsize(file) / (1024 * 1024)339 filename = os.path.basename(file)340 info += f" • {filename} ({size:.2f} MB)\n"341 else:342 info += " (empty)\n"343 344 # HF cache345 info += f"\n═══ Hugging Face Cache ═══\n"346 try:347 cache_info = scan_cache_dir()348 hf_size_mb = cache_info.size_on_disk / (1024 * 1024)349 hf_repos = len(cache_info.repos)350 351 info += f"Total: {hf_repos} repo(s), {hf_size_mb:.2f} MB\n\n"352 353 if hf_repos > 0:354 for repo in cache_info.repos:355 repo_size = repo.size_on_disk / (1024 * 1024)356 info += f" 📦 {repo.repo_id}\n"357 info += f" Size: {repo_size:.2f} MB, Revisions: {len(repo.revisions)}\n"358 info += f" Last accessed: {repo.last_accessed}\n"359 else:360 info += " (empty)\n"361 except Exception as e:362 info += f" Error reading HF cache: {str(e)}\n"363 364 info += f"\n═══ Total: {temp_size_mb + (hf_size_mb if 'hf_size_mb' in locals() else 0):.2f} MB ═══"365 return info366 except Exception as e:367 return f"❌ Error: {str(e)}"368 369 370# Create Gradio interface371custom_css = """372* {373 font-family: Consolas, Monaco, 'Courier New', monospace !important;374}375.compact-row {376 gap: 0.5rem !important;377}378.tensor-preview pre {379 font-size: 0.75rem !important;380 line-height: 1.0 !important;381}382.compact-file {383 max-height: 80px !important;384}385.compact-file > div {386 min-height: 60px !important;387}388"""389 390with gr.Blocks(title="Hugging Face Model Weight Inspector") as demo:391 gr.Markdown("# 🔍 Hugging Face Model Weight Inspector")392 393 # Session state for per-user token394 session_token = gr.State(None)395 396 # HF Login section397 with gr.Accordion("🔐 Hugging Face Login (Per-User Session) [⚠️⚠️⚠️WIP, Do not use⚠️⚠️⚠️]", open=False):398 gr.Markdown("""399 **Note:** This Space uses the default `HF_TOKEN` secret for all users if no session token is provided. 400 Login below with your own token for per-user authentication (affects only your session).401 """)402 with gr.Row():403 with gr.Column(scale=3):404 hf_token_input = gr.Textbox(405 label="HF Token",406 placeholder="hf_...",407 type="password",408 )409 with gr.Column(scale=2):410 initial_status = "✅ Using default HF_TOKEN" if DEFAULT_HF_TOKEN else "Not logged in"411 hf_status = gr.Textbox(412 label="Status",413 value=initial_status,414 interactive=False,415 )416 with gr.Row():417 login_btn = gr.Button("🔑 Login", variant="primary", scale=1)418 logout_btn = gr.Button("🚪 Logout", variant="secondary", scale=1)419 check_status_btn = gr.Button("ℹ️ Check Status", variant="secondary", scale=1)420 login_output = gr.Textbox(label="Login Status", interactive=False, lines=2)421 422 with gr.Row():423 with gr.Column(scale=1):424 model_id_input = gr.Textbox(425 label="Model ID",426 placeholder="e.g., meta-llama/Llama-2-7b-hf",427 value="Qwen/Qwen3-Coder-Next-FP8",428 )429 param_key_input = gr.Textbox(430 label="Parameter Key",431 placeholder="e.g., model.norm.weight",432 value="model.norm.weight",433 )434 with gr.Row():435 list_keys_btn = gr.Button(436 "📋 List Keys", variant="secondary", scale=1)437 fetch_btn = gr.Button("🔎 Fetch", variant="primary", scale=1)438 439 with gr.Column(scale=1):440 keys_output = gr.Textbox(441 label="Available Parameter Keys",442 lines=5,443 max_lines=8,444 )445 446 with gr.Tabs():447 with gr.Tab("Results"):448 with gr.Row():449 with gr.Column(scale=3):450 preview_output = gr.Markdown(label="Tensor Preview", elem_classes="tensor-preview")451 with gr.Column(scale=1):452 info_output = gr.Markdown(label="Tensor Info")453 download_output = gr.File(label="Download Tensor (.pt file)", elem_classes="compact-file")454 log_output = gr.Textbox(455 label="📋 Download Log", lines=1, interactive=False)456 457 with gr.Tab("Cache Management"):458 with gr.Row():459 get_info_btn = gr.Button(460 "📊 Get Cache Info", variant="secondary", scale=1)461 clear_temp_btn = gr.Button(462 "🗑️ Clear Temp Folder", variant="secondary", scale=1)463 clear_hf_btn = gr.Button(464 "🗑️ Clear HF Cache", variant="secondary", scale=1)465 clear_status = gr.Textbox(466 label="Status", interactive=False, lines=6)467 468 # Event handlers469 login_btn.click(470 fn=hf_login,471 inputs=[hf_token_input, session_token],472 outputs=[login_output, hf_status, session_token],473 )474 475 logout_btn.click(476 fn=hf_logout,477 inputs=[session_token],478 outputs=[login_output, hf_status, session_token],479 )480 481 check_status_btn.click(482 fn=check_hf_status,483 inputs=[session_token],484 outputs=[login_output, hf_status, session_token],485 )486 487 list_keys_btn.click(488 fn=list_keys,489 inputs=[model_id_input, session_token],490 outputs=[keys_output],491 )492 493 fetch_btn.click(494 fn=fetch_param,495 inputs=[model_id_input, param_key_input, session_token],496 outputs=[info_output, preview_output, download_output, log_output],497 )498 499 clear_temp_btn.click(500 fn=clear_temp_files,501 inputs=[],502 outputs=[clear_status],503 )504 505 clear_hf_btn.click(506 fn=clear_hf_cache,507 inputs=[],508 outputs=[clear_status],509 )510 511 get_info_btn.click(512 fn=get_cache_info,513 inputs=[],514 outputs=[clear_status],515 )516 517 # Auto-check status on load518 demo.load(519 fn=check_hf_status,520 inputs=[session_token],521 outputs=[login_output, hf_status, session_token],522 )523 524 525if __name__ == "__main__":526 demo.launch(server_name="0.0.0.0", css=custom_css)527 