{"metadata":{"kernelspec":{"language":"python","display_name":"Python 3","name":"python3"},"language_info":{"name":"python","version":"3.12.12","mimetype":"text/x-python","codemirror_mode":{"name":"ipython","version":3},"pygments_lexer":"ipython3","nbconvert_exporter":"python","file_extension":".py"},"kaggle":{"accelerator":"gpu","dataSources":[{"sourceType":"competition","sourceId":6799,"databundleVersionId":4225553}],"dockerImageVersionId":31287,"isInternetEnabled":true,"language":"python","sourceType":"notebook","isGpuEnabled":true}},"nbformat_minor":4,"nbformat":4,"cells":[{"cell_type":"code","source":"import os\nimport json\nimport math\nimport gc\nfrom pathlib import Path\n\nimport torch\nimport torch.nn as nn\nimport torch.nn.functional as F\nfrom torch.utils.data import Dataset, DataLoader, ConcatDataset\nfrom torchvision import transforms\nfrom PIL import Image\nimport numpy as np\n\n# ---------------------------------------------------------------------------\n# FULL RUN CONFIGURATION (Actual Training)\n# ---------------------------------------------------------------------------\nCFG = {\n    # Paths (Kaggle working directory)\n    \"working_dir\": \"/kaggle/working\",\n    \"activations_path\": \"/kaggle/working/activations_full.pt\",\n    \"sae_weights_path\": \"/kaggle/working/micro_sae_1024d_full.pt\",\n    \"sae_config_path\": \"/kaggle/working/config_full.json\",\n\n    # Model\n    \"model_id\": \"llava-hf/llava-1.5-7b-hf\",\n\n    # Dataset sizes -- FULL SCALE\n    \"coco_num_images\": 10_000,        # Broad baseline for diverse features\n    \"concept_num_images\": 500,        # 500 zebras, 500 fire trucks\n\n    # Activation extraction\n    \"vision_hidden_dim\": 1024,        \n    \"num_patches\": 576,               \n    \"extraction_batch_size\": 16,      # Keeps VRAM safe during extraction\n    \"penultimate_layer_idx\": -2,      \n    \"save_every_n_batches\": 200,      # Flushes to disk every ~3200 images to prevent RAM crashes\n\n    # SAE architecture\n    \"sae_input_dim\": 1024,\n    \"sae_dict_size\": 4096,            # 4x expansion factor\n    \"sae_topk\": 32,                 \n\n    # SAE training\n    \"sae_epochs\": 2,                  # 2 passes over ~12 million tokens is plenty for low-level features\n    \"sae_batch_size\": 4096,           # Large batch size for fast stable gradient updates\n    \"sae_lr\": 3e-4,                   # Slightly lower LR for stability on the large dataset\n    \"sae_l1_coeff\": 1e-3,             \n\n    # Hub upload\n    \"hf_repo_name\": \"llava-micro-sae\", # The final production repository name\n}\n\n# Ensure output directory exists (no-op on Kaggle, useful locally)\nos.makedirs(CFG[\"working_dir\"], exist_ok=True)","metadata":{"trusted":true},"outputs":[],"execution_count":null},{"cell_type":"code","source":"# Hugging Face Authentication\ndef step1_authenticate():\n    \"\"\"\n    Retrieve HF_TOKEN from Kaggle secrets and log in.\n    Falls back to environment variable HF_TOKEN if not running on Kaggle.\n    \"\"\"\n    print(\"=\" * 70)\n    print(\"STEP 1: Hugging Face Authentication\")\n    print(\"=\" * 70)\n\n    token = None\n\n    # --- Try Kaggle secrets first ---\n    try:\n        from kaggle_secrets import UserSecretsClient\n        user_secrets = UserSecretsClient()\n        token = user_secrets.get_secret(\"HF_TOKEN\")\n        print(\"[OK] Retrieved HF_TOKEN from Kaggle secrets.\")\n    except Exception as e:\n        print(f\"[WARN] Kaggle secrets unavailable ({e}). Trying env variable...\")\n        token = os.environ.get(\"HF_TOKEN\", None)\n\n    if token is None:\n        raise RuntimeError(\n            \"No HF_TOKEN found. Set it as a Kaggle secret or env variable.\"\n        )\n\n    from huggingface_hub import login\n    login(token=token)\n    print(\"[OK] Logged in to Hugging Face Hub.\\n\")\n    return token","metadata":{"trusted":true},"outputs":[],"execution_count":null},{"cell_type":"code","source":"# Dataset Curation (Disk-Streaming + Kaggle Native ImageNet)\n\nimport shutil\n\nclass DiskImageDataset(Dataset):\n    \"\"\"Reads images one-by-one from local disk to save CPU RAM.\"\"\"\n    def __init__(self, image_dir: Path, transform):\n        self.image_paths = list(image_dir.glob(\"*.jpg\"))\n        self.transform = transform\n\n    def __len__(self):\n        return len(self.image_paths)\n\n    def __getitem__(self, idx):\n        with Image.open(self.image_paths[idx]) as img:\n            img_rgb = img.convert(\"RGB\")\n            tensor = self.transform(img_rgb)\n        return tensor\n\n\ndef step2_curate_dataset(processor):\n    print(\"=\" * 70)\n    print(\"STEP 2: Dataset Curation (Hybrid Local/Streaming Mode)\")\n    print(\"=\" * 70)\n\n    from datasets import load_dataset\n\n    img_dir = Path(CFG[\"working_dir\"]) / \"images\"\n    if img_dir.exists():\n        shutil.rmtree(img_dir)  \n    img_dir.mkdir(parents=True, exist_ok=True)\n\n    def clip_transform(pil_img: Image.Image) -> torch.Tensor:\n        out = processor.image_processor(images=pil_img, return_tensors=\"pt\")\n        return out[\"pixel_values\"].squeeze(0)\n\n    # --- 2a. COCO images (Streaming - this worked perfectly in your last run) ---\n    print(f\"[...] Streaming and saving {CFG['coco_num_images']} COCO images to disk...\")\n    try:\n        ds_stream = load_dataset(\"detection-datasets/coco\", split=\"train\", streaming=True)\n        count = 0\n        for ex in ds_stream:\n            if count >= CFG[\"coco_num_images\"]: break\n            try:\n                img = ex[\"image\"]\n                if not isinstance(img, Image.Image):\n                    img = Image.fromarray(np.array(img))\n                \n                img_resized = img.convert(\"RGB\").resize((336, 336), Image.Resampling.LANCZOS)\n                img_resized.save(img_dir / f\"coco_{count}.jpg\", quality=85)\n                \n                img.close()\n                img_resized.close()\n                del img, img_resized, ex\n            except Exception:\n                continue\n            \n            count += 1\n            if count % 1000 == 0:\n                gc.collect()\n                print(f\"   ... {count}/{CFG['coco_num_images']} saved\")\n        \n        gc.collect()\n        print(f\"[OK] COCO: Saved {count} images.\")\n    except Exception as e:\n        print(f\"[FAIL] COCO download failed: {e}\")\n\n    # --- 2b & 2c. Concept Images (Direct from Kaggle Native ImageNet) ---\n    def fetch_local_imagenet(wnid, name, max_count):\n        print(f\"[...] Fetching {max_count} '{name}' images from local Kaggle ImageNet...\")\n        # Path to the mounted Kaggle Competition Dataset\n        source_dir = Path(f\"/kaggle/input/competitions/imagenet-object-localization-challenge/ILSVRC/Data/CLS-LOC/train/{wnid}\")\n        \n        if not source_dir.exists():\n            print(f\"[FAIL] Directory not found: {source_dir}\")\n            print(\"      Did you click 'Add Data' and attach 'imagenet-object-localization-challenge'?\")\n            return 0\n            \n        jpeg_files = list(source_dir.glob(\"*.JPEG\"))\n        \n        count = 0\n        for file_path in jpeg_files:\n            if count >= max_count: break\n            try:\n                with Image.open(file_path) as img:\n                    img_resized = img.convert(\"RGB\").resize((336, 336), Image.Resampling.LANCZOS)\n                    img_resized.save(img_dir / f\"{name}_{count}.jpg\", quality=85)\n                    img_resized.close()\n                count += 1\n            except Exception:\n                continue\n                \n        print(f\"[OK] {name}: Saved {count} images locally.\")\n        return count\n\n    # n02391049 is the exact WordNet ID for Zebra\n    fetch_local_imagenet(\"n02391049\", \"zebra\", CFG[\"concept_num_images\"])\n    \n    # n03345487 is the exact WordNet ID for Fire engine / Fire truck\n    fetch_local_imagenet(\"n03345487\", \"firetruck\", CFG[\"concept_num_images\"])\n\n    # --- Create DataLoader reading strictly from Disk ---\n    disk_dataset = DiskImageDataset(img_dir, transform=clip_transform)\n    print(f\"\\n[OK] Total images securely stored on disk: {len(disk_dataset)}\")\n\n    loader = DataLoader(\n        disk_dataset,\n        batch_size=CFG[\"extraction_batch_size\"],\n        shuffle=True,\n        num_workers=2,\n        pin_memory=True,\n        drop_last=False,\n    )\n    print(f\"[OK] DataLoader ready ({len(loader)} batches).\\n\")\n    return loader","metadata":{"trusted":true},"outputs":[],"execution_count":null},{"cell_type":"code","source":"# Activation Extraction\ndef step3_extract_activations(loader):\n    print(\"=\" * 70)\n    print(\"STEP 3: Activation Extraction\")\n    print(\"=\" * 70)\n\n    from transformers import LlavaForConditionalGeneration\n\n    print(\"[...] Loading LLaVA model (vision tower only will be used)...\")\n    model = LlavaForConditionalGeneration.from_pretrained(\n        CFG[\"model_id\"],\n        torch_dtype=torch.float16,\n        device_map=\"auto\",\n        low_cpu_mem_usage=True,\n    )\n    model.eval()\n    print(\"[OK] Model loaded.\\n\")\n\n    vision_tower = model.model.vision_tower\n\n    chunk_size = CFG[\"save_every_n_batches\"]\n    chunk_buffer = []\n    chunk_idx = 0\n    total_patches = 0\n    chunk_dir = Path(CFG[\"working_dir\"]) / \"activation_chunks\"\n    chunk_dir.mkdir(parents=True, exist_ok=True)\n\n    print(\"[...] Extracting patch activations from CLIP Layer 23...\")\n    \n    # 1. THE EXTRACTION LOOP (Reads JPEGs from disk)\n    with torch.no_grad():\n        for batch_idx, pixel_values in enumerate(loader):\n            pixel_values = pixel_values.to(\n                device=next(vision_tower.parameters()).device,\n                dtype=torch.float16,\n            )\n\n            vision_outputs = vision_tower(\n                pixel_values,\n                output_hidden_states=True,\n            )\n\n            hidden = vision_outputs.hidden_states[CFG[\"penultimate_layer_idx\"]]\n            patch_tokens = hidden[:, 1:, :]  \n\n            flat = patch_tokens.reshape(-1, CFG[\"vision_hidden_dim\"]).cpu().half()\n            chunk_buffer.append(flat)\n            total_patches += flat.shape[0]\n\n            del vision_outputs, hidden, patch_tokens, flat, pixel_values\n            torch.cuda.empty_cache()\n\n            if (batch_idx + 1) % 50 == 0:\n                print(f\"   Batch {batch_idx + 1}/{len(loader)} -- {total_patches:,} patches so far\")\n\n            if (batch_idx + 1) % chunk_size == 0:\n                chunk_tensor = torch.cat(chunk_buffer, dim=0)\n                chunk_path = chunk_dir / f\"chunk_{chunk_idx:04d}.pt\"\n                torch.save(chunk_tensor, chunk_path)\n                print(f\"   [SAVE] Chunk {chunk_idx} saved: {chunk_tensor.shape[0]:,} patches -> {chunk_path}\")\n                \n                del chunk_tensor, chunk_buffer\n                chunk_buffer = []\n                chunk_idx += 1\n                gc.collect()\n\n    if chunk_buffer:\n        chunk_tensor = torch.cat(chunk_buffer, dim=0)\n        chunk_path = chunk_dir / f\"chunk_{chunk_idx:04d}.pt\"\n        torch.save(chunk_tensor, chunk_path)\n        print(f\"   [SAVE] Final chunk {chunk_idx} saved: {chunk_tensor.shape[0]:,} patches\")\n        del chunk_tensor, chunk_buffer\n        gc.collect()\n\n    print(f\"\\n[OK] Total patches extracted: {total_patches:,}\")\n\n    # 2. DELETE JPEGS TO FREE DISK SPACE \n    # (Safe to do now because DataLoader is finished!)\n    img_dir = Path(CFG[\"working_dir\"]) / \"images\"\n    if img_dir.exists():\n        import shutil\n        shutil.rmtree(img_dir)\n        print(\"[OK] Deleted original JPEG images to free disk space.\")\n\n    # 3. DELETE LLAVA MODEL TO FREE RAM/GPU\n    del model, vision_tower\n    gc.collect()\n    torch.cuda.empty_cache()\n    print(\"[OK] LLaVA model deleted and GPU cache cleared.\")\n\n    # 4. PRE-ALLOCATED MERGE (Extremely safe for RAM/Disk limits)\n    print(\"[...] Merging activation chunks into single file...\")\n    chunk_files = sorted(chunk_dir.glob(\"chunk_*.pt\"))\n    \n    total_rows = 0\n    for cf in chunk_files:\n        c = torch.load(cf, map_location=\"cpu\")\n        total_rows += c.shape[0]\n        del c\n    gc.collect()\n    print(f\"   Total rows to merge: {total_rows:,}\")\n    \n    activations_tensor = torch.empty((total_rows, CFG[\"vision_hidden_dim\"]), dtype=torch.float16)\n    offset = 0\n    for cf in chunk_files:\n        c = torch.load(cf, map_location=\"cpu\")\n        rows = c.shape[0]\n        activations_tensor[offset:offset + rows] = c\n        offset += rows\n        cf.unlink()  # delete chunk IMMEDIATELY to free disk space\n        del c\n        gc.collect()\n    \n    print(f\"[OK] Final activation tensor: {activations_tensor.shape} ({activations_tensor.dtype})\")\n    torch.save(activations_tensor, CFG[\"activations_path\"])\n    print(f\"[SAVE] Saved to {CFG['activations_path']}\")\n\n    chunk_dir.rmdir()\n    del activations_tensor\n    gc.collect()\n    print(\"[OK] Activation extraction complete.\\n\")","metadata":{"trusted":true},"outputs":[],"execution_count":null},{"cell_type":"code","source":"# Train the Micro-SAE\n\nclass SparseAutoencoder(nn.Module):\n    def __init__(self, input_dim: int, dict_size: int):\n        super().__init__()\n        self.input_dim = input_dim\n        self.dict_size = dict_size\n\n        self.encoder = nn.Linear(input_dim, dict_size)\n        self.decoder = nn.Linear(dict_size, input_dim, bias=True)\n\n        nn.init.xavier_uniform_(self.encoder.weight)\n        nn.init.zeros_(self.encoder.bias)\n        nn.init.xavier_uniform_(self.decoder.weight)\n        nn.init.zeros_(self.decoder.bias)\n\n    def encode(self, x: torch.Tensor) -> torch.Tensor:\n        return F.relu(self.encoder(x))\n\n    def decode(self, z: torch.Tensor) -> torch.Tensor:\n        return self.decoder(z)\n\n    def forward(self, x: torch.Tensor):\n        z = self.encode(x)\n        x_hat = self.decode(z)\n        return x_hat, z\n\n    @torch.no_grad()\n    def normalize_decoder(self):\n        norms = self.decoder.weight.norm(dim=0, keepdim=True).clamp(min=1e-8)\n        self.decoder.weight.div_(norms)\n\n\nclass TopKSparseAutoencoder(nn.Module):\n    def __init__(self, input_dim: int, dict_size: int, k: int = 32):\n        super().__init__()\n        self.input_dim = input_dim\n        self.dict_size = dict_size\n        self.k = k\n\n        self.encoder = nn.Linear(input_dim, dict_size)\n        self.decoder = nn.Linear(dict_size, input_dim, bias=True)\n\n        nn.init.xavier_uniform_(self.encoder.weight)\n        nn.init.zeros_(self.encoder.bias)\n        nn.init.xavier_uniform_(self.decoder.weight)\n        nn.init.zeros_(self.decoder.bias)\n\n    def encode(self, x: torch.Tensor) -> torch.Tensor:\n        pre_act = self.encoder(x)\n        topk_vals, topk_idx = pre_act.topk(self.k, dim=-1)\n        z = torch.zeros_like(pre_act)\n        z.scatter_(-1, topk_idx, F.relu(topk_vals))\n        return z\n\n    def decode(self, z: torch.Tensor) -> torch.Tensor:\n        return self.decoder(z)\n\n    def forward(self, x: torch.Tensor):\n        z = self.encode(x)\n        x_hat = self.decode(z)\n        return x_hat, z\n\n    @torch.no_grad()\n    def normalize_decoder(self):\n        norms = self.decoder.weight.norm(dim=0, keepdim=True).clamp(min=1e-8)\n        self.decoder.weight.div_(norms)\n\n\nclass ActivationDataset(Dataset):\n    def __init__(self, tensor: torch.Tensor):\n        self.data = tensor\n\n    def __len__(self):\n        return self.data.shape[0]\n\n    def __getitem__(self, idx):\n        return self.data[idx]\n\n\ndef step4_train_sae():\n    print(\"=\" * 70)\n    print(\"STEP 4: Train the Micro-SAE\")\n    print(\"=\" * 70)\n\n    device = torch.device(\"cuda\" if torch.cuda.is_available() else \"cpu\")\n\n    # 1. LOAD DATA IN FLOAT16 (RAM usage: ~13 GB -> Safe!)\n    print(f\"[...] Loading activations from {CFG['activations_path']}...\")\n    activations = torch.load(CFG[\"activations_path\"], map_location=\"cpu\")\n    print(f\"[OK] Loaded {activations.shape[0]:,} activation vectors of dim {activations.shape[1]}.\")\n    assert activations.shape[1] == CFG[\"sae_input_dim\"]\n\n    # Shuffle once on CPU\n    perm = torch.randperm(activations.shape[0])\n    activations = activations[perm]\n\n    act_dataset = ActivationDataset(activations)\n    act_loader = DataLoader(\n        act_dataset,\n        batch_size=CFG[\"sae_batch_size\"],\n        shuffle=True,\n        num_workers=2,\n        pin_memory=True,\n        drop_last=True,\n    )\n    print(f\"[OK] Training DataLoader: {len(act_loader)} batches per epoch.\\n\")\n\n    # 2. INIT MODEL\n    use_topk = CFG[\"sae_topk\"] is not None\n    if use_topk:\n        print(f\"[...] Using TopK SAE (k={CFG['sae_topk']})\")\n        sae = TopKSparseAutoencoder(\n            input_dim=CFG[\"sae_input_dim\"],\n            dict_size=CFG[\"sae_dict_size\"],\n            k=CFG[\"sae_topk\"],\n        ).to(device)\n    else:\n        print(\"[...] Using ReLU SAE with L1 sparsity penalty\")\n        sae = SparseAutoencoder(\n            input_dim=CFG[\"sae_input_dim\"],\n            dict_size=CFG[\"sae_dict_size\"],\n        ).to(device)\n\n    total_params = sum(p.numel() for p in sae.parameters())\n    print(f\"[OK] SAE parameters: {total_params:,}\")\n\n    optimizer = torch.optim.Adam(sae.parameters(), lr=CFG[\"sae_lr\"])\n\n    # 3. TRAINING LOOP\n    print(\"[...] Starting training...\\n\")\n    for epoch in range(CFG[\"sae_epochs\"]):\n        epoch_mse = 0.0\n        epoch_l1 = 0.0\n        epoch_loss = 0.0\n        n_batches = 0\n\n        sae.train()\n        for batch_idx, batch in enumerate(act_loader):\n            \n            # --- THE MEMORY FIX ---\n            # Cast the float16 batch to float32 strictly on the GPU!\n            batch = batch.to(device).float()\n\n            x_hat, z = sae(batch)\n            mse_loss = F.mse_loss(x_hat, batch)\n\n            if use_topk:\n                l1_loss = torch.tensor(0.0, device=device)\n                loss = mse_loss\n            else:\n                l1_loss = z.abs().mean()\n                loss = mse_loss + CFG[\"sae_l1_coeff\"] * l1_loss\n\n            optimizer.zero_grad()\n            loss.backward()\n            optimizer.step()\n\n            sae.normalize_decoder()\n\n            epoch_mse += mse_loss.item()\n            epoch_l1 += l1_loss.item()\n            epoch_loss += loss.item()\n            n_batches += 1\n\n            if (batch_idx + 1) % 200 == 0:\n                avg_mse = epoch_mse / n_batches\n                avg_l1 = epoch_l1 / n_batches\n                with torch.no_grad():\n                    l0 = (z > 0).float().sum(dim=-1).mean().item()\n                print(\n                    f\"   Epoch {epoch+1} | Batch {batch_idx+1}/{len(act_loader)} | \"\n                    f\"MSE={avg_mse:.6f} | L1={avg_l1:.6f} | L0={l0:.1f}\"\n                )\n\n        avg_mse = epoch_mse / n_batches\n        avg_l1 = epoch_l1 / n_batches\n        avg_loss = epoch_loss / n_batches\n        print(\n            f\"\\n   * Epoch {epoch+1}/{CFG['sae_epochs']} complete -- \"\n            f\"Loss={avg_loss:.6f} | MSE={avg_mse:.6f} | L1={avg_l1:.6f}\\n\"\n        )\n\n    # 4. SAVE ARTIFACTS\n    print(f\"[SAVE] Saving SAE weights to {CFG['sae_weights_path']}...\")\n    torch.save(sae.state_dict(), CFG[\"sae_weights_path\"])\n\n    config = {\n        \"architecture\": \"TopKSparseAutoencoder\" if use_topk else \"SparseAutoencoder\",\n        \"input_dim\": CFG[\"sae_input_dim\"],\n        \"dict_size\": CFG[\"sae_dict_size\"],\n        \"topk\": CFG[\"sae_topk\"],\n        \"activation\": \"topk\" if use_topk else \"relu\",\n        \"source_model\": CFG[\"model_id\"],\n        \"source_layer\": \"vision_tower.hidden_states[-2] (Layer 23)\",\n        \"num_patches_per_image\": CFG[\"num_patches\"],\n        \"training_epochs\": CFG[\"sae_epochs\"],\n        \"training_lr\": CFG[\"sae_lr\"],\n        \"l1_coeff\": CFG[\"sae_l1_coeff\"] if not use_topk else None,\n        \"total_training_vectors\": len(activations),\n    }\n    with open(CFG[\"sae_config_path\"], \"w\") as f:\n        json.dump(config, f, indent=2)\n    print(f\"[SAVE] Saved config to {CFG['sae_config_path']}\")\n    print(f\"[OK] SAE training complete.\\n\")\n\n    del activations, sae\n    gc.collect()\n    torch.cuda.empty_cache()","metadata":{"trusted":true},"outputs":[],"execution_count":null},{"cell_type":"code","source":"# Upload to Hugging Face Hub\ndef step5_upload_to_hub(token: str):\n    \"\"\"\n    Create a HF repo and upload the SAE weights + config.\n    \"\"\"\n    print(\"=\" * 70)\n    print(\"STEP 5: Upload to Hugging Face Hub\")\n    print(\"=\" * 70)\n\n    from huggingface_hub import HfApi\n\n    api = HfApi()\n\n    # Get the authenticated user's username\n    user_info = api.whoami(token=token)\n    username = user_info[\"name\"]\n    repo_id = f\"{username}/{CFG['hf_repo_name']}\"\n\n    # Create repo (no-op if it already exists)\n    print(f\"[...] Creating repo: {repo_id}\")\n    api.create_repo(\n        repo_id=repo_id,\n        repo_type=\"model\",\n        exist_ok=True,\n        token=token,\n    )\n    print(f\"[OK] Repo ready: https://huggingface.co/{repo_id}\")\n\n    # Upload weights\n    print(f\"[...] Uploading {CFG['sae_weights_path']}...\")\n    api.upload_file(\n        path_or_fileobj=CFG[\"sae_weights_path\"],\n        path_in_repo=\"micro_sae_1024d.pt\",\n        repo_id=repo_id,\n        token=token,\n    )\n    print(\"[OK] Weights uploaded.\")\n\n    # Upload config\n    print(f\"[...] Uploading {CFG['sae_config_path']}...\")\n    api.upload_file(\n        path_or_fileobj=CFG[\"sae_config_path\"],\n        path_in_repo=\"config.json\",\n        repo_id=repo_id,\n        token=token,\n    )\n    print(\"[OK] Config uploaded.\")\n    print(f\"\\nSUCCESS: All artifacts uploaded to: https://huggingface.co/{repo_id}\\n\")\n","metadata":{"trusted":true},"outputs":[],"execution_count":null},{"cell_type":"code","source":"print(\"Micro-SAE Training Pipeline for LLaVA VLM Unlearning\")\nprint(\"=\" * 70)\nprint(f\"   Target model : {CFG['model_id']}\")\nprint(f\"   SAE dims     : {CFG['sae_input_dim']} -> {CFG['sae_dict_size']} -> {CFG['sae_input_dim']}\")\nprint(f\"   Device       : {'CUDA' if torch.cuda.is_available() else 'CPU'}\")\nif torch.cuda.is_available():\n    print(f\"   GPU          : {torch.cuda.get_device_name(0)}\")\n    print(f\"   VRAM         : {torch.cuda.get_device_properties(0).total_memory / 1e9:.1f} GB\")\nprint(\"=\" * 70 + \"\\n\")\n\n# Authenticate\ntoken = step1_authenticate()","metadata":{"trusted":true},"outputs":[],"execution_count":null},{"cell_type":"code","source":"# CLIP Processor\nfrom transformers import AutoProcessor\nprocessor = AutoProcessor.from_pretrained(CFG[\"model_id\"],use_fast=False)\nloader = step2_curate_dataset(processor)\ndel processor  # not needed after this","metadata":{"trusted":true},"outputs":[],"execution_count":null},{"cell_type":"code","source":"# Extract activations (loads model -> extracts -> frees GPU)\nstep3_extract_activations(loader)\ndel loader\ngc.collect()","metadata":{"trusted":true,"execution":{"iopub.status.busy":"2026-03-06T23:39:32.007702Z","iopub.execute_input":"2026-03-06T23:39:32.007982Z","iopub.status.idle":"2026-03-06T23:42:39.357389Z","shell.execute_reply.started":"2026-03-06T23:39:32.007957Z","shell.execute_reply":"2026-03-06T23:42:39.356727Z"}},"outputs":[],"execution_count":null},{"cell_type":"code","source":"# Train SAE on saved activations\nstep4_train_sae()\n","metadata":{"trusted":true},"outputs":[],"execution_count":null},{"cell_type":"code","source":"# Upload to Hugging Face Hub\nstep5_upload_to_hub(token)\n\nprint(\"DONE: Pipeline complete!\")","metadata":{"trusted":true},"outputs":[],"execution_count":null}]}