aryadomain commited on
Commit
9ebbe39
·
verified ·
1 Parent(s): 6d6dbbc

Add files using upload-large-folder tool

Browse files
Reward_sd15_idealized/__pycache__/lr_scheduler.cpython-313.pyc ADDED
Binary file (9.36 kB). View file
 
Reward_sd15_idealized/timestep_convergence_analysis.ipynb ADDED
@@ -0,0 +1,1105 @@
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
1
+ {
2
+ "cells": [
3
+ {
4
+ "cell_type": "markdown",
5
+ "id": "513b682a",
6
+ "metadata": {},
7
+ "source": [
8
+ "#### Setup and Imports"
9
+ ]
10
+ },
11
+ {
12
+ "cell_type": "code",
13
+ "execution_count": null,
14
+ "id": "f9f44d61",
15
+ "metadata": {},
16
+ "outputs": [],
17
+ "source": [
18
+ "import os\n",
19
+ "import sys\n",
20
+ "import json\n",
21
+ "import warnings\n",
22
+ "import numpy as np\n",
23
+ "import matplotlib.pyplot as plt\n",
24
+ "import matplotlib\n",
25
+ "import torch\n",
26
+ "import torch.nn as nn\n",
27
+ "from pathlib import Path\n",
28
+ "from PIL import Image\n",
29
+ "from tqdm.auto import tqdm\n",
30
+ "from diffusers import StableDiffusionPipeline, DDIMScheduler, UNet2DConditionModel\n",
31
+ "from transformers import CLIPModel, CLIPProcessor\n",
32
+ "from torchmetrics.image.fid import FrechetInceptionDistance\n",
33
+ "from torchmetrics.multimodal import CLIPScore\n",
34
+ "\n",
35
+ "warnings.filterwarnings(\"ignore\")\n",
36
+ "\n",
37
+ "# Import local modules\n",
38
+ "from models import LRMRewardModel\n",
39
+ "from pipelines.sd15_gradient_ascent_pipeline import StableDiffusionGradientAscentPipeline\n",
40
+ "from grad_ascent_configs import get_config, list_configs\n",
41
+ "\n",
42
+ "# Import evaluation metrics\n",
43
+ "sys.path.append('../evaluation')\n",
44
+ "from pick_score import PickScorer\n",
45
+ "from hpsv2_score import HPSv2Scorer\n",
46
+ "from imagereward_score import load_imagereward\n"
47
+ ]
48
+ },
49
+ {
50
+ "cell_type": "markdown",
51
+ "id": "1740dd7c",
52
+ "metadata": {},
53
+ "source": [
54
+ "#### Configuration"
55
+ ]
56
+ },
57
+ {
58
+ "cell_type": "code",
59
+ "execution_count": null,
60
+ "id": "f1bc2b07",
61
+ "metadata": {},
62
+ "outputs": [],
63
+ "source": [
64
+ "# ============ CONFIGURATION ============\n",
65
+ "\n",
66
+ "# Dataset\n",
67
+ "DATA_DIR = \"./data\"\n",
68
+ "DATASET_TYPE = \"coco\" # \"coco\" or \"pickapic\"\n",
69
+ "NUM_SAMPLES = 20 # Number of samples to analyze\n",
70
+ "\n",
71
+ "# Model\n",
72
+ "BASE_MODEL = \"runwayml/stable-diffusion-v1-5\"\n",
73
+ "MODEL_VARIANT = \"lpo\" # \"origin\", \"spo\", \"diffusion_dpo\", \"lpo\"\n",
74
+ "LRM_MODEL = \"casiatao/LRM\"\n",
75
+ "\n",
76
+ "# Generation\n",
77
+ "NUM_INFERENCE_STEPS = 100\n",
78
+ "CFG_SCALE = 5.0\n",
79
+ "SEED = 42\n",
80
+ "BATCH_SIZE = 1\n",
81
+ "\n",
82
+ "# Gradient Ascent Config\n",
83
+ "GRAD_CONFIG = \"low_to_high_nesterov\" # Use None for manual config, or specify preset name\n",
84
+ "GRAD_RANGE_START = 0\n",
85
+ "GRAD_RANGE_END = 500\n",
86
+ "GRAD_STEPS = 1\n",
87
+ "GRAD_STEP_SIZE = 0.1\n",
88
+ "\n",
89
+ "# Metrics to compute\n",
90
+ "METRICS = [\"reward\", \"clip\", \"aesthetic\", \"pickscore\", \"hpsv2\", \"fid\"] # Add/remove as needed\n",
91
+ "\n",
92
+ "# Device\n",
93
+ "CUDA_DEVICE = 0\n",
94
+ "device = f\"cuda:{CUDA_DEVICE}\" if torch.cuda.is_available() else \"cpu\"\n",
95
+ "dtype = torch.float16 if torch.cuda.is_available() else torch.float32\n",
96
+ "\n",
97
+ "# Output\n",
98
+ "OUTPUT_DIR = \"timestep_analysis_results\"\n",
99
+ "os.makedirs(OUTPUT_DIR, exist_ok=True)\n",
100
+ "\n",
101
+ "print(f\"Device: {device}\")\n",
102
+ "print(f\"Dataset: {DATASET_TYPE}\")\n",
103
+ "print(f\"Samples to analyze: {NUM_SAMPLES}\")\n",
104
+ "print(f\"Metrics: {METRICS}\")\n",
105
+ "print(f\"Output directory: {OUTPUT_DIR}\")"
106
+ ]
107
+ },
108
+ {
109
+ "cell_type": "markdown",
110
+ "id": "1b1b6d02",
111
+ "metadata": {},
112
+ "source": [
113
+ "#### Load Dataset"
114
+ ]
115
+ },
116
+ {
117
+ "cell_type": "code",
118
+ "execution_count": null,
119
+ "id": "a74b2816",
120
+ "metadata": {},
121
+ "outputs": [],
122
+ "source": [
123
+ "def load_validation_data(data_dir, max_samples=None):\n",
124
+ " \"\"\"Load COCO validation prompts and image paths.\"\"\"\n",
125
+ " data_dir = Path(data_dir)\n",
126
+ " val_json = data_dir / \"coco\" / \"caption_val.json\"\n",
127
+ " \n",
128
+ " if not val_json.exists():\n",
129
+ " raise FileNotFoundError(f\"Validation data not found at {val_json}\")\n",
130
+ " \n",
131
+ " with open(val_json, 'r') as f:\n",
132
+ " data = json.load(f)\n",
133
+ " \n",
134
+ " print(f\"Loaded JSON with {len(data)} entries\")\n",
135
+ " \n",
136
+ " # Validate that image folder exists\n",
137
+ " val_img_dir = data_dir / \"coco\" / \"images\" / \"val\"\n",
138
+ " if not val_img_dir.exists():\n",
139
+ " print(f\"Warning: Standard validation directory not found: {val_img_dir}\")\n",
140
+ " \n",
141
+ " # Parse data - img_path already contains \"images/val/\" prefix\n",
142
+ " prompts = []\n",
143
+ " image_paths = []\n",
144
+ " \n",
145
+ " for img_path, caption in data.items():\n",
146
+ " # Try the path as given (relative to data_dir/coco/)\n",
147
+ " full_path = data_dir / \"coco\" / img_path\n",
148
+ " if full_path.exists():\n",
149
+ " prompts.append(caption)\n",
150
+ " image_paths.append(str(full_path))\n",
151
+ " \n",
152
+ " print(f\"Found {len(prompts)} valid image-caption pairs\")\n",
153
+ " \n",
154
+ " if len(prompts) == 0:\n",
155
+ " print(f\"\\n⚠ WARNING: No valid images found!\")\n",
156
+ " print(f\"Debug information:\")\n",
157
+ " print(f\" JSON file: {val_json}\")\n",
158
+ " print(f\" JSON entries: {len(data)}\")\n",
159
+ " print(f\" Sample keys from JSON: {list(data.keys())[:3]}\")\n",
160
+ " \n",
161
+ " # Check if images exist at all\n",
162
+ " coco_dir = data_dir / \"coco\"\n",
163
+ " if coco_dir.exists():\n",
164
+ " print(f\" COCO dir exists: {coco_dir}\")\n",
165
+ " # List subdirectories\n",
166
+ " subdirs = [d.name for d in coco_dir.iterdir() if d.is_dir()]\n",
167
+ " print(f\" Subdirectories in COCO: {subdirs}\")\n",
168
+ " \n",
169
+ " # Try to find images\n",
170
+ " if val_img_dir.exists():\n",
171
+ " img_files = list(val_img_dir.glob(\"*.jpg\"))[:5]\n",
172
+ " print(f\" Sample images in val dir: {[f.name for f in img_files]}\")\n",
173
+ " \n",
174
+ " if max_samples and len(prompts) > 0:\n",
175
+ " prompts = prompts[:max_samples]\n",
176
+ " image_paths = image_paths[:max_samples]\n",
177
+ " \n",
178
+ " return prompts, image_paths\n",
179
+ "\n",
180
+ "# Load data\n",
181
+ "prompts, image_paths = load_validation_data(DATA_DIR, NUM_SAMPLES)\n",
182
+ "print(f\"\\n✓ Loaded {len(prompts)} samples\")\n",
183
+ "\n",
184
+ "if len(prompts) > 0:\n",
185
+ " print(f\"\\nSample prompts:\")\n",
186
+ " for i, prompt in enumerate(prompts[:3]):\n",
187
+ " print(f\" {i+1}. {prompt[:80]}...\")\n",
188
+ " print(f\"\\nSample image paths:\")\n",
189
+ " for i, path in enumerate(image_paths[:3]):\n",
190
+ " print(f\" {i+1}. {path}\")\n",
191
+ "else:\n",
192
+ " print(\"\\n❌ ERROR: No samples loaded! Please check your data directory structure.\")\n",
193
+ " print(\"Expected structure:\")\n",
194
+ " print(\" ./data/coco/caption_val.json\")\n",
195
+ " print(\" ./data/coco/images/val/*.jpg\")"
196
+ ]
197
+ },
198
+ {
199
+ "cell_type": "markdown",
200
+ "id": "5ceae64a",
201
+ "metadata": {},
202
+ "source": [
203
+ "#### Load Models and Scorers"
204
+ ]
205
+ },
206
+ {
207
+ "cell_type": "code",
208
+ "execution_count": null,
209
+ "id": "43ad1f56",
210
+ "metadata": {},
211
+ "outputs": [],
212
+ "source": [
213
+ "# ============ MLP for Aesthetic Scoring ============\n",
214
+ "class MLP(nn.Module):\n",
215
+ " def __init__(self):\n",
216
+ " super().__init__()\n",
217
+ " self.layers = nn.Sequential(\n",
218
+ " nn.Linear(768, 1024),\n",
219
+ " nn.Dropout(0.2),\n",
220
+ " nn.Linear(1024, 128),\n",
221
+ " nn.Dropout(0.2),\n",
222
+ " nn.Linear(128, 64),\n",
223
+ " nn.Dropout(0.1),\n",
224
+ " nn.Linear(64, 16),\n",
225
+ " nn.Linear(16, 1),\n",
226
+ " )\n",
227
+ " \n",
228
+ " @torch.no_grad()\n",
229
+ " def forward(self, embed):\n",
230
+ " return self.layers(embed)\n",
231
+ "\n",
232
+ "class AestheticScorer(torch.nn.Module):\n",
233
+ " def __init__(self, dtype, device):\n",
234
+ " super().__init__()\n",
235
+ " self.clip = CLIPModel.from_pretrained(\"openai/clip-vit-large-patch14\")\n",
236
+ " self.processor = CLIPProcessor.from_pretrained(\"openai/clip-vit-large-patch14\")\n",
237
+ " self.mlp = MLP()\n",
238
+ " \n",
239
+ " aesthetic_path = \"../evaluation/sac+logos+ava1-l14-linearMSE.pth\"\n",
240
+ " if os.path.exists(aesthetic_path):\n",
241
+ " state_dict = torch.load(aesthetic_path, map_location='cpu')\n",
242
+ " self.mlp.load_state_dict(state_dict)\n",
243
+ " \n",
244
+ " self.dtype = dtype\n",
245
+ " self.to(device)\n",
246
+ " self.eval()\n",
247
+ " \n",
248
+ " @torch.no_grad()\n",
249
+ " def __call__(self, images):\n",
250
+ " if not isinstance(images, list):\n",
251
+ " images = [images]\n",
252
+ " inputs = self.processor(images=images, return_tensors=\"pt\", padding=True)\n",
253
+ " inputs = {k: v.to(self.clip.device) for k, v in inputs.items()}\n",
254
+ " image_embeds = self.clip.get_image_features(**inputs)\n",
255
+ " image_embeds = image_embeds / image_embeds.norm(dim=-1, keepdim=True)\n",
256
+ " scores = self.mlp(image_embeds.float())\n",
257
+ " return scores.squeeze().cpu().numpy()\n",
258
+ "\n",
259
+ "print(\"Loading models...\")"
260
+ ]
261
+ },
262
+ {
263
+ "cell_type": "code",
264
+ "execution_count": null,
265
+ "id": "36a70595",
266
+ "metadata": {},
267
+ "outputs": [],
268
+ "source": [
269
+ "# Load Reward Model\n",
270
+ "print(\"Loading reward model...\")\n",
271
+ "reward_model = LRMRewardModel(\n",
272
+ " pretrained_model_name_or_path=BASE_MODEL,\n",
273
+ " lrm_model_path=LRM_MODEL,\n",
274
+ " guidance_scale=CFG_SCALE,\n",
275
+ " device=device\n",
276
+ ")\n",
277
+ "if dtype == torch.float16:\n",
278
+ " reward_model = reward_model.half()\n",
279
+ "reward_model.eval()\n",
280
+ "print(\"✓ Reward model loaded\")\n",
281
+ "\n",
282
+ "# Load Pipeline\n",
283
+ "print(\"\\nLoading diffusion pipeline...\")\n",
284
+ "if MODEL_VARIANT == \"origin\":\n",
285
+ " base_pipeline = StableDiffusionPipeline.from_pretrained(\n",
286
+ " BASE_MODEL, torch_dtype=dtype, safety_checker=None\n",
287
+ " )\n",
288
+ "elif MODEL_VARIANT == \"spo\":\n",
289
+ " base_pipeline = StableDiffusionPipeline.from_pretrained(\n",
290
+ " 'SPO-Diffusion-Models/SPO-SD-v1-5_4k-p_10ep',\n",
291
+ " torch_dtype=dtype, safety_checker=None\n",
292
+ " )\n",
293
+ " CFG_SCALE = 5.0\n",
294
+ "elif MODEL_VARIANT == \"diffusion_dpo\":\n",
295
+ " unet = UNet2DConditionModel.from_pretrained(\n",
296
+ " 'mhdang/dpo-sd1.5-text2image-v1', subfolder=\"unet\", torch_dtype=dtype\n",
297
+ " )\n",
298
+ " base_pipeline = StableDiffusionPipeline.from_pretrained(\n",
299
+ " BASE_MODEL, torch_dtype=dtype, safety_checker=None, unet=unet\n",
300
+ " )\n",
301
+ "elif MODEL_VARIANT == \"lpo\":\n",
302
+ " unet = UNet2DConditionModel.from_pretrained(\n",
303
+ " 'casiatao/LPO', subfolder=\"lpo_sd15_merge/unet\", torch_dtype=dtype\n",
304
+ " )\n",
305
+ " base_pipeline = StableDiffusionPipeline.from_pretrained(\n",
306
+ " BASE_MODEL, torch_dtype=dtype, safety_checker=None, unet=unet\n",
307
+ " )\n",
308
+ " CFG_SCALE = 5.0\n",
309
+ "\n",
310
+ "pipeline = StableDiffusionGradientAscentPipeline(**base_pipeline.components)\n",
311
+ "pipeline.scheduler = DDIMScheduler.from_config(pipeline.scheduler.config)\n",
312
+ "pipeline = pipeline.to(device)\n",
313
+ "pipeline.set_reward_model(reward_model)\n",
314
+ "print(\"✓ Pipeline loaded\")"
315
+ ]
316
+ },
317
+ {
318
+ "cell_type": "code",
319
+ "execution_count": null,
320
+ "id": "4e18f075",
321
+ "metadata": {},
322
+ "outputs": [],
323
+ "source": [
324
+ "# Load Metric Scorers\n",
325
+ "print(\"\\nLoading metric scorers...\")\n",
326
+ "\n",
327
+ "clip_scorer = None\n",
328
+ "aesthetic_scorer = None\n",
329
+ "pick_scorer = None\n",
330
+ "hpsv2_scorer = None\n",
331
+ "imagereward_scorer = None\n",
332
+ "\n",
333
+ "if \"clip\" in METRICS:\n",
334
+ " print(\" Loading CLIP scorer...\")\n",
335
+ " clip_scorer = CLIPScore(model_name_or_path=\"openai/clip-vit-base-patch16\").to(device)\n",
336
+ " print(\" ✓ CLIP scorer loaded\")\n",
337
+ "\n",
338
+ "if \"aesthetic\" in METRICS:\n",
339
+ " print(\" Loading Aesthetic scorer...\")\n",
340
+ " aesthetic_scorer = AestheticScorer(dtype, device)\n",
341
+ " print(\" ✓ Aesthetic scorer loaded\")\n",
342
+ "\n",
343
+ "if \"pickscore\" in METRICS:\n",
344
+ " print(\" Loading PickScore scorer...\")\n",
345
+ " try:\n",
346
+ " pick_scorer = PickScorer(device=device, dtype=dtype)\n",
347
+ " print(\" ✓ PickScore loaded\")\n",
348
+ " except Exception as e:\n",
349
+ " print(f\" ✗ PickScore failed: {e}\")\n",
350
+ " METRICS.remove(\"pickscore\")\n",
351
+ "\n",
352
+ "if \"hpsv2\" in METRICS:\n",
353
+ " print(\" Loading HPSv2 scorer...\")\n",
354
+ " try:\n",
355
+ " hpsv2_scorer = HPSv2Scorer(device=device, dtype=dtype)\n",
356
+ " print(\" ✓ HPSv2 loaded\")\n",
357
+ " except Exception as e:\n",
358
+ " print(f\" ✗ HPSv2 failed: {e}\")\n",
359
+ " METRICS.remove(\"hpsv2\")\n",
360
+ "\n",
361
+ "if \"imagereward\" in METRICS:\n",
362
+ " print(\" Loading ImageReward scorer...\")\n",
363
+ " try:\n",
364
+ " imagereward_scorer = load_imagereward(device=device)\n",
365
+ " print(\" ✓ ImageReward loaded\")\n",
366
+ " except Exception as e:\n",
367
+ " print(f\" ✗ ImageReward failed: {e}\")\n",
368
+ " METRICS.remove(\"imagereward\")\n",
369
+ "\n",
370
+ "print(f\"\\n✓ Active metrics: {METRICS}\")"
371
+ ]
372
+ },
373
+ {
374
+ "cell_type": "markdown",
375
+ "id": "70ac047b",
376
+ "metadata": {},
377
+ "source": [
378
+ "#### Configure Gradient Ascent"
379
+ ]
380
+ },
381
+ {
382
+ "cell_type": "code",
383
+ "execution_count": null,
384
+ "id": "05996448",
385
+ "metadata": {},
386
+ "outputs": [],
387
+ "source": [
388
+ "# Configure gradient ascent\n",
389
+ "if GRAD_CONFIG:\n",
390
+ " print(f\"Loading gradient ascent config: {GRAD_CONFIG}\")\n",
391
+ " grad_config = get_config(GRAD_CONFIG)\n",
392
+ " print(f\"Config: {grad_config}\")\n",
393
+ "else:\n",
394
+ " grad_config = {\n",
395
+ " \"grad_timestep_range\": (GRAD_RANGE_START, GRAD_RANGE_END),\n",
396
+ " \"num_grad_steps\": GRAD_STEPS,\n",
397
+ " \"grad_step_size\": GRAD_STEP_SIZE,\n",
398
+ " }\n",
399
+ " print(f\"Manual gradient ascent configuration: {grad_config}\")\n",
400
+ "\n",
401
+ "pipeline.enable_gradient_ascent(**grad_config)\n",
402
+ "print(\"\\n✓ Gradient ascent enabled\")"
403
+ ]
404
+ },
405
+ {
406
+ "cell_type": "markdown",
407
+ "id": "1f82c3df",
408
+ "metadata": {},
409
+ "source": [
410
+ "#### Timestep Analysis Functions"
411
+ ]
412
+ },
413
+ {
414
+ "cell_type": "code",
415
+ "execution_count": null,
416
+ "id": "e836d8f2",
417
+ "metadata": {},
418
+ "outputs": [],
419
+ "source": [
420
+ "def latents_to_images(latents, vae):\n",
421
+ " \"\"\"Convert latents to PIL images.\"\"\"\n",
422
+ " latents = 1 / 0.18215 * latents\n",
423
+ " with torch.no_grad():\n",
424
+ " images = vae.decode(latents).sample\n",
425
+ " images = (images / 2 + 0.5).clamp(0, 1)\n",
426
+ " images = images.cpu().permute(0, 2, 3, 1).numpy()\n",
427
+ " images = (images * 255).round().astype(\"uint8\")\n",
428
+ " pil_images = [Image.fromarray(image) for image in images]\n",
429
+ " return pil_images\n",
430
+ "\n",
431
+ "\n",
432
+ "def compute_metrics_for_image(image, prompt, reference_image=None):\n",
433
+ " \"\"\"Compute all metrics for a single image.\"\"\"\n",
434
+ " metrics = {}\n",
435
+ " \n",
436
+ " # CLIP Score\n",
437
+ " if clip_scorer is not None:\n",
438
+ " img_tensor = torch.from_numpy(np.array(image)).permute(2, 0, 1).unsqueeze(0).to(device)\n",
439
+ " with torch.no_grad():\n",
440
+ " clip_score = clip_scorer(img_tensor, prompt).item()\n",
441
+ " metrics['clip'] = clip_score\n",
442
+ " \n",
443
+ " # Aesthetic Score\n",
444
+ " if aesthetic_scorer is not None:\n",
445
+ " aesthetic_score = aesthetic_scorer([image])\n",
446
+ " if isinstance(aesthetic_score, np.ndarray):\n",
447
+ " aesthetic_score = aesthetic_score.item()\n",
448
+ " metrics['aesthetic'] = aesthetic_score\n",
449
+ " \n",
450
+ " # PickScore\n",
451
+ " if pick_scorer is not None:\n",
452
+ " pick_score = pick_scorer.score(prompt, [image])[0]\n",
453
+ " metrics['pickscore'] = pick_score\n",
454
+ " \n",
455
+ " # HPSv2\n",
456
+ " if hpsv2_scorer is not None:\n",
457
+ " hpsv2_score = hpsv2_scorer.score(prompt, [image])[0]\n",
458
+ " metrics['hpsv2'] = hpsv2_score\n",
459
+ " \n",
460
+ " # ImageReward\n",
461
+ " if imagereward_scorer is not None:\n",
462
+ " imagereward_score = imagereward_scorer.score(prompt, [image])[0]\n",
463
+ " metrics['imagereward'] = imagereward_score\n",
464
+ " \n",
465
+ " # FID (if reference image provided)\n",
466
+ " if reference_image is not None:\n",
467
+ " try:\n",
468
+ " fid_metric = FrechetInceptionDistance(normalize=True).to(device)\n",
469
+ " \n",
470
+ " # Process reference image\n",
471
+ " ref_img = Image.open(reference_image).convert('RGB').resize((299, 299))\n",
472
+ " ref_tensor = torch.from_numpy(np.array(ref_img)).permute(2, 0, 1).unsqueeze(0).to(device)\n",
473
+ " \n",
474
+ " # Process generated image\n",
475
+ " gen_img = image.resize((299, 299))\n",
476
+ " gen_tensor = torch.from_numpy(np.array(gen_img)).permute(2, 0, 1).unsqueeze(0).to(device)\n",
477
+ " \n",
478
+ " if ref_tensor.size(0) == 1:\n",
479
+ " ref_tensor = ref_tensor.repeat(2, 1, 1, 1)\n",
480
+ " if gen_tensor.size(0) == 1:\n",
481
+ " gen_tensor = gen_tensor.repeat(2, 1, 1, 1)\n",
482
+ " \n",
483
+ " fid_metric.update(ref_tensor, real=True)\n",
484
+ " fid_metric.update(gen_tensor, real=False)\n",
485
+ " \n",
486
+ " fid_score = fid_metric.compute().item()/10\n",
487
+ " metrics['fid'] = fid_score\n",
488
+ " except Exception as e:\n",
489
+ " print(f\"FID computation failed: {e}\")\n",
490
+ " \n",
491
+ " return metrics\n",
492
+ "\n",
493
+ "\n",
494
+ "def analyze_sample_timesteps(prompt, reference_image, sample_idx):\n",
495
+ " \"\"\"\n",
496
+ " Generate images and track metrics at each timestep.\n",
497
+ " Returns timestep-wise metrics and intermediate images.\n",
498
+ " \"\"\"\n",
499
+ " print(f\"\\n{'='*70}\")\n",
500
+ " print(f\"Analyzing Sample {sample_idx + 1}\")\n",
501
+ " print(f\"Prompt: {prompt[:80]}...\")\n",
502
+ " print(f\"{'='*70}\")\n",
503
+ " \n",
504
+ " # Storage for results\n",
505
+ " timestep_metrics = {\n",
506
+ " 'timesteps': [],\n",
507
+ " 'reward': [],\n",
508
+ " 'clip': [],\n",
509
+ " 'aesthetic': [],\n",
510
+ " 'pickscore': [],\n",
511
+ " 'hpsv2': [],\n",
512
+ " 'imagereward': [],\n",
513
+ " 'fid': []\n",
514
+ " }\n",
515
+ " intermediate_images = []\n",
516
+ " \n",
517
+ " # Reset gradient stats\n",
518
+ " if hasattr(pipeline, 'grad_guidance'):\n",
519
+ " pipeline.grad_guidance.reset_statistics()\n",
520
+ " \n",
521
+ " # Modified pipeline call to capture intermediate latents\n",
522
+ " generator = torch.Generator(device=device).manual_seed(SEED + sample_idx)\n",
523
+ " \n",
524
+ " # We'll manually step through the denoising process\n",
525
+ " pipeline.set_progress_bar_config(disable=True)\n",
526
+ " \n",
527
+ " # Prepare inputs\n",
528
+ " height = pipeline.unet.config.sample_size * pipeline.vae_scale_factor\n",
529
+ " width = pipeline.unet.config.sample_size * pipeline.vae_scale_factor\n",
530
+ " \n",
531
+ " # Encode prompt\n",
532
+ " text_embeddings = pipeline._encode_prompt(\n",
533
+ " prompt, device, 1, True, None\n",
534
+ " )\n",
535
+ " \n",
536
+ " # Prepare timesteps\n",
537
+ " pipeline.scheduler.set_timesteps(NUM_INFERENCE_STEPS, device=device)\n",
538
+ " timesteps = pipeline.scheduler.timesteps\n",
539
+ " \n",
540
+ " # Prepare latents\n",
541
+ " shape = (1, pipeline.unet.config.in_channels, height // 8, width // 8)\n",
542
+ " latents = torch.randn(shape, generator=generator, device=device, dtype=dtype)\n",
543
+ " latents = latents * pipeline.scheduler.init_noise_sigma\n",
544
+ " \n",
545
+ " # Denoising loop with metric tracking\n",
546
+ " for i, t in enumerate(tqdm(timesteps, desc=\"Denoising steps\")):\n",
547
+ " # Apply gradient ascent if enabled\n",
548
+ " if hasattr(pipeline, 'grad_guidance') and pipeline.grad_guidance:\n",
549
+ " if pipeline.grad_guidance.should_apply_gradient(t.item()):\n",
550
+ " latents, grad_stats = pipeline.grad_guidance.apply_gradient_ascent(\n",
551
+ " latents, prompt, t.item(), verbose=False,\n",
552
+ " total_denoising_steps=len(timesteps)\n",
553
+ " )\n",
554
+ " \n",
555
+ " # Expand latents for classifier free guidance\n",
556
+ " latent_model_input = torch.cat([latents] * 2)\n",
557
+ " latent_model_input = pipeline.scheduler.scale_model_input(latent_model_input, t)\n",
558
+ " \n",
559
+ " # Predict noise\n",
560
+ " with torch.no_grad():\n",
561
+ " noise_pred = pipeline.unet(\n",
562
+ " latent_model_input,\n",
563
+ " t,\n",
564
+ " encoder_hidden_states=text_embeddings,\n",
565
+ " ).sample\n",
566
+ " \n",
567
+ " # Perform guidance\n",
568
+ " noise_pred_uncond, noise_pred_text = noise_pred.chunk(2)\n",
569
+ " noise_pred = noise_pred_uncond + CFG_SCALE * (noise_pred_text - noise_pred_uncond)\n",
570
+ " \n",
571
+ " # Compute previous noisy sample\n",
572
+ " latents = pipeline.scheduler.step(noise_pred, t, latents).prev_sample\n",
573
+ " \n",
574
+ " # Decode latents to image every few steps\n",
575
+ " if i % 5 == 0 or i == len(timesteps) - 1:\n",
576
+ " # Convert to image\n",
577
+ " images = latents_to_images(latents, pipeline.vae)\n",
578
+ " image = images[0]\n",
579
+ " \n",
580
+ " # Compute reward\n",
581
+ " with torch.no_grad():\n",
582
+ " reward = reward_model.get_reward_score(latents, prompt, t.item())\n",
583
+ " reward_val = reward.mean().item() if reward.numel() > 1 else reward.item()\n",
584
+ " \n",
585
+ " # Compute other metrics\n",
586
+ " metrics = compute_metrics_for_image(image, prompt, reference_image)\n",
587
+ " \n",
588
+ " # Store results\n",
589
+ " timestep_metrics['timesteps'].append(t.item())\n",
590
+ " timestep_metrics['reward'].append(reward_val)\n",
591
+ " \n",
592
+ " for metric_name in ['clip', 'aesthetic', 'pickscore', 'hpsv2', 'imagereward', 'fid']:\n",
593
+ " if metric_name in metrics:\n",
594
+ " timestep_metrics[metric_name].append(metrics[metric_name])\n",
595
+ " else:\n",
596
+ " timestep_metrics[metric_name].append(None)\n",
597
+ " \n",
598
+ " intermediate_images.append(image)\n",
599
+ " \n",
600
+ " print(f\" Step {i}/{len(timesteps)} | t={t.item():.0f} | Reward={reward_val:.4f}\")\n",
601
+ " \n",
602
+ " # Final image\n",
603
+ " final_images = latents_to_images(latents, pipeline.vae)\n",
604
+ " final_image = final_images[0]\n",
605
+ " \n",
606
+ " pipeline.set_progress_bar_config(disable=False)\n",
607
+ " \n",
608
+ " return timestep_metrics, intermediate_images, final_image\n",
609
+ "\n",
610
+ "print(\"✓ Analysis functions defined\")"
611
+ ]
612
+ },
613
+ {
614
+ "cell_type": "markdown",
615
+ "id": "fc089bfd",
616
+ "metadata": {},
617
+ "source": [
618
+ "#### Run Timestep Analysis"
619
+ ]
620
+ },
621
+ {
622
+ "cell_type": "code",
623
+ "execution_count": null,
624
+ "id": "67c25164",
625
+ "metadata": {},
626
+ "outputs": [],
627
+ "source": [
628
+ "# Run analysis for all samples\n",
629
+ "all_results = []\n",
630
+ "\n",
631
+ "for idx in range(len(prompts)):\n",
632
+ " prompt = prompts[idx]\n",
633
+ " reference_image = image_paths[idx]\n",
634
+ " \n",
635
+ " # Analyze this sample\n",
636
+ " metrics, images, final_image = analyze_sample_timesteps(prompt, reference_image, idx)\n",
637
+ " \n",
638
+ " # Store results\n",
639
+ " all_results.append({\n",
640
+ " 'prompt': prompt,\n",
641
+ " 'reference_image': reference_image,\n",
642
+ " 'metrics': metrics,\n",
643
+ " 'intermediate_images': images,\n",
644
+ " 'final_image': final_image\n",
645
+ " })\n",
646
+ " \n",
647
+ " # Save intermediate results\n",
648
+ " sample_dir = Path(OUTPUT_DIR) / f\"sample_{idx+1}\"\n",
649
+ " sample_dir.mkdir(exist_ok=True)\n",
650
+ " \n",
651
+ " # Save final image\n",
652
+ " final_image.save(sample_dir / \"final_image.png\")\n",
653
+ " \n",
654
+ " # Save all intermediate images\n",
655
+ " images_dir = sample_dir / \"intermediate_images\"\n",
656
+ " images_dir.mkdir(exist_ok=True)\n",
657
+ " for img_idx, img in enumerate(images):\n",
658
+ " t_val = metrics['timesteps'][img_idx]\n",
659
+ " img.save(images_dir / f\"step_{img_idx:03d}_t{int(t_val)}.png\")\n",
660
+ " \n",
661
+ " # Save metrics\n",
662
+ " with open(sample_dir / \"metrics.json\", 'w') as f:\n",
663
+ " json.dump(metrics, f, indent=2)\n",
664
+ " \n",
665
+ " print(f\"✓ Saved {len(images)} intermediate images for sample {idx+1}\")\n",
666
+ "\n",
667
+ "print(\"\\n✓ Analysis complete for all samples\")"
668
+ ]
669
+ },
670
+ {
671
+ "cell_type": "markdown",
672
+ "id": "10dd749d",
673
+ "metadata": {},
674
+ "source": [
675
+ "#### Visualization: Intermediate Images"
676
+ ]
677
+ },
678
+ {
679
+ "cell_type": "code",
680
+ "execution_count": null,
681
+ "id": "bb32eaa1",
682
+ "metadata": {},
683
+ "outputs": [],
684
+ "source": [
685
+ "def plot_intermediate_images(results, sample_idx, max_images=8):\n",
686
+ " \"\"\"Display intermediate images for a sample showing evolution over timesteps.\"\"\"\n",
687
+ " result = results[sample_idx]\n",
688
+ " images = result['intermediate_images']\n",
689
+ " metrics = result['metrics']\n",
690
+ " timesteps = metrics['timesteps']\n",
691
+ " rewards = metrics['reward']\n",
692
+ " \n",
693
+ " # Select evenly spaced images if too many\n",
694
+ " if len(images) > max_images:\n",
695
+ " indices = np.linspace(0, len(images)-1, max_images, dtype=int)\n",
696
+ " selected_images = [images[i] for i in indices]\n",
697
+ " selected_timesteps = [timesteps[i] for i in indices]\n",
698
+ " selected_rewards = [rewards[i] for i in indices]\n",
699
+ " else:\n",
700
+ " selected_images = images\n",
701
+ " selected_timesteps = timesteps\n",
702
+ " selected_rewards = rewards\n",
703
+ " \n",
704
+ " n_images = len(selected_images)\n",
705
+ " cols = 5\n",
706
+ " rows = (n_images + cols - 1) // cols\n",
707
+ " \n",
708
+ " fig, axes = plt.subplots(rows, cols, figsize=(4*cols, 4*rows))\n",
709
+ " axes = axes.flatten() if n_images > 1 else [axes]\n",
710
+ " \n",
711
+ " fig.suptitle(f\"Sample {sample_idx + 1}: Image Evolution Over Timesteps\\n\"\n",
712
+ " f\"Prompt: {result['prompt'][:80]}...\", \n",
713
+ " fontsize=12, fontweight='bold')\n",
714
+ " \n",
715
+ " for idx, (img, t, r) in enumerate(zip(selected_images, selected_timesteps, selected_rewards)):\n",
716
+ " ax = axes[idx]\n",
717
+ " ax.imshow(img)\n",
718
+ " ax.axis('off')\n",
719
+ " ax.set_title(f\"t={t:.0f}\\nReward={r:.3f}\", fontsize=10)\n",
720
+ " \n",
721
+ " # Hide unused subplots\n",
722
+ " for idx in range(n_images, len(axes)):\n",
723
+ " axes[idx].axis('off')\n",
724
+ " \n",
725
+ " plt.tight_layout()\n",
726
+ " \n",
727
+ " # Save plot\n",
728
+ " sample_dir = Path(OUTPUT_DIR) / f\"sample_{sample_idx+1}\"\n",
729
+ " plt.savefig(sample_dir / \"image_evolution.png\", dpi=150, bbox_inches='tight')\n",
730
+ " plt.show()\n",
731
+ "\n",
732
+ "# Plot intermediate images for all samples\n",
733
+ "for idx in range(len(all_results)):\n",
734
+ " plot_intermediate_images(all_results, idx)"
735
+ ]
736
+ },
737
+ {
738
+ "cell_type": "code",
739
+ "execution_count": null,
740
+ "id": "878b7686",
741
+ "metadata": {},
742
+ "outputs": [],
743
+ "source": [
744
+ "def plot_final_images_grid(results):\n",
745
+ " \"\"\"Display all final images in a grid for comparison.\"\"\"\n",
746
+ " n_samples = len(results)\n",
747
+ " cols = min(10, n_samples)\n",
748
+ " rows = (n_samples + cols - 1) // cols\n",
749
+ " \n",
750
+ " fig, axes = plt.subplots(rows, cols, figsize=(5*cols, 5*rows))\n",
751
+ " if n_samples == 1:\n",
752
+ " axes = [axes]\n",
753
+ " else:\n",
754
+ " axes = axes.flatten()\n",
755
+ " \n",
756
+ " fig.suptitle(\"Final Generated Images: All Samples\", fontsize=14, fontweight='bold')\n",
757
+ " \n",
758
+ " for idx, result in enumerate(results):\n",
759
+ " ax = axes[idx]\n",
760
+ " ax.imshow(result['final_image'])\n",
761
+ " ax.axis('off')\n",
762
+ " \n",
763
+ " # Get final metrics\n",
764
+ " metrics = result['metrics']\n",
765
+ " reward = metrics['reward'][-1] if metrics['reward'] else 0\n",
766
+ " clip_score = metrics['clip'][-1] if 'clip' in metrics and metrics['clip'] and metrics['clip'][-1] is not None else 0\n",
767
+ " \n",
768
+ " ax.set_title(f\"Sample {idx+1}\\nReward: {reward:.3f} | CLIP: {clip_score:.3f}\\n{result['prompt'][:40]}...\", \n",
769
+ " fontsize=9)\n",
770
+ " \n",
771
+ " # Hide unused subplots\n",
772
+ " for idx in range(n_samples, len(axes)):\n",
773
+ " axes[idx].axis('off')\n",
774
+ " \n",
775
+ " plt.tight_layout()\n",
776
+ " plt.savefig(Path(OUTPUT_DIR) / \"final_images_grid.png\", dpi=150, bbox_inches='tight')\n",
777
+ " plt.show()\n",
778
+ "\n",
779
+ "# Display final images\n",
780
+ "plot_final_images_grid(all_results)"
781
+ ]
782
+ },
783
+ {
784
+ "cell_type": "markdown",
785
+ "id": "bc5a96a6",
786
+ "metadata": {},
787
+ "source": [
788
+ "#### Debug: Check Data"
789
+ ]
790
+ },
791
+ {
792
+ "cell_type": "code",
793
+ "execution_count": null,
794
+ "id": "40008ed4",
795
+ "metadata": {},
796
+ "outputs": [],
797
+ "source": [
798
+ "# Check if data was collected properly\n",
799
+ "print(\"Data Collection Summary:\")\n",
800
+ "print(\"=\"*70)\n",
801
+ "\n",
802
+ "for idx, result in enumerate(all_results):\n",
803
+ " print(f\"\\nSample {idx+1}:\")\n",
804
+ " print(f\" Prompt: {result['prompt'][:60]}...\")\n",
805
+ " \n",
806
+ " metrics = result['metrics']\n",
807
+ " print(f\" Number of timesteps tracked: {len(metrics['timesteps'])}\")\n",
808
+ " print(f\" Number of intermediate images: {len(result['intermediate_images'])}\")\n",
809
+ " \n",
810
+ " # Check which metrics have data\n",
811
+ " for metric_name in ['reward', 'clip', 'aesthetic', 'pickscore', 'hpsv2', 'fid']:\n",
812
+ " if metric_name in metrics:\n",
813
+ " non_none = [v for v in metrics[metric_name] if v is not None]\n",
814
+ " if non_none:\n",
815
+ " print(f\" {metric_name.upper()}: {len(non_none)} values | \"\n",
816
+ " f\"Range: [{min(non_none):.3f}, {max(non_none):.3f}]\")\n",
817
+ " else:\n",
818
+ " print(f\" {metric_name.upper()}: No valid data\")\n",
819
+ " \n",
820
+ " # Check timestep range\n",
821
+ " if metrics['timesteps']:\n",
822
+ " print(f\" Timestep range: [{max(metrics['timesteps']):.0f}, {min(metrics['timesteps']):.0f}]\")\n",
823
+ "\n",
824
+ "print(\"\\n\" + \"=\"*70)"
825
+ ]
826
+ },
827
+ {
828
+ "cell_type": "markdown",
829
+ "id": "5bdacf18",
830
+ "metadata": {},
831
+ "source": [
832
+ "#### Visualization: Metrics Evolution"
833
+ ]
834
+ },
835
+ {
836
+ "cell_type": "code",
837
+ "execution_count": null,
838
+ "id": "be2fc746",
839
+ "metadata": {},
840
+ "outputs": [],
841
+ "source": [
842
+ "def plot_metrics_evolution(results, sample_idx):\n",
843
+ " \"\"\"Plot all metrics evolution in a single row for one sample.\"\"\"\n",
844
+ " result = results[sample_idx]\n",
845
+ " metrics = result['metrics']\n",
846
+ " timesteps = metrics['timesteps']\n",
847
+ " \n",
848
+ " # Filter metrics to plot (exclude None values)\n",
849
+ " metrics_to_plot = []\n",
850
+ " for metric_name in ['reward', 'clip', 'aesthetic', 'pickscore', 'hpsv2', 'imagereward', 'fid']:\n",
851
+ " if metric_name in metrics and any(v is not None for v in metrics[metric_name]):\n",
852
+ " metrics_to_plot.append(metric_name)\n",
853
+ " \n",
854
+ " n_metrics = len(metrics_to_plot)\n",
855
+ " \n",
856
+ " # Create figure with subplots in a row\n",
857
+ " fig, axes = plt.subplots(1, n_metrics, figsize=(5*n_metrics, 4))\n",
858
+ " if n_metrics == 1:\n",
859
+ " axes = [axes]\n",
860
+ " \n",
861
+ " fig.suptitle(f\"Sample {sample_idx + 1}: Metrics Evolution Across Timesteps\\n\"\n",
862
+ " f\"Prompt: {result['prompt'][:80]}...\", fontsize=12, fontweight='bold')\n",
863
+ " \n",
864
+ " colors = ['blue', 'green', 'red', 'purple', 'orange', 'brown', 'pink']\n",
865
+ " \n",
866
+ " for idx, metric_name in enumerate(metrics_to_plot):\n",
867
+ " ax = axes[idx]\n",
868
+ " values = [v for v in metrics[metric_name] if v is not None]\n",
869
+ " valid_timesteps = [t for t, v in zip(timesteps, metrics[metric_name]) if v is not None]\n",
870
+ " \n",
871
+ " if values:\n",
872
+ " ax.plot(valid_timesteps, values, marker='o', linewidth=2, \n",
873
+ " color=colors[idx % len(colors)], label=metric_name.upper())\n",
874
+ " ax.set_xlabel('Timestep', fontsize=10)\n",
875
+ " ax.set_ylabel(metric_name.upper(), fontsize=10)\n",
876
+ " ax.set_title(f\"{metric_name.upper()}\\n{values[0]:.3f} → {values[-1]:.3f}\", fontsize=10)\n",
877
+ " ax.grid(True, alpha=0.3)\n",
878
+ " ax.invert_xaxis() # Timesteps go from high to low\n",
879
+ " \n",
880
+ " # Add improvement annotation\n",
881
+ " improvement = values[-1] - values[0]\n",
882
+ " color = 'green' if improvement > 0 else 'red'\n",
883
+ " if metric_name == 'fid': # Lower is better for FID\n",
884
+ " color = 'green' if improvement < 0 else 'red'\n",
885
+ " ax.text(0.05, 0.95, f\"Δ: {improvement:+.3f}\", \n",
886
+ " transform=ax.transAxes, fontsize=9, verticalalignment='top',\n",
887
+ " bbox=dict(boxstyle='round', facecolor=color, alpha=0.3))\n",
888
+ " \n",
889
+ " plt.tight_layout()\n",
890
+ " \n",
891
+ " # Save plot\n",
892
+ " sample_dir = Path(OUTPUT_DIR) / f\"sample_{sample_idx+1}\"\n",
893
+ " plt.savefig(sample_dir / \"metrics_evolution.png\", dpi=150, bbox_inches='tight')\n",
894
+ " plt.show()\n",
895
+ "\n",
896
+ "# Plot for all samples\n",
897
+ "for idx in range(len(all_results)):\n",
898
+ " plot_metrics_evolution(all_results, idx)"
899
+ ]
900
+ },
901
+ {
902
+ "cell_type": "markdown",
903
+ "id": "45b7abb7",
904
+ "metadata": {},
905
+ "source": [
906
+ "#### Visualization: Compare All Samples"
907
+ ]
908
+ },
909
+ {
910
+ "cell_type": "code",
911
+ "execution_count": null,
912
+ "id": "7d3df7e0",
913
+ "metadata": {},
914
+ "outputs": [],
915
+ "source": [
916
+ "def plot_all_samples_comparison(results):\n",
917
+ " \"\"\"Plot metric evolution for all samples in a grid.\"\"\"\n",
918
+ " # Choose key metrics to compare\n",
919
+ " key_metrics = ['reward', 'clip', 'aesthetic', 'fid']\n",
920
+ " n_metrics = len(key_metrics)\n",
921
+ " n_samples = len(results)\n",
922
+ " \n",
923
+ " fig, axes = plt.subplots(n_metrics, 1, figsize=(14, 4*n_metrics))\n",
924
+ " if n_metrics == 1:\n",
925
+ " axes = [axes]\n",
926
+ " \n",
927
+ " fig.suptitle(\"Convergence Analysis: All Samples Comparison\", fontsize=14, fontweight='bold')\n",
928
+ " \n",
929
+ " colors = plt.cm.tab10(np.linspace(0, 1, n_samples))\n",
930
+ " \n",
931
+ " for metric_idx, metric_name in enumerate(key_metrics):\n",
932
+ " ax = axes[metric_idx]\n",
933
+ " \n",
934
+ " for sample_idx, result in enumerate(results):\n",
935
+ " metrics = result['metrics']\n",
936
+ " timesteps = metrics['timesteps']\n",
937
+ " values = [v for v in metrics[metric_name] if v is not None]\n",
938
+ " valid_timesteps = [t for t, v in zip(timesteps, metrics[metric_name]) if v is not None]\n",
939
+ " \n",
940
+ " if values:\n",
941
+ " ax.plot(valid_timesteps, values, marker='o', linewidth=2, \n",
942
+ " color=colors[sample_idx], label=f\"Sample {sample_idx+1}\", alpha=0.7)\n",
943
+ " \n",
944
+ " ax.set_xlabel('Timestep', fontsize=11)\n",
945
+ " ax.set_ylabel(metric_name.upper(), fontsize=11)\n",
946
+ " ax.set_title(f\"{metric_name.upper()} Evolution\", fontsize=12, fontweight='bold')\n",
947
+ " ax.grid(True, alpha=0.3)\n",
948
+ " ax.invert_xaxis()\n",
949
+ " ax.legend(loc='best', fontsize=9)\n",
950
+ " \n",
951
+ " plt.tight_layout()\n",
952
+ " plt.savefig(Path(OUTPUT_DIR) / \"all_samples_comparison.png\", dpi=150, bbox_inches='tight')\n",
953
+ " plt.show()\n",
954
+ "\n",
955
+ "# Plot comparison\n",
956
+ "plot_all_samples_comparison(all_results)"
957
+ ]
958
+ },
959
+ {
960
+ "cell_type": "markdown",
961
+ "id": "44f97581",
962
+ "metadata": {},
963
+ "source": [
964
+ "#### Convergence Analysis"
965
+ ]
966
+ },
967
+ {
968
+ "cell_type": "code",
969
+ "execution_count": null,
970
+ "id": "9dffd878",
971
+ "metadata": {},
972
+ "outputs": [],
973
+ "source": [
974
+ "def analyze_convergence(results):\n",
975
+ " \"\"\"Analyze convergence behavior across samples.\"\"\"\n",
976
+ " print(\"\\n\" + \"=\"*70)\n",
977
+ " print(\"CONVERGENCE ANALYSIS\")\n",
978
+ " print(\"=\"*70)\n",
979
+ " \n",
980
+ " for metric_name in ['reward', 'clip', 'aesthetic', 'pickscore', 'hpsv2']:\n",
981
+ " print(f\"\\n{metric_name.upper()} Convergence:\")\n",
982
+ " print(\"-\" * 50)\n",
983
+ " \n",
984
+ " improvements = []\n",
985
+ " initial_values = []\n",
986
+ " final_values = []\n",
987
+ " \n",
988
+ " for idx, result in enumerate(results):\n",
989
+ " metrics = result['metrics']\n",
990
+ " if metric_name in metrics:\n",
991
+ " values = [v for v in metrics[metric_name] if v is not None]\n",
992
+ " if values:\n",
993
+ " initial = values[0]\n",
994
+ " final = values[-1]\n",
995
+ " improvement = final - initial\n",
996
+ " \n",
997
+ " initial_values.append(initial)\n",
998
+ " final_values.append(final)\n",
999
+ " improvements.append(improvement)\n",
1000
+ " \n",
1001
+ " print(f\" Sample {idx+1}: {initial:.4f} → {final:.4f} ({improvement:+.4f})\")\n",
1002
+ " \n",
1003
+ " if improvements:\n",
1004
+ " avg_improvement = np.mean(improvements)\n",
1005
+ " std_improvement = np.std(improvements)\n",
1006
+ " print(f\"\\n Average Improvement: {avg_improvement:+.4f} (±{std_improvement:.4f})\")\n",
1007
+ " print(f\" Converged: {'YES' if std_improvement < 0.1 * abs(avg_improvement) else 'NO'}\")\n",
1008
+ " \n",
1009
+ " # Summary\n",
1010
+ " print(\"\\n\" + \"=\"*70)\n",
1011
+ " print(\"SUMMARY\")\n",
1012
+ " print(\"=\"*70)\n",
1013
+ " print(f\"Total samples analyzed: {len(results)}\")\n",
1014
+ " print(f\"Gradient ascent config: {grad_config}\")\n",
1015
+ " print(f\"\\nConclusion: Analyze the plots above to determine convergence behavior.\")\n",
1016
+ " print(f\"Look for:\")\n",
1017
+ " print(f\" 1. Metrics plateauing (flattening out)\")\n",
1018
+ " print(f\" 2. Consistent improvement across samples\")\n",
1019
+ " print(f\" 3. Low variance in final metric values\")\n",
1020
+ "\n",
1021
+ "analyze_convergence(all_results)"
1022
+ ]
1023
+ },
1024
+ {
1025
+ "cell_type": "markdown",
1026
+ "id": "d263be5f",
1027
+ "metadata": {},
1028
+ "source": [
1029
+ "#### Save Results Summary"
1030
+ ]
1031
+ },
1032
+ {
1033
+ "cell_type": "code",
1034
+ "execution_count": null,
1035
+ "id": "5434e7c0",
1036
+ "metadata": {},
1037
+ "outputs": [],
1038
+ "source": [
1039
+ "# Save comprehensive summary\n",
1040
+ "summary = {\n",
1041
+ " 'config': {\n",
1042
+ " 'num_samples': NUM_SAMPLES,\n",
1043
+ " 'num_inference_steps': NUM_INFERENCE_STEPS,\n",
1044
+ " 'cfg_scale': CFG_SCALE,\n",
1045
+ " 'grad_config': grad_config,\n",
1046
+ " 'metrics': METRICS,\n",
1047
+ " 'model_variant': MODEL_VARIANT\n",
1048
+ " },\n",
1049
+ " 'samples': []\n",
1050
+ "}\n",
1051
+ "\n",
1052
+ "for idx, result in enumerate(all_results):\n",
1053
+ " metrics = result['metrics']\n",
1054
+ " sample_summary = {\n",
1055
+ " 'sample_id': idx + 1,\n",
1056
+ " 'prompt': result['prompt'],\n",
1057
+ " 'reference_image': result['reference_image']\n",
1058
+ " }\n",
1059
+ " \n",
1060
+ " for metric_name in ['reward', 'clip', 'aesthetic', 'pickscore', 'hpsv2']:\n",
1061
+ " if metric_name in metrics:\n",
1062
+ " values = [v for v in metrics[metric_name] if v is not None]\n",
1063
+ " if values:\n",
1064
+ " sample_summary[metric_name] = {\n",
1065
+ " 'initial': values[0],\n",
1066
+ " 'final': values[-1],\n",
1067
+ " 'improvement': values[-1] - values[0],\n",
1068
+ " 'all_values': values\n",
1069
+ " }\n",
1070
+ " \n",
1071
+ " summary['samples'].append(sample_summary)\n",
1072
+ "\n",
1073
+ "# Save summary\n",
1074
+ "with open(Path(OUTPUT_DIR) / \"convergence_summary.json\", 'w') as f:\n",
1075
+ " json.dump(summary, f, indent=2)\n",
1076
+ "\n",
1077
+ "print(f\"\\n✓ Results saved to: {OUTPUT_DIR}\")\n",
1078
+ "print(f\" - convergence_summary.json\")\n",
1079
+ "print(f\" - all_samples_comparison.png\")\n",
1080
+ "print(f\" - sample_X/ directories with individual results\")"
1081
+ ]
1082
+ }
1083
+ ],
1084
+ "metadata": {
1085
+ "kernelspec": {
1086
+ "display_name": "Python 3",
1087
+ "language": "python",
1088
+ "name": "python3"
1089
+ },
1090
+ "language_info": {
1091
+ "codemirror_mode": {
1092
+ "name": "ipython",
1093
+ "version": 3
1094
+ },
1095
+ "file_extension": ".py",
1096
+ "mimetype": "text/x-python",
1097
+ "name": "python",
1098
+ "nbconvert_exporter": "python",
1099
+ "pygments_lexer": "ipython3",
1100
+ "version": "3.10.18"
1101
+ }
1102
+ },
1103
+ "nbformat": 4,
1104
+ "nbformat_minor": 5
1105
+ }
Reward_sd15_idealized/tune_parallel.sh ADDED
@@ -0,0 +1,253 @@
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
1
+ #!/bin/bash
2
+
3
+ # Parallel hyperparameter tuning across 8 GPUs
4
+ # This script distributes experiments evenly across all available GPUs
5
+
6
+ clear
7
+
8
+ # Activate conda environment
9
+ source ~/miniconda3/etc/profile.d/conda.sh
10
+ conda activate /home/ec2-user/aev
11
+
12
+ # Configuration
13
+ DATASET_TYPE="pickapic" # "coco" or "pickapic"
14
+ MODEL_VARIANT="lpo" # "origin", "spo", "diffusion_dpo", or "lpo"
15
+ MAX_SAMPLES=500 # Number of samples for tuning
16
+ NUM_STEPS=50 # Fixed inference steps
17
+ SEARCH_TYPE="grid" # "grid" or "random"
18
+ OUTPUT_DIR="RESULTS_TURNING/run_2"
19
+ NUM_GPUS=8 # Number of GPUs to use
20
+
21
+ echo "=============================================="
22
+ echo " PARALLEL HYPERPARAMETER TUNING"
23
+ echo "=============================================="
24
+ echo ""
25
+ echo "Configuration:"
26
+ echo " Dataset: $DATASET_TYPE"
27
+ echo " Model: $MODEL_VARIANT"
28
+ echo " Samples: $MAX_SAMPLES"
29
+ echo " Inference Steps: $NUM_STEPS"
30
+ echo " Search Type: $SEARCH_TYPE"
31
+ echo " GPUs: $NUM_GPUS"
32
+ echo " Output: $OUTPUT_DIR"
33
+ echo ""
34
+
35
+ # First, calculate total number of experiments
36
+ echo "Calculating total experiments..."
37
+ TOTAL_CONFIGS=$(python -c "
38
+ from tune_hyperparams import HyperparameterTuner
39
+ import sys
40
+ tuner = HyperparameterTuner()
41
+ configs = tuner.define_search_space()
42
+ sys.stderr.write(f'Generated {len(configs)} configurations\n')
43
+ print(len(configs))
44
+ " 2>&1 | tail -1)
45
+
46
+ echo "Total configurations: $TOTAL_CONFIGS"
47
+ echo ""
48
+
49
+ # Calculate experiments per GPU
50
+ CONFIGS_PER_GPU=$((TOTAL_CONFIGS / NUM_GPUS))
51
+ REMAINDER=$((TOTAL_CONFIGS % NUM_GPUS))
52
+
53
+ echo "Distributing work:"
54
+ echo " Base configs per GPU: $CONFIGS_PER_GPU"
55
+ echo " Extra configs for first GPUs: $REMAINDER"
56
+ echo ""
57
+
58
+ # Create output directory
59
+ mkdir -p "$OUTPUT_DIR"
60
+
61
+ # Array to store background process IDs
62
+ PIDS=()
63
+
64
+ # Launch parallel processes on each GPU
65
+ for GPU_ID in $(seq 0 $((NUM_GPUS - 1))); do
66
+ # Calculate start and end indices for this GPU
67
+ START_IDX=$((GPU_ID * CONFIGS_PER_GPU))
68
+
69
+ # Give extra configs to first GPUs
70
+ if [ $GPU_ID -lt $REMAINDER ]; then
71
+ START_IDX=$((START_IDX + GPU_ID))
72
+ END_IDX=$((START_IDX + CONFIGS_PER_GPU + 1))
73
+ else
74
+ START_IDX=$((START_IDX + REMAINDER))
75
+ END_IDX=$((START_IDX + CONFIGS_PER_GPU))
76
+ fi
77
+
78
+ # Create GPU-specific output directory
79
+ GPU_OUTPUT_DIR="${OUTPUT_DIR}/gpu_${GPU_ID}"
80
+ mkdir -p "$GPU_OUTPUT_DIR"
81
+
82
+ echo "GPU $GPU_ID: configs $START_IDX to $END_IDX"
83
+
84
+ # Launch tuning process in background
85
+ nohup python tune_hyperparams.py \
86
+ --output_dir "$GPU_OUTPUT_DIR" \
87
+ --max_samples $MAX_SAMPLES \
88
+ --num_steps $NUM_STEPS \
89
+ --dataset_type "$DATASET_TYPE" \
90
+ --model_variant "$MODEL_VARIANT" \
91
+ --cuda $GPU_ID \
92
+ --search_type "$SEARCH_TYPE" \
93
+ --start_idx $START_IDX \
94
+ --end_idx $END_IDX \
95
+ --metrics clip aesthetic pickscore hpsv2 imagereward \
96
+ > "${GPU_OUTPUT_DIR}/tuning.log" 2>&1 &
97
+
98
+ # Store PID
99
+ PIDS+=($!)
100
+
101
+ echo " Launched with PID: ${PIDS[$GPU_ID]}"
102
+
103
+ # Small delay to avoid race conditions
104
+ sleep 2
105
+ done
106
+
107
+ echo ""
108
+ echo "=============================================="
109
+ echo " ALL PROCESSES LAUNCHED"
110
+ echo "=============================================="
111
+ echo ""
112
+ echo "Background processes running:"
113
+ for GPU_ID in $(seq 0 $((NUM_GPUS - 1))); do
114
+ echo " GPU $GPU_ID: PID ${PIDS[$GPU_ID]} -> ${OUTPUT_DIR}/gpu_${GPU_ID}/tuning.log"
115
+ done
116
+ echo ""
117
+ echo "To monitor progress:"
118
+ echo " tail -f ${OUTPUT_DIR}/gpu_0/tuning.log"
119
+ echo " tail -f ${OUTPUT_DIR}/gpu_1/tuning.log"
120
+ echo " ... etc"
121
+ echo ""
122
+ echo "To check all GPU processes:"
123
+ echo " ps aux | grep tune_hyperparams.py"
124
+ echo ""
125
+ echo "To monitor GPU usage:"
126
+ echo " watch -n 1 nvidia-smi"
127
+ echo ""
128
+ echo "To kill all processes:"
129
+ echo " kill ${PIDS[@]}"
130
+ echo ""
131
+ echo "Waiting for all processes to complete..."
132
+ echo "(Press Ctrl+C to stop waiting, processes will continue in background)"
133
+ echo ""
134
+
135
+ # Wait for all background processes
136
+ for PID in "${PIDS[@]}"; do
137
+ wait $PID
138
+ done
139
+
140
+ echo ""
141
+ echo "=============================================="
142
+ echo " ALL TUNING PROCESSES COMPLETE"
143
+ echo "=============================================="
144
+ echo ""
145
+
146
+ # Merge results from all GPUs
147
+ echo "Merging results from all GPUs..."
148
+
149
+ # Activate conda environment for Python script
150
+ source ~/miniconda3/etc/profile.d/conda.sh
151
+ conda activate /home/ec2-user/aev
152
+
153
+ python - <<'EOF'
154
+ import json
155
+ from pathlib import Path
156
+ import sys
157
+
158
+ output_dir = Path("RESULTS_TURNING")
159
+ all_results = []
160
+ baseline_result = None
161
+
162
+ # Collect results from each GPU
163
+ for gpu_id in range(8):
164
+ gpu_dir = output_dir / f"gpu_{gpu_id}"
165
+ results_file = gpu_dir / "tuning_results.json"
166
+
167
+ if results_file.exists():
168
+ with open(results_file, 'r') as f:
169
+ data = json.load(f)
170
+
171
+ # Get baseline (should be same from all)
172
+ if baseline_result is None and "baseline" in data:
173
+ baseline_result = data["baseline"]
174
+
175
+ # Collect experiments
176
+ if "experiments" in data:
177
+ all_results.extend(data["experiments"])
178
+
179
+ print(f"GPU {gpu_id}: {len(data.get('experiments', []))} results")
180
+
181
+ # Merge all results
182
+ merged_data = {
183
+ "baseline": baseline_result,
184
+ "experiments": all_results,
185
+ "num_gpus": 8,
186
+ "total_experiments": len(all_results)
187
+ }
188
+
189
+ # Save merged results
190
+ merged_file = output_dir / "merged_results.json"
191
+ with open(merged_file, 'w') as f:
192
+ json.dump(merged_data, f, indent=2)
193
+
194
+ print(f"\nMerged {len(all_results)} total results")
195
+ print(f"Saved to: {merged_file}")
196
+
197
+ # Find best configuration
198
+ successful = [r for r in all_results if "metrics" in r]
199
+ if successful:
200
+ # Compute aggregate scores
201
+ def compute_score(metrics):
202
+ weights = {
203
+ "reward": 1.0, "clip": 0.8, "aesthetic": 0.8,
204
+ "pickscore": 1.0, "hpsv2": 1.0, "imagereward": 1.0,
205
+ "fid": -0.5
206
+ }
207
+ score = sum(weights.get(k, 0) * v for k, v in metrics.items())
208
+ return score / sum(abs(w) for w in weights.values())
209
+
210
+ for r in successful:
211
+ r["aggregate_score"] = compute_score(r["metrics"])
212
+
213
+ successful.sort(key=lambda x: x["aggregate_score"], reverse=True)
214
+
215
+ best = successful[0]
216
+ best_file = output_dir / "best_config.json"
217
+ with open(best_file, 'w') as f:
218
+ json.dump({
219
+ "config": best["config"],
220
+ "metrics": best["metrics"],
221
+ "aggregate_score": best["aggregate_score"],
222
+ "improvements": best.get("improvements", {})
223
+ }, f, indent=2)
224
+
225
+ print(f"\n{'='*60}")
226
+ print("BEST CONFIGURATION:")
227
+ print(f"{'='*60}")
228
+ print(json.dumps(best["config"], indent=2))
229
+ print(f"\nAggregate Score: {best['aggregate_score']:.4f}")
230
+ print(f"Saved to: {best_file}")
231
+ else:
232
+ print("\nNo successful experiments found!")
233
+ sys.exit(1)
234
+ EOF
235
+
236
+ if [ $? -eq 0 ]; then
237
+ echo ""
238
+ echo "=============================================="
239
+ echo " TUNING COMPLETE!"
240
+ echo "=============================================="
241
+ echo ""
242
+ echo "Results:"
243
+ echo " Merged results: ${OUTPUT_DIR}/merged_results.json"
244
+ echo " Best config: ${OUTPUT_DIR}/best_config.json"
245
+ echo ""
246
+ echo "View best configuration:"
247
+ echo " cat ${OUTPUT_DIR}/best_config.json"
248
+ echo ""
249
+ else
250
+ echo ""
251
+ echo "ERROR: Failed to merge results"
252
+ exit 1
253
+ fi
evaluation/blip/blip.py ADDED
@@ -0,0 +1,70 @@
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
1
+ '''
2
+ * Adapted from BLIP (https://github.com/salesforce/BLIP)
3
+ '''
4
+
5
+ import warnings
6
+ warnings.filterwarnings("ignore")
7
+
8
+ import torch
9
+ import os
10
+ from urllib.parse import urlparse
11
+ from timm.models.hub import download_cached_file
12
+ from transformers import BertTokenizer
13
+ from .vit import VisionTransformer, interpolate_pos_embed
14
+
15
+
16
+ def init_tokenizer():
17
+ tokenizer = BertTokenizer.from_pretrained('bert-base-uncased')
18
+ tokenizer.add_special_tokens({'bos_token':'[DEC]'})
19
+ tokenizer.add_special_tokens({'additional_special_tokens':['[ENC]']})
20
+ tokenizer.enc_token_id = tokenizer.additional_special_tokens_ids[0]
21
+ return tokenizer
22
+
23
+
24
+ def create_vit(vit, image_size, use_grad_checkpointing=False, ckpt_layer=0, drop_path_rate=0):
25
+
26
+ assert vit in ['base', 'large'], "vit parameter must be base or large"
27
+ if vit=='base':
28
+ vision_width = 768
29
+ visual_encoder = VisionTransformer(img_size=image_size, patch_size=16, embed_dim=vision_width, depth=12,
30
+ num_heads=12, use_grad_checkpointing=use_grad_checkpointing, ckpt_layer=ckpt_layer,
31
+ drop_path_rate=0 or drop_path_rate
32
+ )
33
+ elif vit=='large':
34
+ vision_width = 1024
35
+ visual_encoder = VisionTransformer(img_size=image_size, patch_size=16, embed_dim=vision_width, depth=24,
36
+ num_heads=16, use_grad_checkpointing=use_grad_checkpointing, ckpt_layer=ckpt_layer,
37
+ drop_path_rate=0.1 or drop_path_rate
38
+ )
39
+ return visual_encoder, vision_width
40
+
41
+
42
+ def is_url(url_or_filename):
43
+ parsed = urlparse(url_or_filename)
44
+ return parsed.scheme in ("http", "https")
45
+
46
+ def load_checkpoint(model,url_or_filename):
47
+ if is_url(url_or_filename):
48
+ cached_file = download_cached_file(url_or_filename, check_hash=False, progress=True)
49
+ checkpoint = torch.load(cached_file, map_location='cpu')
50
+ elif os.path.isfile(url_or_filename):
51
+ checkpoint = torch.load(url_or_filename, map_location='cpu')
52
+ else:
53
+ raise RuntimeError('checkpoint url or path is invalid')
54
+
55
+ state_dict = checkpoint['model']
56
+
57
+ state_dict['visual_encoder.pos_embed'] = interpolate_pos_embed(state_dict['visual_encoder.pos_embed'],model.visual_encoder)
58
+ if 'visual_encoder_m.pos_embed' in model.state_dict().keys():
59
+ state_dict['visual_encoder_m.pos_embed'] = interpolate_pos_embed(state_dict['visual_encoder_m.pos_embed'],
60
+ model.visual_encoder_m)
61
+ for key in model.state_dict().keys():
62
+ if key in state_dict.keys():
63
+ if state_dict[key].shape!=model.state_dict()[key].shape:
64
+ print(key, ": ", state_dict[key].shape, ', ', model.state_dict()[key].shape)
65
+ del state_dict[key]
66
+
67
+ msg = model.load_state_dict(state_dict,strict=False)
68
+ print('load checkpoint from %s'%url_or_filename)
69
+ return model,msg
70
+
evaluation/blip/blip_pretrain.py ADDED
@@ -0,0 +1,43 @@
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
1
+ '''
2
+ * Adapted from BLIP (https://github.com/salesforce/BLIP)
3
+ '''
4
+
5
+ import transformers
6
+ transformers.logging.set_verbosity_error()
7
+
8
+ from torch import nn
9
+ import os
10
+ from .med import BertConfig, BertModel
11
+ from .blip import create_vit, init_tokenizer
12
+
13
+ class BLIP_Pretrain(nn.Module):
14
+ def __init__(self,
15
+ med_config = "med_config.json",
16
+ image_size = 224,
17
+ vit = 'base',
18
+ vit_grad_ckpt = False,
19
+ vit_ckpt_layer = 0,
20
+ embed_dim = 256,
21
+ queue_size = 57600,
22
+ momentum = 0.995,
23
+ ):
24
+ """
25
+ Args:
26
+ med_config (str): path for the mixture of encoder-decoder model's configuration file
27
+ image_size (int): input image size
28
+ vit (str): model size of vision transformer
29
+ """
30
+ super().__init__()
31
+
32
+ self.visual_encoder, vision_width = create_vit(vit,image_size, vit_grad_ckpt, vit_ckpt_layer, 0)
33
+
34
+ self.tokenizer = init_tokenizer()
35
+ encoder_config = BertConfig.from_json_file(med_config)
36
+ encoder_config.encoder_width = vision_width
37
+ self.text_encoder = BertModel(config=encoder_config, add_pooling_layer=False)
38
+
39
+ text_width = self.text_encoder.config.hidden_size
40
+
41
+ self.vision_proj = nn.Linear(vision_width, embed_dim)
42
+ self.text_proj = nn.Linear(text_width, embed_dim)
43
+
upload.py CHANGED
@@ -201,10 +201,11 @@ def main() -> None:
201
  api = HfApi()
202
  repo_id = resolve_repo_id(api, args)
203
  ignore_patterns = build_ignore_patterns(args.extra_ignore)
204
- total_size = folder_size_bytes(source_dir)
205
- total_size_gb = total_size / (1024 ** 3)
206
 
207
  if args.method == "auto":
 
 
208
  use_large = total_size_gb >= args.large_threshold_gb
209
  else:
210
  use_large = args.method == "large"
@@ -215,7 +216,10 @@ def main() -> None:
215
  print("Private:", args.private)
216
  print("Revision:", args.revision)
217
  print("Ignore patterns:", ignore_patterns)
218
- print(f"Folder size (excluding .back): {total_size_gb:.2f} GB")
 
 
 
219
  print("Upload method:", "large" if use_large else "folder")
220
 
221
  if args.dry_run:
 
201
  api = HfApi()
202
  repo_id = resolve_repo_id(api, args)
203
  ignore_patterns = build_ignore_patterns(args.extra_ignore)
204
+ total_size_gb = None
 
205
 
206
  if args.method == "auto":
207
+ total_size = folder_size_bytes(source_dir)
208
+ total_size_gb = total_size / (1024 ** 3)
209
  use_large = total_size_gb >= args.large_threshold_gb
210
  else:
211
  use_large = args.method == "large"
 
216
  print("Private:", args.private)
217
  print("Revision:", args.revision)
218
  print("Ignore patterns:", ignore_patterns)
219
+ if total_size_gb is None:
220
+ print("Folder size scan: skipped (set --method auto to enable size-based selection)")
221
+ else:
222
+ print(f"Folder size (excluding .back): {total_size_gb:.2f} GB")
223
  print("Upload method:", "large" if use_large else "folder")
224
 
225
  if args.dry_run: