Team Ai
Datasetpublic

codeShare/lora-training-data

sourceHugging Faceupdated 1mo agoView on Hugging Face
2likes785downloads
prune_lora.ipynb306 linesDownload Raw Back to root
1{2  "nbformat": 4,3  "nbformat_minor": 0,4  "metadata": {5    "colab": {6      "provenance": [],7      "gpuType": "T4"8    },9    "kernelspec": {10      "name": "python3",11      "display_name": "Python 3"12    },13    "language_info": {14      "name": "python"15    },16    "accelerator": "GPU"17  },18  "cells": [19    {20      "cell_type": "markdown",21      "source": [22        "# Process all LoRa safetensors in huggingface repo"23      ],24      "metadata": {25        "id": "AlccmYv2yH5V"26      }27    },28    {29      "cell_type": "code",30      "source": [31        "#@markdown **Cell 1** - Install dependencies\n",32        "!pip install -q safetensors torch matplotlib numpy huggingface_hub\n",33        "\n",34        "#@markdown **Cell 2** - Import libraries + Login with HF_TOKEN\n",35        "import torch\n",36        "import safetensors.torch\n",37        "import matplotlib.pyplot as plt\n",38        "import numpy as np\n",39        "from collections import defaultdict\n",40        "import os\n",41        "import json\n",42        "import shutil\n",43        "from google.colab import drive, files\n",44        "from huggingface_hub import snapshot_download, HfApi\n",45        "import getpass\n",46        "\n",47        "print(\"✅ Libraries imported!\")\n",48        "\n",49        "# Use HF_TOKEN from Colab Secrets for faster downloads\n",50        "from google.colab import userdata\n",51        "hf_token = userdata.get('HF_TOKEN')\n",52        "if hf_token:\n",53        "    print(\"✅ HF_TOKEN loaded from secrets\")\n",54        "else:\n",55        "    print(\"⚠️ HF_TOKEN not found in secrets. You may need to add it.\")\n",56        "\n",57        "#@markdown **Cell 3** - Mount Drive + Create Working Folders\n",58        "drive.mount('/content/drive')\n",59        "\n",60        "base_dir = \"/content/Klein_LoRAs\"\n",61        "os.makedirs(base_dir, exist_ok=True)\n",62        "os.makedirs(f\"{base_dir}/processed\", exist_ok=True)\n",63        "os.makedirs(f\"{base_dir}/metadata\", exist_ok=True)\n",64        "\n",65        "print(f\"Working directory: {base_dir}\")\n"66      ],67      "metadata": {68        "cellView": "form",69        "id": "ZkAw7Xp_zCcY"70      },71      "execution_count": null,72      "outputs": []73    },74    {75      "cell_type": "code",76      "source": [77        "#@markdown **Cell 4** - Download All LoRAs from Hugging Face\n",78        "#repo_id = \"Winnougan/Must_Have_Klein_9b_loras\"\n",79        "\n",80        "repo_id = '' #@param {type:'string'}\n",81        "\n",82        "print(f\"Downloading all files from {repo_id} ... (this may take a while)\")\n",83        "\n",84        "# Download all .safetensors files\n",85        "snapshot_download(\n",86        "    repo_id=repo_id,\n",87        "    local_dir=base_dir,\n",88        "    allow_patterns=\"*.safetensors\",\n",89        "    token=hf_token,\n",90        "    ignore_patterns=[\"*.md\", \"*.txt\", \"*.json\"]  # skip non-lora files if any\n",91        ")\n",92        "\n",93        "# Get list of downloaded safetensors\n",94        "lora_files = sorted([f for f in os.listdir(base_dir) if f.endswith(\".safetensors\")])\n",95        "print(f\"✅ Downloaded {len(lora_files)} LoRA files:\")\n",96        "for i, f in enumerate(lora_files):\n",97        "    print(f\"   {i+1:2d}. {f}\")"98      ],99      "metadata": {100        "cellView": "form",101        "id": "8Y2gIsvRzMLl"102      },103      "execution_count": null,104      "outputs": []105    },106    {107      "cell_type": "code",108      "source": [109        "#@markdown **Cell 5** - Main Processing: SVD Pruning at 90% Accuracy\n",110        "\n",111        "desired_accuracy = 90  #@param {type:'slider', min:80, step:1, max:100}\n",112        "\n",113        "max_allowed_error = 100 - desired_accuracy\n",114        "print(f\"🎯 Target Accuracy: {desired_accuracy}% | Max Error: {max_allowed_error}%\")\n",115        "\n",116        "device = torch.device(\"cuda\" if torch.cuda.is_available() else \"cpu\")\n",117        "print(f\"Using device: {device}\")\n",118        "\n",119        "import random\n",120        "import datetime\n",121        "from safetensors import safe_open\n",122        "from safetensors.torch import save_file\n",123        "\n",124        "processed_files = []\n",125        "mapping = {}  # original -> new letter name\n",126        "\n",127        "letters = [chr(65 + i) for i in range(30)]  # A to Z, then more if needed\n",128        "\n",129        "for idx, filename in enumerate(lora_files):\n",130        "    input_path = os.path.join(base_dir, filename)\n",131        "    letter = letters[idx]\n",132        "    random_suffix = random.randint(10000, 99999)\n",133        "    output_name = f\"{letter}_{random_suffix}.safetensors\"\n",134        "    output_path = f\"{base_dir}/processed/{output_name}\"\n",135        "\n",136        "    print(f\"\\n🔄 Processing [{letter}] {filename} ...\")\n",137        "\n",138        "    # === Load original metadata if it exists ===\n",139        "    original_metadata = {}\n",140        "    try:\n",141        "        with safe_open(input_path, framework=\"pt\", device=\"cpu\") as f:\n",142        "            original_metadata = f.metadata() or {}\n",143        "        if original_metadata:\n",144        "            print(f\"   📋 Found original metadata ({len(original_metadata)} keys)\")\n",145        "        else:\n",146        "            print(f\"   📋 No metadata found in original file\")\n",147        "    except Exception as e:\n",148        "        print(f\"   ⚠️ Could not read metadata: {e}\")\n",149        "\n",150        "    # Load tensors\n",151        "    state_dict = safetensors.torch.load_file(input_path)\n",152        "    lora_pairs = defaultdict(dict)\n",153        "    original_ranks = {}\n",154        "\n",155        "    for key, tensor in state_dict.items():\n",156        "        if \"lora_A.weight\" in key:\n",157        "            base_key = key.replace(\".lora_A.weight\", \"\")\n",158        "            lora_pairs[base_key][\"A\"] = tensor\n",159        "            original_ranks[base_key] = tensor.shape[0]\n",160        "        elif \"lora_B.weight\" in key:\n",161        "            base_key = key.replace(\".lora_B.weight\", \"\")\n",162        "            lora_pairs[base_key][\"B\"] = tensor\n",163        "\n",164        "    new_state_dict = state_dict.copy()\n",165        "\n",166        "    for base_key, matrices in lora_pairs.items():\n",167        "        if \"A\" not in matrices or \"B\" not in matrices:\n",168        "            continue\n",169        "\n",170        "        A = matrices[\"A\"].to(device)\n",171        "        B = matrices[\"B\"].to(device)\n",172        "        orig_rank = original_ranks[base_key]\n",173        "\n",174        "        delta_W = torch.matmul(B, A)\n",175        "        delta_W_fp32 = delta_W.to(torch.float32)\n",176        "        original_norm = torch.norm(delta_W_fp32, p='fro').item()\n",177        "\n",178        "        U, S, Vh = torch.linalg.svd(delta_W_fp32, full_matrices=False)\n",179        "\n",180        "        # Find lowest rank meeting accuracy\n",181        "        best_rank = orig_rank\n",182        "        best_error = 0.0\n",183        "        best_acc = 100.0\n",184        "\n",185        "        candidates = [max(1, int(orig_rank * f)) for f in [0.9,0.8,0.7,0.6,0.5,0.4,0.3,0.2,0.1]]\n",186        "        test_ranks = sorted(set(r for r in candidates if r < orig_rank))\n",187        "\n",188        "        for r in test_ranks:\n",189        "            if r >= len(S):\n",190        "                continue\n",191        "            S_trunc = torch.zeros_like(S)\n",192        "            S_trunc[:r] = S[:r]\n",193        "            approx = U @ torch.diag(S_trunc) @ Vh\n",194        "            error = torch.norm(delta_W_fp32 - approx, p='fro').item()\n",195        "            rel_error = (error / original_norm * 100) if original_norm > 1e-8 else 0.0\n",196        "            acc = 100 - rel_error\n",197        "\n",198        "            if acc >= desired_accuracy:\n",199        "                best_rank = r\n",200        "                best_error = rel_error\n",201        "                best_acc = acc\n",202        "                break\n",203        "\n",204        "        # Reconstruct low-rank A/B\n",205        "        r = best_rank\n",206        "        sqrt_S = torch.sqrt(S[:r])\n",207        "        new_B = (U[:, :r] * sqrt_S.unsqueeze(0)).to(A.dtype)\n",208        "        new_A = (sqrt_S.unsqueeze(1) * Vh[:r, :]).to(A.dtype)\n",209        "\n",210        "        a_key = base_key + \".lora_A.weight\"\n",211        "        b_key = base_key + \".lora_B.weight\"\n",212        "        new_state_dict[a_key] = new_A.contiguous().cpu()\n",213        "        new_state_dict[b_key] = new_B.contiguous().cpu()\n",214        "\n",215        "        print(f\"   {base_key[-45:]:45} | {orig_rank:3}→{best_rank:3} | Acc {best_acc:5.2f}%\")\n",216        "\n",217        "    # === Prepare new metadata ===\n",218        "    new_metadata = dict(original_metadata) if original_metadata else {}\n",219        "    new_metadata.update({\n",220        "        \"pruned_with\": f\"SVD Rank Reduction at {desired_accuracy}% accuracy\",\n",221        "        \"original_file\": filename,\n",222        "        \"pruning_date\": datetime.datetime.now().isoformat(),\n",223        "        \"random_id\": str(random_suffix)\n",224        "    })\n",225        "\n",226        "    # Make all tensors contiguous before saving\n",227        "    for k, v in new_state_dict.items():\n",228        "        if isinstance(v, torch.Tensor):\n",229        "            new_state_dict[k] = v.contiguous()\n",230        "\n",231        "    # Save with metadata\n",232        "    safetensors.torch.save_file(new_state_dict, output_path, metadata=new_metadata)\n",233        "\n",234        "    processed_files.append(output_path)\n",235        "    mapping[filename] = letter\n",236        "\n",237        "    # === Create & Encrypt Metadata Text File ===\n",238        "    meta_dir = f\"{base_dir}/metadata/{letter}\"\n",239        "    os.makedirs(meta_dir, exist_ok=True)\n",240        "    plain_meta_path = f\"{meta_dir}/{letter}_metadata.txt\"\n",241        "    encrypted_meta_path = f\"{meta_dir}/{letter}_metadata.7z\"\n",242        "\n",243        "    # Write plain text first\n",244        "    with open(plain_meta_path, \"w\", encoding=\"utf-8\") as f:\n",245        "        f.write(f\"Original filename: {filename}\\n\")\n",246        "        f.write(f\"Pruned filename : {output_name}\\n\")\n",247        "        f.write(f\"Target Accuracy : {desired_accuracy}%\\n\")\n",248        "        f.write(f\"Pruning Date    : {datetime.datetime.now()}\\n\")\n",249        "        f.write(f\"Random ID       : {random_suffix}\\n\")\n",250        "        f.write(\"\\n--- Original Metadata ---\\n\")\n",251        "        if original_metadata:\n",252        "            for k, v in original_metadata.items():\n",253        "                f.write(f\"{k}: {v}\\n\")\n",254        "        else:\n",255        "            f.write(\"No original metadata found in safetensors file.\\n\")\n",256        "\n",257        "    # Encrypt with password 'banana'\n",258        "    password = \"banana\"  #@param {type:'string'}\n",259        "    !7z a -p{password} -mhe=on \"{encrypted_meta_path}\" \"{plain_meta_path}\" > /dev/null 2>&1\n",260        "\n",261        "    # Optional: Remove plain text after encryption (uncomment if you want only encrypted version)\n",262        "    # os.remove(plain_meta_path)\n",263        "\n",264        "    print(f\"   → Saved as {output_name} (with preserved metadata)\")"265      ],266      "metadata": {267        "cellView": "form",268        "id": "dbt3RipEzPA5"269      },270      "execution_count": null,271      "outputs": []272    },273    {274      "cell_type": "code",275      "metadata": {276        "cellView": "form",277        "id": "249224b7"278      },279      "source": [280        "#@markdown **Cell 6** - Save Processed LoRA to Google Drive\n",281        "\n",282        "#@markdown ---\n",283        "#@markdown **Specify the target folder in your Google Drive (e.g., `/content/drive/MyDrive/MyLoRAs`)**\n",284        "drive_save_path = \"/content/drive/MyDrive/Processed_LoRAs\" #@param {type:\"string\"}\n",285        "\n",286        "import shutil\n",287        "import os\n",288        "\n",289        "# Ensure the target directory exists in Google Drive\n",290        "os.makedirs(drive_save_path, exist_ok=True)\n",291        "\n",292        "print(f\"Saving processed LoRA files to: {drive_save_path}\")\n",293        "\n",294        "for processed_file in processed_files:\n",295        "    file_name = os.path.basename(processed_file)\n",296        "    target_file_path = os.path.join(drive_save_path, file_name)\n",297        "    shutil.copy(processed_file, target_file_path)\n",298        "    print(f\"✅ Copied '{file_name}' to '{target_file_path}'\")\n",299        "\n",300        "print(\"All processed LoRA files saved to Google Drive!\")"301      ],302      "execution_count": null,303      "outputs": []304    }305  ]306}