codeShare/lora-training-data
2785
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",