Team Ai
Datasetpublic

codeShare/lora-training-data

sourceHugging Faceupdated 1mo agoView on Hugging Face
2likes785downloads
train_klein_SDNQ_test.ipynb3812 linesDownload Raw Back to root
1{2  "cells": [3    {4      "cell_type": "code",5      "execution_count": 1,6      "metadata": {7        "id": "OXGhsHQwWue_",8        "colab": {9          "base_uri": "https://localhost:8080/",10          "height": 839,11          "referenced_widgets": [12            "5fb522aff98949e4ae7df3c20ed54703",13            "28d8bb0a08104a25b440c5d5e8a9bd0f",14            "b41a3725067646cf9a3c60ccefc64a97",15            "8a09831ec0e84884b086df43f593ce11",16            "f95662023f5e45c99c4ae704a8c539e9",17            "4c06cd79dd7647658d93120b1bf42e28",18            "ac3e92d8fb1944ff8d93d27092925b35",19            "5f78e9e6854b4ff183f47f4b2a418608",20            "9e103e65b32b4ad0abaa512f97b1ee3c",21            "c078ccfc70c446a8b9ca506483ba1b0d",22            "de393e6b05c14f589e05e2d79968cdfd",23            "9c3333d7b393436da41d607f997a1b87",24            "bd64e1b1949a45b28592e9449acd5d71",25            "19d870948d314da5afc5010e889dc5ef",26            "6fd66e23535e470f81d12cbe17ac42e7",27            "6929567cfb864542b34f87bb17af6999",28            "62b87f3dd7904b8293416dc3f4315517",29            "14d35378a8ea40daabf7582c24141d21",30            "423922fbba2d4c3aa9314ba4331bd0a7",31            "ad61dce7a32f40c38c5c0235620d329b",32            "331f1f508d23434fb5f6cf4bd37975cc",33            "c8684352f6e449168a6cdaddeaa3acfe",34            "491f54937819404d891578bc9cd659db",35            "1da32d981bae4e6c8d4a384734ef39ef",36            "534220ea214d4048a93000e942f7991d",37            "10a62a86e8d04c9784f07ded08919a57",38            "88346c0b0a374449a49911c395a6e001",39            "756a0c170b034a30a14e0c52f299e8af",40            "1c4361237a0940b08d91fe9c36ddb6ff",41            "6b960da08bc64275b283ff2de3fbf864",42            "eb57d6161b194692863969ca77a44f27",43            "0a7a32e457f144b5b247559b5efd1855",44            "ef720746be424a98944112b58cb0f00a",45            "d74d70d64c43401d9636797bb5405ab9",46            "dbeaa864d98645f083939873f2e1c64e",47            "4d162183fe494f5e8589f2efcb3b709c",48            "fed04e5e9f63488f97b780dc2599c0ab",49            "49726df6963e484f80dccd9dad3f561d",50            "b668d6c9da30444ba481764f58a1d613",51            "aa4e6608fffd411ba0897d6106b4f551",52            "1dce4b0d95d1452d982d9b6fb07c0f90",53            "cfdd467912d54b2e8a72059ae9cacdc9",54            "d0848635653141028bb665cf58ba412b",55            "0a0e651f0f7a470fbca8662664c115fc",56            "59b3aee270954902bba4c2bebec11683",57            "1dbdd5a0e7ac47ccbedf8e7926abfe7c",58            "f1738bed749c4414b566a963b2379b21",59            "e5e2932e6e654e44a5b6f6d0b181d288",60            "84a98b60fac441589d63bd86227dc17d",61            "86a0b89247cc41bfa4bc751d2d8eff53",62            "9861f89efbaa432c86ea381968338451",63            "28155ec9896843528718308d0c47a839",64            "7d96df2110f64ef4b572f2c8a57aaa43",65            "6eb73b9d073742af9ddd59f54ff56841",66            "a842e265159542fba4edb2f61a1b6a86"67          ]68        },69        "outputId": "072fe4f3-aab5-4b5d-e9ac-6a6324204fdf"70      },71      "outputs": [72        {73          "output_type": "stream",74          "name": "stdout",75          "text": [76            "Mounted at /content/drive\n",77            "๐Ÿงน Removing old diffusers...\n",78            "๐Ÿ”„ Installing latest diffusers...\n",79            "  Installing build dependencies ... \u001b[?25l\u001b[?25hdone\n",80            "  Getting requirements to build wheel ... \u001b[?25l\u001b[?25hdone\n",81            "  Preparing metadata (pyproject.toml) ... \u001b[?25l\u001b[?25hdone\n",82            "  Building wheel for diffusers (pyproject.toml) ... \u001b[?25l\u001b[?25hdone\n",83            "Files removed: 6\n",84            "โœ… Cell 1 complete!\n",85            "Using zip from Drive: /content/drive/MyDrive/PROCESSED_klein_processed.zip\n",86            "Widget upload is disabled. Using zip from Drive as specified in the previous cell.\n",87            "Final zip_path for processing: /content/drive/MyDrive/PROCESSED_klein_processed.zip\n",88            "โœ… Cell 2 settings loaded\n",89            "   Resolution: 1024ร—1024\n",90            "   Use .txt prompts: False\n",91            "   Debug mode: False\n",92            "\n",93            "Now run Cell 3 (model load), then Cell 4 (inference)\n"94          ]95        },96        {97          "output_type": "stream",98          "name": "stderr",99          "text": [100            "Flax classes are deprecated and will be removed in Diffusers v1.0.0. We recommend migrating to PyTorch classes or pinning your version of Diffusers.\n",101            "Flax classes are deprecated and will be removed in Diffusers v1.0.0. We recommend migrating to PyTorch classes or pinning your version of Diffusers.\n"102          ]103        },104        {105          "output_type": "stream",106          "name": "stdout",107          "text": [108            "๐Ÿ“ฆ Installing SDNQ...\n",109            "\u001b[2K   \u001b[90mโ”โ”โ”โ”โ”โ”โ”โ”โ”โ”โ”โ”โ”โ”โ”โ”โ”โ”โ”โ”โ”โ”โ”โ”โ”โ”โ”โ”โ”โ”โ”โ”โ”โ”โ”โ”โ”โ”โ”โ”\u001b[0m \u001b[32m104.0/104.0 kB\u001b[0m \u001b[31m5.2 MB/s\u001b[0m eta \u001b[36m0:00:00\u001b[0m\n",110            "\u001b[2K   \u001b[90mโ”โ”โ”โ”โ”โ”โ”โ”โ”โ”โ”โ”โ”โ”โ”โ”โ”โ”โ”โ”โ”โ”โ”โ”โ”โ”โ”โ”โ”โ”โ”โ”โ”โ”โ”โ”โ”โ”โ”โ”\u001b[0m \u001b[32m509.1/509.1 kB\u001b[0m \u001b[31m27.5 MB/s\u001b[0m eta \u001b[36m0:00:00\u001b[0m\n",111            "\u001b[?25h๐Ÿ”„ Loading model from: codeShare/FLUX.2-klein-AIO-SDNQ-4bit-dynamic\n"112          ]113        },114        {115          "output_type": "display_data",116          "data": {117            "text/plain": [118              "model_index.json:   0%|          | 0.00/509 [00:00<?, ?B/s]"119            ],120            "application/vnd.jupyter.widget-view+json": {121              "version_major": 2,122              "version_minor": 0,123              "model_id": "5fb522aff98949e4ae7df3c20ed54703"124            }125          },126          "metadata": {}127        },128        {129          "output_type": "display_data",130          "data": {131            "text/plain": [132              "Downloading (incomplete total...): 0.00B [00:00, ?B/s]"133            ],134            "application/vnd.jupyter.widget-view+json": {135              "version_major": 2,136              "version_minor": 0,137              "model_id": "9c3333d7b393436da41d607f997a1b87"138            }139          },140          "metadata": {}141        },142        {143          "output_type": "display_data",144          "data": {145            "text/plain": [146              "Fetching 11 files:   0%|          | 0/11 [00:00<?, ?it/s]"147            ],148            "application/vnd.jupyter.widget-view+json": {149              "version_major": 2,150              "version_minor": 0,151              "model_id": "491f54937819404d891578bc9cd659db"152            }153          },154          "metadata": {}155        },156        {157          "output_type": "display_data",158          "data": {159            "text/plain": [160              "Loading pipeline components...:   0%|          | 0/5 [00:00<?, ?it/s]"161            ],162            "application/vnd.jupyter.widget-view+json": {163              "version_major": 2,164              "version_minor": 0,165              "model_id": "d74d70d64c43401d9636797bb5405ab9"166            }167          },168          "metadata": {}169        },170        {171          "output_type": "display_data",172          "data": {173            "text/plain": [174              "Loading weights:   0%|          | 0/901 [00:00<?, ?it/s]"175            ],176            "application/vnd.jupyter.widget-view+json": {177              "version_major": 2,178              "version_minor": 0,179              "model_id": "59b3aee270954902bba4c2bebec11683"180            }181          },182          "metadata": {}183        },184        {185          "output_type": "stream",186          "name": "stdout",187          "text": [188            "โœ… Base pipeline loaded on CPU\n",189            "๐Ÿ”ค Encoding fixed prompt on GPU: 'remove the white background. the background is dark gray. '\n",190            "โœ… Embedding computed on GPU. Shape: torch.Size([1, 512, 7680])\n",191            "โœ… Original text_encoder fully unloaded and replaced with fixed embedding\n",192            "   โœ… SDNQ quantized matmul applied to transformer\n",193            "โœ… CELL 3B COMPLETE - Fixed embedding computed on GPU\n",194            "VRAM usage: 0.01 GB\n",195            "๐Ÿ‘‰ Original text_encoder has been unloaded.\n",196            "๐Ÿ”„ Converting transformer to SDNQ training mode...\n",197            "โœ… CELL 4B COMPLETE - Transformer converted to SDNQ training model\n",198            "VRAM: 0.01 GB\n",199            "Model is ready for training with fixed edit prompt embedding.\n",200            "๐Ÿงน Cleaning up before training...\n"201          ]202        }203      ],204      "source": [205        "# =============================================================================\n",206        "#@markdown # **CELL 1**: Mount Drive + HF auth\n",207        "# =============================================================================\n",208        "\n",209        "from google.colab import drive, userdata\n",210        "from huggingface_hub import login\n",211        "import torch\n",212        "import os\n",213        "import gc\n",214        "import shutil\n",215        "\n",216        "drive.mount('/content/drive')\n",217        "\n",218        "hf_token = userdata.get('HF_TOKEN')\n",219        "if hf_token:\n",220        "    login(token=hf_token)\n",221        "else:\n",222        "    print(\"โš ๏ธ No HF_TOKEN found in secrets.\")\n",223        "\n",224        "print(\"๐Ÿงน Removing old diffusers...\")\n",225        "!pip uninstall -y diffusers > /dev/null 2>&1\n",226        "!rm -rf /usr/local/lib/python3.12/dist-packages/diffusers* ~/.cache/pip/*diffusers*\n",227        "\n",228        "print(\"๐Ÿ”„ Installing latest diffusers...\")\n",229        "!pip install -q git+https://github.com/huggingface/diffusers.git --force-reinstall --no-deps\n",230        "!python -m pip cache purge\n",231        "\n",232        "print(\"โœ… Cell 1 complete!\")\n",233        "\n",234        "#-----#\n",235        "\n",236        "#@title Set path to zip file on your drive , and set klein edit prompt\n",237        "upload_from_widget = False #@param {type:'boolean'}\n",238        "input_zip_path = '/content/drive/MyDrive/PROCESSED_klein_processed.zip' #@param {type:'string'}\n",239        "edit_prompt = 'remove the white background. the background is dark gray. ' #@param {type:'string'}\n",240        "\n",241        "# Initialize zip_path; it will be set definitively by this cell or the next.\n",242        "zip_path = None\n",243        "\n",244        "if not upload_from_widget:\n",245        "    zip_path = input_zip_path\n",246        "    print(f\"Using zip from Drive: {zip_path}\")\n",247        "else:\n",248        "    print(\"Widget upload enabled. Please run the next cell to upload your files.\")\n",249        "\n",250        "#---#\n",251        "#@title Upload files via widget\n",252        "import os\n",253        "from google.colab import files\n",254        "import zipfile\n",255        "import shutil\n",256        "\n",257        "if upload_from_widget:\n",258        "    print(\"Please upload your zip file or individual image files now.\")\n",259        "    uploaded = files.upload()\n",260        "\n",261        "    if not uploaded:\n",262        "        raise ValueError(\"No files uploaded. Please upload a zip file or images.\")\n",263        "\n",264        "    if len(uploaded) == 1 and list(uploaded.keys())[0].endswith('.zip'):\n",265        "        # If a single zip file is uploaded\n",266        "        uploaded_filename = list(uploaded.keys())[0]\n",267        "        shutil.move(uploaded_filename, '/content/' + uploaded_filename)\n",268        "        zip_path = '/content/' + uploaded_filename\n",269        "        print(f\"Using uploaded zip file: {zip_path}\")\n",270        "    else:\n",271        "        # If multiple image files are uploaded, create a zip file\n",272        "        temp_img_dir = '/content/uploaded_images_temp'\n",273        "        os.makedirs(temp_img_dir, exist_ok=True)\n",274        "\n",275        "        image_count = 0\n",276        "        for fname, content in uploaded.items():\n",277        "            if fname.lower().endswith(('.png', '.jpg', '.jpeg', '.webp', '.gif')):\n",278        "                with open(os.path.join(temp_img_dir, fname), 'wb') as f:\n",279        "                    f.write(content)\n",280        "                image_count += 1\n",281        "            else:\n",282        "                print(f\"Skipping non-image file: {fname}\")\n",283        "\n",284        "        if image_count == 0:\n",285        "            raise ValueError(\"No valid image files uploaded. Please upload images.\")\n",286        "\n",287        "        zip_path = '/content/uploaded_images.zip'\n",288        "        with zipfile.ZipFile(zip_path, 'w') as zf:\n",289        "            for root, _, files_in_dir in os.walk(temp_img_dir):\n",290        "                for file_in_dir in files_in_dir:\n",291        "                    zf.write(os.path.join(root, file_in_dir), os.path.basename(file_in_dir))\n",292        "        shutil.rmtree(temp_img_dir)\n",293        "        print(f\"Created zip from {image_count} uploaded images: {zip_path}\")\n",294        "else:\n",295        "    print(\"Widget upload is disabled. Using zip from Drive as specified in the previous cell.\")\n",296        "\n",297        "print(f\"Final zip_path for processing: {zip_path}\")\n",298        "\n",299        "#-----#\n",300        "\n",301        "# =============================================================================\n",302        "#@markdown # **CELL 2**: Fixed settings (resolution + options)\n",303        "# =============================================================================\n",304        "\n",305        "resolution = '1024 x 1024 (Square)' #@param [\"1024 x 1024 (Square)\", \"512 x 1024 (Portrait)\", \"768 x 1024 (Slight Portrait)\", \"1536 x 1024 (Landscape)\", \"2048 x 1024 (Wide Landscape)\"] {type:\"string\"}\n",306        "use_txt_prompts = False #@param {type:\"boolean\"}\n",307        "debug = False #@param {type:\"boolean\"}\n",308        "\n",309        "res_map = {\n",310        "    \"1024 x 1024 (Square)\": (1024, 1024),\n",311        "    \"512 x 1024 (Portrait)\": (512, 1024),\n",312        "    \"768 x 1024 (Slight Portrait)\": (768, 1024),\n",313        "    \"1536 x 1024 (Landscape)\": (1536, 1024),\n",314        "    \"2048 x 1024 (Wide Landscape)\": (2048, 1024)\n",315        "}\n",316        "target_width, target_height = res_map[resolution]\n",317        "\n",318        "print(\"โœ… Cell 2 settings loaded\")\n",319        "print(f\"   Resolution: {target_width}ร—{target_height}\")\n",320        "print(f\"   Use .txt prompts: {use_txt_prompts}\")\n",321        "print(f\"   Debug mode: {debug}\")\n",322        "print(\"\\nNow run Cell 3 (model load), then Cell 4 (inference)\")\n",323        "\n",324        "#----#\n",325        "\n",326        "# =============================================================================\n",327        "#@markdown # **CELL 3B**: Load SDNQ MODEL + Fixed Text Encoder (GPU Embedding)\n",328        "# =============================================================================\n",329        "\n",330        "import torch\n",331        "import gc\n",332        "import os\n",333        "from diffusers import Flux2KleinPipeline\n",334        "\n",335        "print(\"๐Ÿ“ฆ Installing SDNQ...\")\n",336        "!pip install -q sdnq\n",337        "\n",338        "from sdnq.common import use_torch_compile as triton_is_available\n",339        "from sdnq.loader import apply_sdnq_options_to_model\n",340        "\n",341        "gc.collect()\n",342        "torch.cuda.empty_cache()\n",343        "\n",344        "# =========================================================\n",345        "# ๐Ÿ”ฅ LOAD MODEL ON CPU FIRST (Safe loading)\n",346        "# =========================================================\n",347        "MODEL_ID = \"codeShare/FLUX.2-klein-AIO-SDNQ-4bit-dynamic\" #@param ['codeShare/Flux-Klein-SDNQ-4bit','codeShare/FLUX.2-klein-AIO-SDNQ-4bit-dynamic', 'codeShare/unstableRevolution_SDNQ']\n",348        "\n",349        "print(f\"๐Ÿ”„ Loading model from: {MODEL_ID}\")\n",350        "\n",351        "pipe = Flux2KleinPipeline.from_pretrained(\n",352        "    MODEL_ID,\n",353        "    torch_dtype=torch.float16,\n",354        "    low_cpu_mem_usage=True,\n",355        "    device_map=\"cpu\"\n",356        ")\n",357        "\n",358        "print(\"โœ… Base pipeline loaded on CPU\")\n",359        "\n",360        "# =========================================================\n",361        "# ๐Ÿ”ฅ COMPUTE FIXED EMBEDDING ON GPU\n",362        "# =========================================================\n",363        "edit_prompt = 'remove the white background. the background is dark gray. '\n",364        "\n",365        "print(f\"๐Ÿ”ค Encoding fixed prompt on GPU: '{edit_prompt}'\")\n",366        "\n",367        "# Temporarily move text_encoder to GPU for faster embedding computation\n",368        "pipe.text_encoder = pipe.text_encoder.to(\"cuda\")\n",369        "\n",370        "with torch.no_grad():\n",371        "    prompt_embeds, pooled_prompt_embeds = pipe.encode_prompt(\n",372        "        prompt=edit_prompt,\n",373        "        num_images_per_prompt=1,\n",374        "        device=\"cuda\",\n",375        "        max_sequence_length=512,\n",376        "    )\n",377        "\n",378        "# Move embedding back to CPU for storage (lower memory)\n",379        "prompt_embeds = prompt_embeds.cpu()\n",380        "if pooled_prompt_embeds is not None:\n",381        "    pooled_prompt_embeds = pooled_prompt_embeds.cpu()\n",382        "\n",383        "print(\"โœ… Embedding computed on GPU. Shape:\", prompt_embeds.shape)\n",384        "\n",385        "# =========================================================\n",386        "# ๐Ÿ”ฅ UNLOAD ORIGINAL TEXT ENCODER + REPLACE WITH FIXED ONE\n",387        "# =========================================================\n",388        "# Delete original text encoder to free VRAM\n",389        "if hasattr(pipe, 'text_encoder'):\n",390        "    del pipe.text_encoder\n",391        "    gc.collect()\n",392        "    torch.cuda.empty_cache()\n",393        "\n",394        "class FixedEmbeddingEncoder:\n",395        "    def __init__(self, prompt_embeds, pooled=None):\n",396        "        self.prompt_embeds = prompt_embeds\n",397        "        self.pooled_prompt_embeds = pooled\n",398        "\n",399        "    def __call__(self, *args, **kwargs):\n",400        "        # Return object that mimics what the pipeline expects\n",401        "        class DummyOutput:\n",402        "            last_hidden_state = self.prompt_embeds\n",403        "            # Add pooled_prompt_embeds if pipeline accesses it\n",404        "            pooled_prompt_embeds = self.pooled_prompt_embeds\n",405        "        return DummyOutput()\n",406        "\n",407        "    def to(self, device=None, *args, **kwargs):\n",408        "        return self\n",409        "    def cuda(self): return self\n",410        "    def cpu(self): return self\n",411        "\n",412        "# Replace with fixed embedding\n",413        "pipe.text_encoder = FixedEmbeddingEncoder(prompt_embeds, pooled_prompt_embeds)\n",414        "\n",415        "# Clean tokenizer\n",416        "if hasattr(pipe, \"tokenizer\"):\n",417        "    pipe.tokenizer = None\n",418        "\n",419        "print(\"โœ… Original text_encoder fully unloaded and replaced with fixed embedding\")\n",420        "\n",421        "# =========================================================\n",422        "# ๐Ÿ”ฅ SDNQ OPTIMIZATIONS + MEMORY SETTINGS\n",423        "# =========================================================\n",424        "if torch.cuda.is_available() and triton_is_available:\n",425        "    pipe.transformer = apply_sdnq_options_to_model(\n",426        "        pipe.transformer,\n",427        "        use_quantized_matmul=True\n",428        "    )\n",429        "    print(\"   โœ… SDNQ quantized matmul applied to transformer\")\n",430        "\n",431        "# Enable memory optimizations\n",432        "pipe.enable_model_cpu_offload()\n",433        "pipe.vae.enable_slicing()\n",434        "pipe.vae.enable_tiling()\n",435        "\n",436        "gc.collect()\n",437        "torch.cuda.empty_cache()\n",438        "\n",439        "print(\"โœ… CELL 3B COMPLETE - Fixed embedding computed on GPU\")\n",440        "print(\"VRAM usage:\", round(torch.cuda.memory_allocated() / 1e9, 2), \"GB\")\n",441        "print(\"๐Ÿ‘‰ Original text_encoder has been unloaded.\")\n",442        "\n",443        "#-----#\n",444        "\n",445        "# =============================================================================\n",446        "#@markdown # **CELL 4B**: Convert to SDNQ Quantized Training Model\n",447        "# =============================================================================\n",448        "\n",449        "from sdnq.training import sdnq_training_post_load_quant, convert_sdnq_model_to_training\n",450        "from sdnq.common import use_torch_compile as triton_is_available\n",451        "import torch\n",452        "import gc\n",453        "\n",454        "gc.collect()\n",455        "torch.cuda.empty_cache()\n",456        "\n",457        "print(\"๐Ÿ”„ Converting transformer to SDNQ training mode...\")\n",458        "\n",459        "# Option 1: If the loaded model is already SDNQ-quantized, use convert_...\n",460        "quantized_model = convert_sdnq_model_to_training(\n",461        "    pipe.transformer,                    # Usually the heaviest part\n",462        "    quantized_matmul_dtype=\"int8\",\n",463        "    use_grad_ckpt=True,                  # Recommended for training\n",464        "    use_quantized_matmul=triton_is_available,\n",465        "    use_stochastic_rounding=True,\n",466        "    dequantize_fp32=True,\n",467        ")\n",468        "\n",469        "# Option 2: Alternative - post-load quant (use one or the other)\n",470        "# quantized_model = sdnq_training_post_load_quant(\n",471        "#     pipe.transformer,\n",472        "#     weights_dtype=\"uint8\",\n",473        "#     quantized_matmul_dtype=\"int8\",\n",474        "#     group_size=32,\n",475        "#     svd_rank=32,\n",476        "#     svd_steps=8,\n",477        "#     use_svd=False,\n",478        "#     use_grad_ckpt=True,\n",479        "#     use_quantized_matmul=triton_is_available,\n",480        "#     use_static_quantization=True,\n",481        "#     use_stochastic_rounding=True,\n",482        "#     dequantize_fp32=True,\n",483        "#     non_blocking=False,\n",484        "#     add_skip_keys=True,\n",485        "#     quantization_device=\"cuda\",\n",486        "#     return_device=\"cuda\",\n",487        "#     modules_to_not_convert=[\"correction_coefs\", \"prediction_coefs\", \"lm_head\", \"embedding_projection\"],\n",488        "# )\n",489        "\n",490        "# Replace in pipeline\n",491        "pipe.transformer = quantized_model\n",492        "\n",493        "# Move to training mode\n",494        "pipe.transformer.train()\n",495        "\n",496        "# Optional: quantized optimizer later when you set up training loop\n",497        "# from sdnq.optim import AdamW\n",498        "# optimizer = AdamW(pipe.transformer.parameters(), ...)\n",499        "\n",500        "gc.collect()\n",501        "torch.cuda.empty_cache()\n",502        "\n",503        "print(\"โœ… CELL 4B COMPLETE - Transformer converted to SDNQ training model\")\n",504        "print(\"VRAM:\", round(torch.cuda.memory_allocated() / 1e9, 2), \"GB\")\n",505        "print(\"Model is ready for training with fixed edit prompt embedding.\")\n",506        "\n",507        "#----#\n",508        "\n",509        "# ------------------- CLEANUP -------------------\n",510        "print(\"๐Ÿงน Cleaning up before training...\")\n",511        "\n",512        "gc.collect()\n",513        "torch.cuda.empty_cache()\n",514        "torch.cuda.reset_peak_memory_stats()\n",515        "\n"516      ]517    },518    {519      "cell_type": "code",520      "source": [521        "# =============================================================================\n",522        "# ๐Ÿ”ฅ CELL 5B: SINGLE IMAGE TRAINING (FINAL FINAL FIX)\n",523        "# =============================================================================\n",524        "\n",525        "import torch\n",526        "import torch.nn.functional as F\n",527        "from torchvision import transforms\n",528        "from PIL import Image\n",529        "import gc\n",530        "from tqdm import tqdm\n",531        "import os\n",532        "\n",533        "from sdnq.optim import AdamW\n",534        "\n",535        "device = \"cuda\"\n",536        "\n",537        "# =========================================================\n",538        "# ๐Ÿ“ท LOAD IMAGE\n",539        "# =========================================================\n",540        "\n",541        "image_path = \"/content/drive/MyDrive/test_image.jpg\"\n",542        "\n",543        "image = Image.open(image_path).convert(\"RGB\")\n",544        "\n",545        "transform = transforms.Compose([\n",546        "    transforms.Resize((target_height, target_width)),\n",547        "    transforms.CenterCrop((target_height, target_width)),\n",548        "    transforms.ToTensor(),\n",549        "    transforms.Normalize([0.5], [0.5])\n",550        "])\n",551        "\n",552        "image = transform(image).unsqueeze(0).to(device, dtype=torch.float16)\n",553        "print(\"๐Ÿ“ท Loaded training image:\", image.shape)\n",554        "\n",555        "# =========================================================\n",556        "# โš™๏ธ SETTINGS\n",557        "# =========================================================\n",558        "\n",559        "steps = 1000\n",560        "lr = 1e-5\n",561        "save_every = 200\n",562        "\n",563        "# =========================================================\n",564        "# ๐Ÿ”ง OPTIMIZER\n",565        "# =========================================================\n",566        "\n",567        "optimizer = AdamW(pipe.transformer.parameters(), lr=lr)\n",568        "pipe.transformer.train()\n",569        "\n",570        "prompt_embeds_gpu = prompt_embeds.to(device)\n",571        "\n",572        "# =========================================================\n",573        "# ๐Ÿ”ง BUILD STATIC txt_ids (once)\n",574        "# =========================================================\n",575        "\n",576        "seq_len = prompt_embeds_gpu.shape[1]\n",577        "txt_ids = torch.arange(seq_len, device=device).unsqueeze(0)  # [1, seq_len]\n",578        "\n",579        "# =========================================================\n",580        "# ๐Ÿ”ฅ TRAINING LOOP\n",581        "# =========================================================\n",582        "\n",583        "pbar = tqdm(range(steps))\n",584        "\n",585        "for step in pbar:\n",586        "\n",587        "    # ----------------------------------------\n",588        "    # ๐Ÿ”น Encode โ†’ latents\n",589        "    # ----------------------------------------\n",590        "    with torch.no_grad():\n",591        "        latents = pipe.vae.encode(image).latent_dist.sample()\n",592        "\n",593        "    B, C, H, W = latents.shape\n",594        "\n",595        "    # ----------------------------------------\n",596        "    # ๐Ÿ”น Build img_ids (CRITICAL FIX)\n",597        "    # ----------------------------------------\n",598        "    # create 2D grid of positions\n",599        "    y, x = torch.meshgrid(\n",600        "        torch.arange(H, device=device),\n",601        "        torch.arange(W, device=device),\n",602        "        indexing=\"ij\"\n",603        "    )\n",604        "\n",605        "    img_ids = torch.stack([y, x], dim=-1)  # [H, W, 2]\n",606        "    img_ids = img_ids.view(1, H * W, 2)    # [1, tokens, 2]\n",607        "\n",608        "    # ----------------------------------------\n",609        "    # ๐Ÿ”น Sample noise + t\n",610        "    # ----------------------------------------\n",611        "    noise = torch.randn_like(latents)\n",612        "\n",613        "    t = torch.rand((1,), device=device, dtype=torch.float16)\n",614        "    t_broadcast = t.view(1, 1, 1, 1)\n",615        "\n",616        "    # Flow interpolation\n",617        "    noisy_latents = (1 - t_broadcast) * latents + t_broadcast * noise\n",618        "    target = noise - latents\n",619        "\n",620        "    # ----------------------------------------\n",621        "    # ๐Ÿ”น Forward\n",622        "    # ----------------------------------------\n",623        "    model_pred = pipe.transformer(\n",624        "        hidden_states=noisy_latents,\n",625        "        timestep=t,\n",626        "        encoder_hidden_states=prompt_embeds_gpu,\n",627        "        img_ids=img_ids,\n",628        "        txt_ids=txt_ids,\n",629        "        return_dict=False\n",630        "    )[0]\n",631        "\n",632        "    # ----------------------------------------\n",633        "    # ๐Ÿ”น LOSS\n",634        "    # ----------------------------------------\n",635        "    loss = F.mse_loss(model_pred.float(), target.float(), reduction=\"mean\")\n",636        "\n",637        "    loss.backward()\n",638        "\n",639        "    # ----------------------------------------\n",640        "    # ๐Ÿ”น STEP\n",641        "    # ----------------------------------------\n",642        "    optimizer.step()\n",643        "    optimizer.zero_grad(set_to_none=True)\n",644        "\n",645        "    torch.nn.utils.clip_grad_norm_(pipe.transformer.parameters(), 1.0)\n",646        "\n",647        "    # ----------------------------------------\n",648        "    # ๐Ÿ”น LOGGING\n",649        "    # ----------------------------------------\n",650        "    pbar.set_postfix(loss=loss.item())\n",651        "\n",652        "    # ----------------------------------------\n",653        "    # ๐Ÿ’พ SAVE\n",654        "    # ----------------------------------------\n",655        "    if step % save_every == 0 and step > 0:\n",656        "        save_path = f\"/content/drive/MyDrive/sdnq_single_image_{step}\"\n",657        "        os.makedirs(save_path, exist_ok=True)\n",658        "\n",659        "        pipe.transformer.save_pretrained(save_path)\n",660        "        print(f\"\\n๐Ÿ’พ Saved checkpoint at step {step}\")\n",661        "\n",662        "    # ----------------------------------------\n",663        "    # ๐Ÿงน CLEANUP\n",664        "    # ----------------------------------------\n",665        "    if step % 20 == 0:\n",666        "        gc.collect()\n",667        "        torch.cuda.empty_cache()\n",668        "\n",669        "print(\"โœ… TRAINING COMPLETE\")"670      ],671      "metadata": {672        "colab": {673          "base_uri": "https://localhost:8080/",674          "height": 391675        },676        "id": "2mK44niN52a3",677        "outputId": "9a38b0be-f00e-4d43-9ca8-bb49716d027e"678      },679      "execution_count": 6,680      "outputs": [681        {682          "output_type": "stream",683          "name": "stdout",684          "text": [685            "๐Ÿ“ท Loaded training image: torch.Size([1, 3, 1024, 1024])\n"686          ]687        },688        {689          "output_type": "stream",690          "name": "stderr",691          "text": [692            "  0%|          | 0/1000 [00:02<?, ?it/s]\n"693          ]694        },695        {696          "output_type": "error",697          "ename": "IndexError",698          "evalue": "index 2 is out of bounds for dimension 1 with size 2",699          "traceback": [700            "\u001b[0;31m---------------------------------------------------------------------------\u001b[0m",701            "\u001b[0;31mIndexError\u001b[0m                                Traceback (most recent call last)",702            "\u001b[0;32m/tmp/ipykernel_8575/1505866536.py\u001b[0m in \u001b[0;36m<cell line: 0>\u001b[0;34m()\u001b[0m\n\u001b[1;32m    101\u001b[0m     \u001b[0;31m# ๐Ÿ”น Forward\u001b[0m\u001b[0;34m\u001b[0m\u001b[0;34m\u001b[0m\u001b[0m\n\u001b[1;32m    102\u001b[0m     \u001b[0;31m# ----------------------------------------\u001b[0m\u001b[0;34m\u001b[0m\u001b[0;34m\u001b[0m\u001b[0m\n\u001b[0;32m--> 103\u001b[0;31m     model_pred = pipe.transformer(\n\u001b[0m\u001b[1;32m    104\u001b[0m         \u001b[0mhidden_states\u001b[0m\u001b[0;34m=\u001b[0m\u001b[0mnoisy_latents\u001b[0m\u001b[0;34m,\u001b[0m\u001b[0;34m\u001b[0m\u001b[0;34m\u001b[0m\u001b[0m\n\u001b[1;32m    105\u001b[0m         \u001b[0mtimestep\u001b[0m\u001b[0;34m=\u001b[0m\u001b[0mt\u001b[0m\u001b[0;34m,\u001b[0m\u001b[0;34m\u001b[0m\u001b[0;34m\u001b[0m\u001b[0m\n",703            "\u001b[0;32m/usr/local/lib/python3.12/dist-packages/torch/nn/modules/module.py\u001b[0m in \u001b[0;36m_wrapped_call_impl\u001b[0;34m(self, *args, **kwargs)\u001b[0m\n\u001b[1;32m   1774\u001b[0m             \u001b[0;32mreturn\u001b[0m \u001b[0mself\u001b[0m\u001b[0;34m.\u001b[0m\u001b[0m_compiled_call_impl\u001b[0m\u001b[0;34m(\u001b[0m\u001b[0;34m*\u001b[0m\u001b[0margs\u001b[0m\u001b[0;34m,\u001b[0m \u001b[0;34m**\u001b[0m\u001b[0mkwargs\u001b[0m\u001b[0;34m)\u001b[0m  \u001b[0;31m# type: ignore[misc]\u001b[0m\u001b[0;34m\u001b[0m\u001b[0;34m\u001b[0m\u001b[0m\n\u001b[1;32m   1775\u001b[0m         \u001b[0;32melse\u001b[0m\u001b[0;34m:\u001b[0m\u001b[0;34m\u001b[0m\u001b[0;34m\u001b[0m\u001b[0m\n\u001b[0;32m-> 1776\u001b[0;31m             \u001b[0;32mreturn\u001b[0m \u001b[0mself\u001b[0m\u001b[0;34m.\u001b[0m\u001b[0m_call_impl\u001b[0m\u001b[0;34m(\u001b[0m\u001b[0;34m*\u001b[0m\u001b[0margs\u001b[0m\u001b[0;34m,\u001b[0m \u001b[0;34m**\u001b[0m\u001b[0mkwargs\u001b[0m\u001b[0;34m)\u001b[0m\u001b[0;34m\u001b[0m\u001b[0;34m\u001b[0m\u001b[0m\n\u001b[0m\u001b[1;32m   1777\u001b[0m \u001b[0;34m\u001b[0m\u001b[0m\n\u001b[1;32m   1778\u001b[0m     \u001b[0;31m# torchrec tests the code consistency with the following code\u001b[0m\u001b[0;34m\u001b[0m\u001b[0;34m\u001b[0m\u001b[0m\n",704            "\u001b[0;32m/usr/local/lib/python3.12/dist-packages/torch/nn/modules/module.py\u001b[0m in \u001b[0;36m_call_impl\u001b[0;34m(self, *args, **kwargs)\u001b[0m\n\u001b[1;32m   1785\u001b[0m                 \u001b[0;32mor\u001b[0m \u001b[0m_global_backward_pre_hooks\u001b[0m \u001b[0;32mor\u001b[0m \u001b[0m_global_backward_hooks\u001b[0m\u001b[0;34m\u001b[0m\u001b[0;34m\u001b[0m\u001b[0m\n\u001b[1;32m   1786\u001b[0m                 or _global_forward_hooks or _global_forward_pre_hooks):\n\u001b[0;32m-> 1787\u001b[0;31m             \u001b[0;32mreturn\u001b[0m \u001b[0mforward_call\u001b[0m\u001b[0;34m(\u001b[0m\u001b[0;34m*\u001b[0m\u001b[0margs\u001b[0m\u001b[0;34m,\u001b[0m \u001b[0;34m**\u001b[0m\u001b[0mkwargs\u001b[0m\u001b[0;34m)\u001b[0m\u001b[0;34m\u001b[0m\u001b[0;34m\u001b[0m\u001b[0m\n\u001b[0m\u001b[1;32m   1788\u001b[0m \u001b[0;34m\u001b[0m\u001b[0m\n\u001b[1;32m   1789\u001b[0m         \u001b[0mresult\u001b[0m \u001b[0;34m=\u001b[0m \u001b[0;32mNone\u001b[0m\u001b[0;34m\u001b[0m\u001b[0;34m\u001b[0m\u001b[0m\n",705            "\u001b[0;32m/usr/local/lib/python3.12/dist-packages/accelerate/hooks.py\u001b[0m in \u001b[0;36mnew_forward\u001b[0;34m(module, *args, **kwargs)\u001b[0m\n\u001b[1;32m    190\u001b[0m                 \u001b[0moutput\u001b[0m \u001b[0;34m=\u001b[0m \u001b[0mmodule\u001b[0m\u001b[0;34m.\u001b[0m\u001b[0m_old_forward\u001b[0m\u001b[0;34m(\u001b[0m\u001b[0;34m*\u001b[0m\u001b[0margs\u001b[0m\u001b[0;34m,\u001b[0m \u001b[0;34m**\u001b[0m\u001b[0mkwargs\u001b[0m\u001b[0;34m)\u001b[0m\u001b[0;34m\u001b[0m\u001b[0;34m\u001b[0m\u001b[0m\n\u001b[1;32m    191\u001b[0m         \u001b[0;32melse\u001b[0m\u001b[0;34m:\u001b[0m\u001b[0;34m\u001b[0m\u001b[0;34m\u001b[0m\u001b[0m\n\u001b[0;32m--> 192\u001b[0;31m             \u001b[0moutput\u001b[0m \u001b[0;34m=\u001b[0m \u001b[0mmodule\u001b[0m\u001b[0;34m.\u001b[0m\u001b[0m_old_forward\u001b[0m\u001b[0;34m(\u001b[0m\u001b[0;34m*\u001b[0m\u001b[0margs\u001b[0m\u001b[0;34m,\u001b[0m \u001b[0;34m**\u001b[0m\u001b[0mkwargs\u001b[0m\u001b[0;34m)\u001b[0m\u001b[0;34m\u001b[0m\u001b[0;34m\u001b[0m\u001b[0m\n\u001b[0m\u001b[1;32m    193\u001b[0m         \u001b[0;32mreturn\u001b[0m \u001b[0mmodule\u001b[0m\u001b[0;34m.\u001b[0m\u001b[0m_hf_hook\u001b[0m\u001b[0;34m.\u001b[0m\u001b[0mpost_forward\u001b[0m\u001b[0;34m(\u001b[0m\u001b[0mmodule\u001b[0m\u001b[0;34m,\u001b[0m \u001b[0moutput\u001b[0m\u001b[0;34m)\u001b[0m\u001b[0;34m\u001b[0m\u001b[0;34m\u001b[0m\u001b[0m\n\u001b[1;32m    194\u001b[0m \u001b[0;34m\u001b[0m\u001b[0m\n",706            "\u001b[0;32m/usr/local/lib/python3.12/dist-packages/diffusers/utils/peft_utils.py\u001b[0m in \u001b[0;36mwrapper\u001b[0;34m(self, *args, **kwargs)\u001b[0m\n\u001b[1;32m    313\u001b[0m             \u001b[0;32mtry\u001b[0m\u001b[0;34m:\u001b[0m\u001b[0;34m\u001b[0m\u001b[0;34m\u001b[0m\u001b[0m\n\u001b[1;32m    314\u001b[0m                 \u001b[0;31m# Execute the forward pass\u001b[0m\u001b[0;34m\u001b[0m\u001b[0;34m\u001b[0m\u001b[0m\n\u001b[0;32m--> 315\u001b[0;31m                 \u001b[0mresult\u001b[0m \u001b[0;34m=\u001b[0m \u001b[0mforward_fn\u001b[0m\u001b[0;34m(\u001b[0m\u001b[0mself\u001b[0m\u001b[0;34m,\u001b[0m \u001b[0;34m*\u001b[0m\u001b[0margs\u001b[0m\u001b[0;34m,\u001b[0m \u001b[0;34m**\u001b[0m\u001b[0mkwargs\u001b[0m\u001b[0;34m)\u001b[0m\u001b[0;34m\u001b[0m\u001b[0;34m\u001b[0m\u001b[0m\n\u001b[0m\u001b[1;32m    316\u001b[0m                 \u001b[0;32mreturn\u001b[0m \u001b[0mresult\u001b[0m\u001b[0;34m\u001b[0m\u001b[0;34m\u001b[0m\u001b[0m\n\u001b[1;32m    317\u001b[0m             \u001b[0;32mfinally\u001b[0m\u001b[0;34m:\u001b[0m\u001b[0;34m\u001b[0m\u001b[0;34m\u001b[0m\u001b[0m\n",707            "\u001b[0;32m/usr/local/lib/python3.12/dist-packages/diffusers/models/transformers/transformer_flux2.py\u001b[0m in \u001b[0;36mforward\u001b[0;34m(self, hidden_states, encoder_hidden_states, timestep, img_ids, txt_ids, guidance, joint_attention_kwargs, return_dict, kv_cache, kv_cache_mode, num_ref_tokens, ref_fixed_timestep)\u001b[0m\n\u001b[1;32m   1272\u001b[0m             \u001b[0mtxt_ids\u001b[0m \u001b[0;34m=\u001b[0m \u001b[0mtxt_ids\u001b[0m\u001b[0;34m[\u001b[0m\u001b[0;36m0\u001b[0m\u001b[0;34m]\u001b[0m\u001b[0;34m\u001b[0m\u001b[0;34m\u001b[0m\u001b[0m\n\u001b[1;32m   1273\u001b[0m \u001b[0;34m\u001b[0m\u001b[0m\n\u001b[0;32m-> 1274\u001b[0;31m         \u001b[0mimage_rotary_emb\u001b[0m \u001b[0;34m=\u001b[0m \u001b[0mself\u001b[0m\u001b[0;34m.\u001b[0m\u001b[0mpos_embed\u001b[0m\u001b[0;34m(\u001b[0m\u001b[0mimg_ids\u001b[0m\u001b[0;34m)\u001b[0m\u001b[0;34m\u001b[0m\u001b[0;34m\u001b[0m\u001b[0m\n\u001b[0m\u001b[1;32m   1275\u001b[0m         \u001b[0mtext_rotary_emb\u001b[0m \u001b[0;34m=\u001b[0m \u001b[0mself\u001b[0m\u001b[0;34m.\u001b[0m\u001b[0mpos_embed\u001b[0m\u001b[0;34m(\u001b[0m\u001b[0mtxt_ids\u001b[0m\u001b[0;34m)\u001b[0m\u001b[0;34m\u001b[0m\u001b[0;34m\u001b[0m\u001b[0m\n\u001b[1;32m   1276\u001b[0m         concat_rotary_emb = (\n",708            "\u001b[0;32m/usr/local/lib/python3.12/dist-packages/torch/nn/modules/module.py\u001b[0m in \u001b[0;36m_wrapped_call_impl\u001b[0;34m(self, *args, **kwargs)\u001b[0m\n\u001b[1;32m   1774\u001b[0m             \u001b[0;32mreturn\u001b[0m \u001b[0mself\u001b[0m\u001b[0;34m.\u001b[0m\u001b[0m_compiled_call_impl\u001b[0m\u001b[0;34m(\u001b[0m\u001b[0;34m*\u001b[0m\u001b[0margs\u001b[0m\u001b[0;34m,\u001b[0m \u001b[0;34m**\u001b[0m\u001b[0mkwargs\u001b[0m\u001b[0;34m)\u001b[0m  \u001b[0;31m# type: ignore[misc]\u001b[0m\u001b[0;34m\u001b[0m\u001b[0;34m\u001b[0m\u001b[0m\n\u001b[1;32m   1775\u001b[0m         \u001b[0;32melse\u001b[0m\u001b[0;34m:\u001b[0m\u001b[0;34m\u001b[0m\u001b[0;34m\u001b[0m\u001b[0m\n\u001b[0;32m-> 1776\u001b[0;31m             \u001b[0;32mreturn\u001b[0m \u001b[0mself\u001b[0m\u001b[0;34m.\u001b[0m\u001b[0m_call_impl\u001b[0m\u001b[0;34m(\u001b[0m\u001b[0;34m*\u001b[0m\u001b[0margs\u001b[0m\u001b[0;34m,\u001b[0m \u001b[0;34m**\u001b[0m\u001b[0mkwargs\u001b[0m\u001b[0;34m)\u001b[0m\u001b[0;34m\u001b[0m\u001b[0;34m\u001b[0m\u001b[0m\n\u001b[0m\u001b[1;32m   1777\u001b[0m \u001b[0;34m\u001b[0m\u001b[0m\n\u001b[1;32m   1778\u001b[0m     \u001b[0;31m# torchrec tests the code consistency with the following code\u001b[0m\u001b[0;34m\u001b[0m\u001b[0;34m\u001b[0m\u001b[0m\n",709            "\u001b[0;32m/usr/local/lib/python3.12/dist-packages/torch/nn/modules/module.py\u001b[0m in \u001b[0;36m_call_impl\u001b[0;34m(self, *args, **kwargs)\u001b[0m\n\u001b[1;32m   1785\u001b[0m                 \u001b[0;32mor\u001b[0m \u001b[0m_global_backward_pre_hooks\u001b[0m \u001b[0;32mor\u001b[0m \u001b[0m_global_backward_hooks\u001b[0m\u001b[0;34m\u001b[0m\u001b[0;34m\u001b[0m\u001b[0m\n\u001b[1;32m   1786\u001b[0m                 or _global_forward_hooks or _global_forward_pre_hooks):\n\u001b[0;32m-> 1787\u001b[0;31m             \u001b[0;32mreturn\u001b[0m \u001b[0mforward_call\u001b[0m\u001b[0;34m(\u001b[0m\u001b[0;34m*\u001b[0m\u001b[0margs\u001b[0m\u001b[0;34m,\u001b[0m \u001b[0;34m**\u001b[0m\u001b[0mkwargs\u001b[0m\u001b[0;34m)\u001b[0m\u001b[0;34m\u001b[0m\u001b[0;34m\u001b[0m\u001b[0m\n\u001b[0m\u001b[1;32m   1788\u001b[0m \u001b[0;34m\u001b[0m\u001b[0m\n\u001b[1;32m   1789\u001b[0m         \u001b[0mresult\u001b[0m \u001b[0;34m=\u001b[0m \u001b[0;32mNone\u001b[0m\u001b[0;34m\u001b[0m\u001b[0;34m\u001b[0m\u001b[0m\n",710            "\u001b[0;32m/usr/local/lib/python3.12/dist-packages/diffusers/models/transformers/transformer_flux2.py\u001b[0m in \u001b[0;36mforward\u001b[0;34m(self, ids)\u001b[0m\n\u001b[1;32m    967\u001b[0m             cos, sin = get_1d_rotary_pos_embed(\n\u001b[1;32m    968\u001b[0m                 \u001b[0mself\u001b[0m\u001b[0;34m.\u001b[0m\u001b[0maxes_dim\u001b[0m\u001b[0;34m[\u001b[0m\u001b[0mi\u001b[0m\u001b[0;34m]\u001b[0m\u001b[0;34m,\u001b[0m\u001b[0;34m\u001b[0m\u001b[0;34m\u001b[0m\u001b[0m\n\u001b[0;32m--> 969\u001b[0;31m                 \u001b[0mpos\u001b[0m\u001b[0;34m[\u001b[0m\u001b[0;34m...\u001b[0m\u001b[0;34m,\u001b[0m \u001b[0mi\u001b[0m\u001b[0;34m]\u001b[0m\u001b[0;34m,\u001b[0m\u001b[0;34m\u001b[0m\u001b[0;34m\u001b[0m\u001b[0m\n\u001b[0m\u001b[1;32m    970\u001b[0m                 \u001b[0mtheta\u001b[0m\u001b[0;34m=\u001b[0m\u001b[0mself\u001b[0m\u001b[0;34m.\u001b[0m\u001b[0mtheta\u001b[0m\u001b[0;34m,\u001b[0m\u001b[0;34m\u001b[0m\u001b[0;34m\u001b[0m\u001b[0m\n\u001b[1;32m    971\u001b[0m                 \u001b[0mrepeat_interleave_real\u001b[0m\u001b[0;34m=\u001b[0m\u001b[0;32mTrue\u001b[0m\u001b[0;34m,\u001b[0m\u001b[0;34m\u001b[0m\u001b[0;34m\u001b[0m\u001b[0m\n",711            "\u001b[0;31mIndexError\u001b[0m: index 2 is out of bounds for dimension 1 with size 2"712          ]713        }714      ]715    },716    {717      "cell_type": "markdown",718      "source": [719        "# Run Klein finetune (select which one to use)"720      ],721      "metadata": {722        "id": "6UhKYCWzPyrm"723      }724    },725    {726      "cell_type": "markdown",727      "source": [728        "# OLD CODE"729      ],730      "metadata": {731        "id": "Mj8UDNOyqdKL"732      }733    },734    {735      "cell_type": "code",736      "source": [737        "# =============================================================================\n",738        "#@markdown # **CELL 4**: MULTI-IMAGE INFERENCE (Klein Edit with custom/gray reference)\n",739        "# =============================================================================\n",740        "import zipfile\n",741        "import os\n",742        "import shutil\n",743        "import glob\n",744        "from PIL import Image, ImageDraw\n",745        "import torch\n",746        "import gc\n",747        "import datetime\n",748        "from google.colab import files\n",749        "\n",750        "#edit_prompt = 'remove the background. the background is gray.' #@param {type:\"string\"}\n",751        "#zip_path = '/content/drive/MyDrive/aiotest.zip' #@param {type:\"string\"}\n",752        "\n",753        "upload_custom_reference = False #@param {type:\"boolean\"}\n",754        "max_image_dimension = 2048 #@param {type:\"slider\", min:512, max:4096, step:256}\n",755        "add_gray_corners = False #@param {type:\"boolean\"}\n",756        "gray_square_size = 50 #@param {type:\"slider\", min:10, max:200, step:10}\n",757        "\n",758        "# ================= NEW CONTROLS =================\n",759        "save_to_drive_every_n = False #@param {type:\"boolean\"}\n",760        "save_every_n = 30 #@param {type:\"slider\", min:1, max:100, step:1}\n",761        "# ===============================================\n",762        "\n",763        "output_folder = '/content/edited_images_multi'\n",764        "checkpoint_folder = '/content/drive/MyDrive/klein_checkpoints'\n",765        "\n",766        "print(\"๐Ÿงน Clearing old temporary folders...\")\n",767        "for p in ['/content/input_images', output_folder]:\n",768        "    if os.path.exists(p):\n",769        "        shutil.rmtree(p)\n",770        "\n",771        "os.makedirs(output_folder, exist_ok=True)\n",772        "os.makedirs(checkpoint_folder, exist_ok=True)\n",773        "\n",774        "# ====================== Unzip ======================\n",775        "print(f\"๐Ÿ“ฆ Unzipping: {zip_path}...\")\n",776        "with zipfile.ZipFile(zip_path, 'r') as z:\n",777        "    z.extractall('/content/input_images')\n",778        "\n",779        "image_files = sorted(glob.glob('/content/input_images/*.*'))\n",780        "image_files = [f for f in image_files if f.lower().endswith(('.png', '.jpg', '.jpeg', '.webp'))]\n",781        "print(f\"Found {len(image_files)} images.\")\n",782        "\n",783        "# ====================== PROMPT MAP ======================\n",784        "prompt_map = {}\n",785        "txt_count = 0\n",786        "for img_path in image_files:\n",787        "    base = os.path.splitext(os.path.basename(img_path))[0]\n",788        "    txt_path = os.path.join('/content/input_images', f\"{base}.txt\")\n",789        "\n",790        "    if use_txt_prompts and os.path.exists(txt_path):\n",791        "        with open(txt_path, 'r', encoding='utf-8') as f:\n",792        "            prompt_map[img_path] = f.read().strip()\n",793        "        txt_count += 1\n",794        "    else:\n",795        "        prompt_map[img_path] = None\n",796        "\n",797        "print(f\"Found {txt_count} matching .txt files\")\n",798        "\n",799        "# ====================== REFERENCE IMAGE ======================\n",800        "uploaded_reference_filepath = None\n",801        "reference_image_to_pair = None\n",802        "\n",803        "if upload_custom_reference:\n",804        "    print(\"Please upload your reference image now.\")\n",805        "    uploaded_files = files.upload()\n",806        "    if uploaded_files:\n",807        "        uploaded_reference_filepath = list(uploaded_files.keys())[0]\n",808        "        reference_image_to_pair = Image.open(uploaded_reference_filepath).convert(\"RGB\")\n",809        "else:\n",810        "    reference_image_to_pair = Image.new(\"RGB\", (target_width, target_height), \"#181818\")\n",811        "\n",812        "if uploaded_reference_filepath and os.path.exists(uploaded_reference_filepath):\n",813        "    os.remove(uploaded_reference_filepath)\n",814        "\n",815        "# ====================== HELPERS ======================\n",816        "def add_corner_squares(image: Image.Image, square_size: int, color=(24, 24, 24)):\n",817        "    img_copy = image.copy()\n",818        "    draw = ImageDraw.Draw(img_copy)\n",819        "    w, h = img_copy.size\n",820        "    draw.rectangle([(0, 0), (square_size, square_size)], fill=color)\n",821        "    draw.rectangle([(w-square_size, 0), (w, square_size)], fill=color)\n",822        "    draw.rectangle([(0, h-square_size), (square_size, h)], fill=color)\n",823        "    draw.rectangle([(w-square_size, h-square_size), (w, h)], fill=color)\n",824        "    return img_copy\n",825        "\n",826        "def save_checkpoint(batch_files, checkpoint_id):\n",827        "    if not save_to_drive_every_n:\n",828        "        return\n",829        "    zip_path = f\"{checkpoint_folder}/checkpoint_{checkpoint_id}.zip\"\n",830        "    temp_folder = f\"{checkpoint_folder}/tmp_{checkpoint_id}\"\n",831        "    os.makedirs(temp_folder, exist_ok=True)\n",832        "\n",833        "    for f in batch_files:\n",834        "        shutil.copy(f, temp_folder)\n",835        "\n",836        "    shutil.make_archive(zip_path.replace('.zip',''), 'zip', temp_folder)\n",837        "    shutil.rmtree(temp_folder)\n",838        "    print(f\"๐Ÿ’พ Saved checkpoint: {zip_path}\")\n",839        "\n",840        "# ====================== INFERENCE ======================\n",841        "print(f\"\\n๐Ÿš€ Starting batch on {len(image_files)} images...\")\n",842        "\n",843        "batch_outputs = []\n",844        "checkpoint_id = 0\n",845        "\n",846        "for i, img_path in enumerate(image_files):\n",847        "    filename = os.path.basename(img_path)\n",848        "\n",849        "    gc.collect()\n",850        "    torch.cuda.empty_cache()\n",851        "\n",852        "    input_image = Image.open(img_path).convert(\"RGB\")\n",853        "\n",854        "    # resize\n",855        "    if max(input_image.width, input_image.height) > max_image_dimension:\n",856        "        aspect = input_image.width / input_image.height\n",857        "        if input_image.width > input_image.height:\n",858        "            input_image = input_image.resize(\n",859        "                (max_image_dimension, int(max_image_dimension / aspect)),\n",860        "                Image.LANCZOS\n",861        "            )\n",862        "        else:\n",863        "            input_image = input_image.resize(\n",864        "                (int(max_image_dimension * aspect), max_image_dimension),\n",865        "                Image.LANCZOS\n",866        "            )\n",867        "\n",868        "    if add_gray_corners:\n",869        "        input_image = add_corner_squares(input_image, gray_square_size)\n",870        "\n",871        "    reference_images = [input_image, reference_image_to_pair]\n",872        "\n",873        "    current_prompt = (\n",874        "        prompt_map[img_path]\n",875        "        if use_txt_prompts and prompt_map[img_path] is not None\n",876        "        else edit_prompt\n",877        "    )\n",878        "\n",879        "    result = pipe(\n",880        "        prompt=current_prompt,\n",881        "        image=reference_images,\n",882        "        height=target_height,\n",883        "        width=target_width,\n",884        "        guidance_scale=1.0,\n",885        "        num_inference_steps=4,\n",886        "        generator=torch.Generator(\"cuda\").manual_seed(42),\n",887        "        output_type=\"pil\",\n",888        "    ).images[0]\n",889        "\n",890        "    out_path = os.path.join(output_folder, f\"edited_{filename}\")\n",891        "    result.save(out_path)\n",892        "    batch_outputs.append(out_path)\n",893        "\n",894        "    print(f\"[{i+1}/{len(image_files)}] saved โ†’ {out_path}\")\n",895        "\n",896        "    # ====================== SAVE EVERY N ======================\n",897        "    if save_to_drive_every_n and (len(batch_outputs) >= save_every_n):\n",898        "        checkpoint_id += 1\n",899        "        save_checkpoint(batch_outputs, checkpoint_id)\n",900        "        batch_outputs = []  # reset buffer\n",901        "\n",902        "# ====================== FINAL FLUSH SAVE ======================\n",903        "if save_to_drive_every_n and len(batch_outputs) > 0:\n",904        "    checkpoint_id += 1\n",905        "    save_checkpoint(batch_outputs, checkpoint_id)\n",906        "\n",907        "print(\"\\nโœ… BATCH COMPLETE!\")\n",908        "\n",909        "# ====================== FINAL ZIP ======================\n",910        "timestamp = datetime.datetime.now().strftime(\"%Y%m%d_%H%M%S\")\n",911        "final_zip = f\"/content/drive/MyDrive/klein_edited_{timestamp}.zip\"\n",912        "\n",913        "print(f\"๐Ÿ“ฆ Final zip โ†’ {final_zip}\")\n",914        "shutil.make_archive(final_zip.replace('.zip',''), 'zip', output_folder)\n",915        "\n",916        "print(\"๐ŸŽฏ DONE\")"917      ],918      "metadata": {919        "cellView": "form",920        "id": "Q14frr2en161"921      },922      "execution_count": null,923      "outputs": []924    },925    {926      "cell_type": "markdown",927      "source": [928        "# Cell 4 version with Encrypt (wip)"929      ],930      "metadata": {931        "id": "TEupY-WNeMHu"932      }933    },934    {935      "cell_type": "code",936      "source": [937        "# =============================================================================\n",938        "#@markdown # **CELL 4 INFERENCE\n",939        "# =============================================================================\n",940        "import zipfile\n",941        "import os\n",942        "import shutil\n",943        "import glob\n",944        "from PIL import Image\n",945        "import torch\n",946        "import gc\n",947        "import datetime\n",948        "import hashlib\n",949        "import io\n",950        "\n",951        "!pip install -q pynacl\n",952        "from nacl.secret import SecretBox\n",953        "from nacl.utils import random\n",954        "\n",955        "# ================= SETTINGS =================\n",956        "encryption_password = \"banana\"  #@param {type:\"string\"}\n",957        "\n",958        "upload_custom_reference = False #@param {type:\"boolean\"}\n",959        "max_image_dimension = 2048\n",960        "\n",961        "output_folder = '/content/encrypted_outputs'\n",962        "\n",963        "# ================= CLEAN =================\n",964        "for p in ['/content/input_images', output_folder]:\n",965        "    if os.path.exists(p):\n",966        "        shutil.rmtree(p)\n",967        "\n",968        "os.makedirs(output_folder, exist_ok=True)\n",969        "\n",970        "# ================= KEY =================\n",971        "def derive_key(password):\n",972        "    return hashlib.sha256(password.encode()).digest()\n",973        "\n",974        "box = SecretBox(derive_key(encryption_password))\n",975        "\n",976        "def pil_to_bytes(img):\n",977        "    buf = io.BytesIO()\n",978        "    img.save(buf, format=\"JPEG\", quality=95)\n",979        "    return buf.getvalue()\n",980        "\n",981        "def encrypt_bytes(data):\n",982        "    nonce = random(SecretBox.NONCE_SIZE)\n",983        "    enc = box.encrypt(data, nonce)\n",984        "    return enc.nonce + enc.ciphertext\n",985        "\n",986        "# ================= UNZIP =================\n",987        "print(f\"๐Ÿ“ฆ Unzipping: {zip_path}\")\n",988        "with zipfile.ZipFile(zip_path, 'r') as z:\n",989        "    z.extractall('/content/input_images')\n",990        "\n",991        "image_files = sorted(glob.glob('/content/input_images/*.*'))\n",992        "image_files = [f for f in image_files if f.lower().endswith(('.png','.jpg','.jpeg','.webp'))]\n",993        "\n",994        "print(f\"Found {len(image_files)} images\")\n",995        "\n",996        "# ================= REFERENCE =================\n",997        "if upload_custom_reference:\n",998        "    from google.colab import files\n",999        "    uploaded = files.upload()\n",1000        "    ref_path = list(uploaded.keys())[0]\n",1001        "    reference_image = Image.open(ref_path).convert(\"RGB\")\n",1002        "else:\n",1003        "    reference_image = Image.new(\"RGB\", (target_width, target_height), \"#181818\")\n",1004        "\n",1005        "# ================= INFERENCE =================\n",1006        "print(\"๐Ÿš€ Running inference + encryption...\")\n",1007        "\n",1008        "for i, img_path in enumerate(image_files):\n",1009        "    gc.collect()\n",1010        "    torch.cuda.empty_cache()\n",1011        "\n",1012        "    filename = os.path.basename(img_path)\n",1013        "    img = Image.open(img_path).convert(\"RGB\")\n",1014        "\n",1015        "    # resize\n",1016        "    if max(img.width, img.height) > max_image_dimension:\n",1017        "        aspect = img.width / img.height\n",1018        "        if img.width > img.height:\n",1019        "            img = img.resize((max_image_dimension, int(max_image_dimension/aspect)), Image.LANCZOS)\n",1020        "        else:\n",1021        "            img = img.resize((int(max_image_dimension*aspect), max_image_dimension), Image.LANCZOS)\n",1022        "\n",1023        "    result = pipe(\n",1024        "        prompt=edit_prompt,\n",1025        "        image=[img, reference_image],\n",1026        "        height=target_height,\n",1027        "        width=target_width,\n",1028        "        guidance_scale=1.0,\n",1029        "        num_inference_steps=4,\n",1030        "        generator=torch.Generator(\"cuda\").manual_seed(42),\n",1031        "        output_type=\"pil\",\n",1032        "    ).images[0]\n",1033        "\n",1034        "    # ===== ENCRYPT SAVE =====\n",1035        "    img_bytes = pil_to_bytes(result)\n",1036        "    encrypted_data = encrypt_bytes(img_bytes)\n",1037        "\n",1038        "    out_name = f\"{os.path.splitext(filename)[0]}_encrypted.bin\"\n",1039        "    out_path = os.path.join(output_folder, out_name)\n",1040        "\n",1041        "    with open(out_path, \"wb\") as f:\n",1042        "        f.write(encrypted_data)\n",1043        "\n",1044        "    print(f\"[{i+1}/{len(image_files)}] ๐Ÿ” saved โ†’ {out_path}\")\n",1045        "\n",1046        "# ================= ZIP =================\n",1047        "timestamp = datetime.datetime.now().strftime(\"%Y%m%d_%H%M%S\")\n",1048        "final_zip = f\"/content/drive/MyDrive/klein_encrypted_{timestamp}.zip\"\n",1049        "\n",1050        "shutil.make_archive(final_zip.replace('.zip',''), 'zip', output_folder)\n",1051        "\n",1052        "print(f\"\\n๐Ÿ“ฆ Final encrypted zip โ†’ {final_zip}\")\n",1053        "print(\"โœ… DONE\")"1054      ],1055      "metadata": {1056        "cellView": "form",1057        "id": "x_8QixuPcNhL"1058      },1059      "execution_count": null,1060      "outputs": []1061    },1062    {1063      "cell_type": "code",1064      "source": [1065        "# Clean unassigned junk files from VRAM  , run this if interrupting the cell script\n",1066        "import torch , gc\n",1067        "torch.cuda.empty_cache()\n",1068        "gc.collect()"1069      ],1070      "metadata": {1071        "id": "e46WUV7jMInu"1072      },1073      "execution_count": null,1074      "outputs": []1075    },1076    {1077      "cell_type": "markdown",1078      "source": [1079        "# ๐Ÿ”Œ Auto Disconnect from Drive\n",1080        "\n",1081        "Results from earlier cell have been saved to your drive as\n",1082        "'content/drive/MyDrive/klein_processed_(date_number).zip'\n",1083        "\n",1084        "You can disconnect from the session , and reconnect to new runtime to prevent conflict in installed libraries. you can run this code on CPU"1085      ],1086      "metadata": {1087        "id": "7eeX3_qpYpwl"1088      }1089    },1090    {1091      "cell_type": "code",1092      "execution_count": null,1093      "metadata": {1094        "id": "FQF71-mvmlc1",1095        "cellView": "form"1096      },1097      "outputs": [],1098      "source": [1099        "# ================================================\n",1100        "#@title ๐Ÿ”Œ Auto Disconnect Colab Session\n",1101        "# ================================================\n",1102        "\n",1103        "enable = False #@param {type:'boolean'}\n",1104        "if enable:\n",1105        "  print(\"๐Ÿ”Œ Disconnecting Colab session in 3 seconds...\")\n",1106        "  import time\n",1107        "  time.sleep(3)\n",1108        "\n",1109        "  from google.colab import runtime\n",1110        "  runtime.unassign()\n",1111        "\n",1112        "  print(\"Session disconnected.\")"1113      ]1114    },1115    {1116      "cell_type": "markdown",1117      "source": [1118        "# ๐Ÿ”‘ Decrypt (wip)"1119      ],1120      "metadata": {1121        "id": "pke7YfSHd4u6"1122      }1123    },1124    {1125      "cell_type": "code",1126      "source": [1127        "# =============================================================================\n",1128        "# ๐Ÿ”‘ Session C: Decrypt most recent encrypted ZIP from Drive (FULLY COMPATIBLE)\n",1129        "# =============================================================================\n",1130        "\n",1131        "from google.colab import drive, files\n",1132        "import os\n",1133        "import zipfile\n",1134        "import io\n",1135        "import hashlib\n",1136        "from PIL import Image\n",1137        "\n",1138        "!pip install -q pynacl\n",1139        "from nacl.secret import SecretBox\n",1140        "\n",1141        "# =============================================================================\n",1142        "# 1. Mount Drive\n",1143        "# =============================================================================\n",1144        "drive.mount('/content/drive')\n",1145        "\n",1146        "# =============================================================================\n",1147        "# 2. Config\n",1148        "# =============================================================================\n",1149        "decryption_password = \"banana\"  #@param{type:'string'}\n",1150        "output_zip_name = \"decrypted_results.zip\"  #@param {type:\"string\"}\n",1151        "image_format = \"JPG\"  #@param [\"PNG\", \"JPG\"]\n",1152        "jpeg_quality = 95 #@param {type:\"slider\", min:1, max:100, step:1}\n",1153        "\n",1154        "output_base_folder = \"/content/zip_outputs\"\n",1155        "os.makedirs(output_base_folder, exist_ok=True)\n",1156        "\n",1157        "# =============================================================================\n",1158        "# 3. Key setup\n",1159        "# =============================================================================\n",1160        "def derive_key(password):\n",1161        "    return hashlib.sha256(password.encode()).digest()\n",1162        "\n",1163        "box = SecretBox(derive_key(decryption_password))\n",1164        "\n",1165        "def decrypt_bytes(enc):\n",1166        "    nonce = enc[:24]\n",1167        "    ciphertext = enc[24:]\n",1168        "    return box.decrypt(ciphertext, nonce)\n",1169        "\n",1170        "# =============================================================================\n",1171        "# 4. Find MOST RECENT matching ZIP\n",1172        "# =============================================================================\n",1173        "drive_base_folder = \"/content/drive/MyDrive\"\n",1174        "search_folders = [\n",1175        "    drive_base_folder,\n",1176        "    os.path.join(drive_base_folder, \"Saved from Chrome\")\n",1177        "]\n",1178        "\n",1179        "print(\"๐Ÿ” Searching for encrypted ZIPs...\")\n",1180        "\n",1181        "matching_zips = []\n",1182        "\n",1183        "for folder in search_folders:\n",1184        "    if not os.path.exists(folder):\n",1185        "        continue\n",1186        "\n",1187        "    for f in os.listdir(folder):\n",1188        "        f_lower = f.lower()\n",1189        "\n",1190        "        is_zip = f_lower.endswith(\".zip\")\n",1191        "\n",1192        "        # ===== SUPPORT OLD + NEW NAMING =====\n",1193        "        contains_old = \"klein_encrypted_outputs\" in f_lower\n",1194        "        contains_new = \"klein_encrypted_\" in f_lower\n",1195        "        contains_checkpoint = \"checkpoint_gpu\" in f_lower\n",1196        "\n",1197        "        if is_zip and (contains_old or contains_new or contains_checkpoint):\n",1198        "            full_path = os.path.join(folder, f)\n",1199        "            matching_zips.append(full_path)\n",1200        "            print(f\"โœ… Found: {f}\")\n",

Showing the first 1,200 of 3812 lines. Download the file for the rest.