Team Ai
Apppublic

tiny-random/model-weight-inspector

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