haritetala commited on
Commit
5f5b8d2
·
verified ·
1 Parent(s): 587073f

Upload pipeline.py

Browse files
Files changed (1) hide show
  1. pipeline.py +978 -0
pipeline.py ADDED
@@ -0,0 +1,978 @@
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
1
+ #!/usr/bin/env python3
2
+ """
3
+ Standalone CV Synthetic Data Engine
4
+
5
+ This module implements an API-first synthetic data pipeline for few-shot object
6
+ conditioning, prompt-driven synthetic frame generation, Grounding DINO zero-shot
7
+ annotation, and dataset compilation for YOLO or COCO-style training workflows.
8
+
9
+ The script is intentionally self-contained. It can be used directly from a
10
+ GPU host, wrapped by Modal serverless functions, or launched inside RunPod.
11
+ """
12
+
13
+ from __future__ import annotations
14
+
15
+ import argparse
16
+ import gc
17
+ import json
18
+ import logging
19
+ import math
20
+ import os
21
+ import random
22
+ import shutil
23
+ import sys
24
+ import time
25
+ from dataclasses import asdict, dataclass, field
26
+ from pathlib import Path
27
+ from typing import Any, Dict, Iterable, List, Literal, Optional, Sequence, Tuple, Union
28
+
29
+ import numpy as np
30
+ from PIL import Image, ImageEnhance, ImageFilter, ImageOps
31
+
32
+ try:
33
+ import cv2
34
+ except Exception as exc: # pragma: no cover - dependency guard
35
+ cv2 = None
36
+ _CV2_IMPORT_ERROR = exc
37
+ else:
38
+ _CV2_IMPORT_ERROR = None
39
+
40
+ try:
41
+ import torch
42
+ import torch.nn.functional as F
43
+ from torch.utils.data import DataLoader, Dataset
44
+ except Exception as exc: # pragma: no cover - dependency guard
45
+ torch = None
46
+ F = None
47
+ DataLoader = object
48
+ Dataset = object
49
+ _TORCH_IMPORT_ERROR = exc
50
+ else:
51
+ _TORCH_IMPORT_ERROR = None
52
+
53
+ try:
54
+ from diffusers import DDPMScheduler, StableDiffusionPipeline, StableDiffusionXLPipeline
55
+ from peft.utils import get_peft_model_state_dict
56
+ except Exception as exc: # pragma: no cover - dependency guard
57
+ DDPMScheduler = None
58
+ StableDiffusionPipeline = None
59
+ StableDiffusionXLPipeline = None
60
+ get_peft_model_state_dict = None
61
+ _DIFFUSERS_IMPORT_ERROR = exc
62
+ else:
63
+ _DIFFUSERS_IMPORT_ERROR = None
64
+
65
+ try:
66
+ from peft import LoraConfig
67
+ except Exception as exc: # pragma: no cover - dependency guard
68
+ LoraConfig = None
69
+ _PEFT_IMPORT_ERROR = exc
70
+ else:
71
+ _PEFT_IMPORT_ERROR = None
72
+
73
+ try:
74
+ from transformers import AutoModelForZeroShotObjectDetection, AutoProcessor
75
+ except Exception as exc: # pragma: no cover - dependency guard
76
+ AutoModelForZeroShotObjectDetection = None
77
+ AutoProcessor = None
78
+ _TRANSFORMERS_IMPORT_ERROR = exc
79
+ else:
80
+ _TRANSFORMERS_IMPORT_ERROR = None
81
+
82
+ try:
83
+ import modal
84
+ except Exception: # pragma: no cover - optional platform integration
85
+ modal = None
86
+
87
+ LOGGER = logging.getLogger("synthetic_cv_pipeline")
88
+
89
+ LabelFormat = Literal["yolo", "coco"]
90
+ ImageLike = Union[str, Path, Image.Image, np.ndarray]
91
+
92
+
93
+ @dataclass
94
+ class TrainingConfig:
95
+ """Configuration for high-velocity few-shot LoRA optimization."""
96
+
97
+ pretrained_model: str = "runwayml/stable-diffusion-v1-5"
98
+ output_dir: str = "conditioned_lora"
99
+ instance_token: str = "sksobj"
100
+ resolution: int = 512
101
+ train_steps: int = 180
102
+ validation_interval: int = 20
103
+ patience: int = 4
104
+ min_delta: float = 0.0025
105
+ learning_rate: float = 1e-4
106
+ batch_size: int = 1
107
+ gradient_accumulation_steps: int = 1
108
+ rank: int = 8
109
+ seed: int = 1337
110
+ mixed_precision: Literal["fp16", "bf16", "no"] = "fp16"
111
+ negative_prompt: str = "low quality, blurry, warped object, extra object, text, watermark"
112
+ num_validation_images: int = 1
113
+ train_text_encoder_lora: bool = False
114
+ max_grad_norm: float = 1.0
115
+
116
+
117
+ @dataclass
118
+ class SynthesisConfig:
119
+ """Configuration for conditioned batch synthesis."""
120
+
121
+ prompt: str
122
+ target_object: str
123
+ output_dir: str = "output_batch"
124
+ lora_dir: Optional[str] = "conditioned_lora"
125
+ pretrained_model: str = "runwayml/stable-diffusion-v1-5"
126
+ num_images: int = 24
127
+ batch_size: int = 1
128
+ width: int = 512
129
+ height: int = 512
130
+ inference_steps: int = 32
131
+ guidance_scale: float = 7.0
132
+ lora_scale: float = 0.85
133
+ seed: int = 1337
134
+ label_format: LabelFormat = "yolo"
135
+ class_id: int = 0
136
+ category_id: int = 1
137
+ detector_model: str = "IDEA-Research/grounding-dino-tiny"
138
+ detection_threshold: float = 0.35
139
+ text_threshold: float = 0.25
140
+ debug_preview_count: int = 5
141
+ negative_prompt: str = "low quality, blurry, duplicate object, malformed, noisy, text, watermark"
142
+
143
+
144
+ @dataclass
145
+ class DetectionRecord:
146
+ """Normalized internal representation of one detector output."""
147
+
148
+ image_id: int
149
+ label: str
150
+ score: float
151
+ box_xyxy: Tuple[float, float, float, float]
152
+ width: int
153
+ height: int
154
+
155
+ def clipped(self) -> "DetectionRecord":
156
+ xmin, ymin, xmax, ymax = self.box_xyxy
157
+ xmin = float(max(0.0, min(xmin, self.width - 1)))
158
+ ymin = float(max(0.0, min(ymin, self.height - 1)))
159
+ xmax = float(max(0.0, min(xmax, self.width - 1)))
160
+ ymax = float(max(0.0, min(ymax, self.height - 1)))
161
+ if xmax < xmin:
162
+ xmin, xmax = xmax, xmin
163
+ if ymax < ymin:
164
+ ymin, ymax = ymax, ymin
165
+ return DetectionRecord(
166
+ image_id=self.image_id,
167
+ label=self.label,
168
+ score=self.score,
169
+ box_xyxy=(xmin, ymin, xmax, ymax),
170
+ width=self.width,
171
+ height=self.height,
172
+ )
173
+
174
+
175
+ def configure_logging(level: str = "INFO") -> None:
176
+ numeric_level = getattr(logging, level.upper(), logging.INFO)
177
+ logging.basicConfig(
178
+ level=numeric_level,
179
+ format="%(asctime)s | %(levelname)s | %(name)s | %(message)s",
180
+ datefmt="%Y-%m-%d %H:%M:%S",
181
+ )
182
+
183
+
184
+ def require_dependency(name: str, import_error: Optional[BaseException]) -> None:
185
+ if import_error is not None:
186
+ raise RuntimeError(
187
+ f"Required dependency '{name}' could not be imported. Install the expected GPU stack "
188
+ f"before running this pipeline. Original error: {import_error}"
189
+ ) from import_error
190
+
191
+
192
+ def resolve_device() -> str:
193
+ require_dependency("torch", _TORCH_IMPORT_ERROR)
194
+ if torch.cuda.is_available():
195
+ return "cuda"
196
+ if hasattr(torch.backends, "mps") and torch.backends.mps.is_available():
197
+ return "mps"
198
+ return "cpu"
199
+
200
+
201
+ def clear_vram() -> None:
202
+ gc.collect()
203
+ if torch is not None and torch.cuda.is_available():
204
+ torch.cuda.empty_cache()
205
+ torch.cuda.ipc_collect()
206
+
207
+
208
+ def load_source_images(images: Sequence[ImageLike]) -> List[Image.Image]:
209
+ loaded: List[Image.Image] = []
210
+ for idx, item in enumerate(images):
211
+ try:
212
+ if isinstance(item, Image.Image):
213
+ image = item.convert("RGB")
214
+ elif isinstance(item, np.ndarray):
215
+ arr = item
216
+ if arr.ndim == 2:
217
+ arr = np.stack([arr] * 3, axis=-1)
218
+ if arr.shape[-1] == 4:
219
+ arr = arr[..., :3]
220
+ image = Image.fromarray(arr.astype(np.uint8)).convert("RGB")
221
+ else:
222
+ path = Path(item).expanduser().resolve()
223
+ if not path.exists():
224
+ raise FileNotFoundError(f"Source image does not exist: {path}")
225
+ image = Image.open(path).convert("RGB")
226
+ loaded.append(image)
227
+ except Exception as exc:
228
+ raise ValueError(f"Failed to load source image at index {idx}: {exc}") from exc
229
+ if not loaded:
230
+ raise ValueError("At least one source target object image is required.")
231
+ return loaded
232
+
233
+
234
+ def load_image_paths(directory: Union[str, Path]) -> List[Path]:
235
+ root = Path(directory).expanduser().resolve()
236
+ if not root.exists() or not root.is_dir():
237
+ raise FileNotFoundError(f"Source image directory not found: {root}")
238
+ allowed = {".jpg", ".jpeg", ".png", ".webp", ".bmp"}
239
+ paths = sorted(path for path in root.iterdir() if path.suffix.lower() in allowed)
240
+ if not paths:
241
+ raise FileNotFoundError(f"No supported images found in {root}")
242
+ return paths
243
+
244
+
245
+ def pil_to_tensor(image: Image.Image, resolution: int) -> "torch.Tensor":
246
+ image = ImageOps.exif_transpose(image).convert("RGB")
247
+ image = ImageOps.fit(image, (resolution, resolution), method=Image.Resampling.LANCZOS)
248
+ arr = np.asarray(image).astype(np.float32) / 127.5 - 1.0
249
+ tensor = torch.from_numpy(arr).permute(2, 0, 1)
250
+ return tensor
251
+
252
+
253
+ class FewShotImageDataset(Dataset):
254
+ """A deterministic plus stochastic few-shot image dataset for LoRA training."""
255
+
256
+ def __init__(self, images: Sequence[Image.Image], prompt: str, resolution: int, length: int = 2048) -> None:
257
+ self.images = list(images)
258
+ self.prompt = prompt
259
+ self.resolution = resolution
260
+ self.length = max(length, len(self.images))
261
+
262
+ def __len__(self) -> int:
263
+ return self.length
264
+
265
+ def _augment(self, image: Image.Image) -> Image.Image:
266
+ image = ImageOps.exif_transpose(image).convert("RGB")
267
+ if random.random() < 0.5:
268
+ image = ImageOps.mirror(image)
269
+ brightness = random.uniform(0.82, 1.18)
270
+ contrast = random.uniform(0.85, 1.15)
271
+ saturation = random.uniform(0.82, 1.2)
272
+ image = ImageEnhance.Brightness(image).enhance(brightness)
273
+ image = ImageEnhance.Contrast(image).enhance(contrast)
274
+ image = ImageEnhance.Color(image).enhance(saturation)
275
+ angle = random.uniform(-8, 8)
276
+ image = image.rotate(angle, resample=Image.Resampling.BICUBIC, expand=False, fillcolor=(127, 127, 127))
277
+ return image
278
+
279
+ def __getitem__(self, index: int) -> Dict[str, Any]:
280
+ image = self.images[index % len(self.images)]
281
+ return {"pixel_values": pil_to_tensor(self._augment(image), self.resolution), "prompt": self.prompt}
282
+
283
+
284
+ def compute_structural_loss(candidate: Image.Image, references: Sequence[Image.Image], resolution: int) -> float:
285
+ """
286
+ Compute a lightweight structural validation loss from edge maps and luminance.
287
+
288
+ The validation signal intentionally emphasizes shape retention instead of exact
289
+ background fidelity. This lets early stopping preserve target object features
290
+ while avoiding overfitting to source lighting and environment.
291
+ """
292
+ require_dependency("opencv-python", _CV2_IMPORT_ERROR)
293
+ candidate_gray = np.asarray(ImageOps.fit(candidate.convert("L"), (resolution, resolution), Image.Resampling.LANCZOS))
294
+ candidate_edges = cv2.Canny(candidate_gray, 80, 160).astype(np.float32) / 255.0
295
+ candidate_luma = candidate_gray.astype(np.float32) / 255.0
296
+ best_loss = float("inf")
297
+ for ref in references:
298
+ ref_gray = np.asarray(ImageOps.fit(ref.convert("L"), (resolution, resolution), Image.Resampling.LANCZOS))
299
+ ref_edges = cv2.Canny(ref_gray, 80, 160).astype(np.float32) / 255.0
300
+ ref_luma = ref_gray.astype(np.float32) / 255.0
301
+ edge_loss = float(np.mean((candidate_edges - ref_edges) ** 2))
302
+ luma_loss = float(np.mean((candidate_luma - ref_luma) ** 2))
303
+ best_loss = min(best_loss, 0.75 * edge_loss + 0.25 * luma_loss)
304
+ return best_loss
305
+
306
+
307
+ def create_prompt_variation(base_prompt: str, target_object: str, instance_token: str, index: int) -> str:
308
+ camera_angles = [
309
+ "front three-quarter view",
310
+ "low angle macro shot",
311
+ "high angle inspection view",
312
+ "side profile perspective",
313
+ "telephoto compressed perspective",
314
+ "wide-angle close pass",
315
+ ]
316
+ illumination = [
317
+ "soft diffuse overcast lighting",
318
+ "hard rim light with long shadows",
319
+ "cool fluorescent industrial illumination",
320
+ "warm golden hour side light",
321
+ "dramatic backlight and controlled reflections",
322
+ "mixed practical lights with subtle glare",
323
+ ]
324
+ materials = [
325
+ "mild specular reflections",
326
+ "matte surface response",
327
+ "gloss highlights on nearby surfaces",
328
+ "wet floor reflections",
329
+ "dusty atmospheric scattering",
330
+ "clean studio-grade clarity",
331
+ ]
332
+ distance = [
333
+ "object occupying 12 percent of the frame",
334
+ "object occupying 25 percent of the frame",
335
+ "object occupying 40 percent of the frame",
336
+ "object in the near foreground",
337
+ "object at medium distance",
338
+ "object partially framed by environmental structures",
339
+ ]
340
+ occlusion = [
341
+ "unoccluded target object",
342
+ "subtle foreground occlusion at one edge",
343
+ "partial shadow crossing the target object",
344
+ "thin cable-like occluder in foreground",
345
+ "minor motion blur in the environment only",
346
+ "clean silhouette with no occlusion",
347
+ ]
348
+ rng = random.Random(index * 7919 + len(base_prompt))
349
+ descriptors = [
350
+ rng.choice(camera_angles),
351
+ rng.choice(illumination),
352
+ rng.choice(materials),
353
+ rng.choice(distance),
354
+ rng.choice(occlusion),
355
+ ]
356
+ descriptor_text = ", ".join(descriptors)
357
+ return f"a photo of {instance_token} {target_object} in {base_prompt}, {descriptor_text}, realistic, high detail"
358
+
359
+
360
+ def encode_prompt_for_training(pipe: Any, prompt: Union[str, List[str]], device: str) -> "torch.Tensor":
361
+ tokens = pipe.tokenizer(
362
+ prompt,
363
+ padding="max_length",
364
+ max_length=pipe.tokenizer.model_max_length,
365
+ truncation=True,
366
+ return_tensors="pt",
367
+ )
368
+ input_ids = tokens.input_ids.to(device)
369
+ return pipe.text_encoder(input_ids)[0]
370
+
371
+
372
+ class FewShotLoRATrainer:
373
+ """Few-shot LoRA trainer with early stopping on validation structural loss."""
374
+
375
+ def __init__(self, config: TrainingConfig) -> None:
376
+ self.config = config
377
+ self.device = resolve_device()
378
+ self.dtype = self._resolve_dtype()
379
+ self.pipe: Optional[Any] = None
380
+
381
+ def _resolve_dtype(self) -> "torch.dtype":
382
+ if torch is None:
383
+ raise RuntimeError("PyTorch is required for training.")
384
+ if self.config.mixed_precision == "bf16" and torch.cuda.is_available() and torch.cuda.is_bf16_supported():
385
+ return torch.bfloat16
386
+ if self.config.mixed_precision == "fp16" and self.device == "cuda":
387
+ return torch.float32
388
+ return torch.float32
389
+
390
+ def _load_pipeline(self) -> Any:
391
+ require_dependency("diffusers", _DIFFUSERS_IMPORT_ERROR)
392
+ require_dependency("peft", _PEFT_IMPORT_ERROR)
393
+ LOGGER.info("Loading diffusion pipeline for LoRA training: %s", self.config.pretrained_model)
394
+ try:
395
+ pipe = StableDiffusionPipeline.from_pretrained(
396
+ self.config.pretrained_model,
397
+ torch_dtype=self.dtype,
398
+ safety_checker=None,
399
+ requires_safety_checker=False,
400
+ )
401
+ pipe.scheduler = DDPMScheduler.from_config(pipe.scheduler.config)
402
+ pipe.to(self.device)
403
+ pipe.vae.requires_grad_(False)
404
+ pipe.text_encoder.requires_grad_(False)
405
+ pipe.unet.requires_grad_(False)
406
+ lora_config = LoraConfig(
407
+ r=self.config.rank,
408
+ lora_alpha=self.config.rank,
409
+ init_lora_weights="gaussian",
410
+ target_modules=["to_k", "to_q", "to_v", "to_out.0"],
411
+ )
412
+ pipe.unet.add_adapter(lora_config)
413
+ if self.config.train_text_encoder_lora:
414
+ text_lora_config = LoraConfig(
415
+ r=self.config.rank,
416
+ lora_alpha=self.config.rank,
417
+ init_lora_weights="gaussian",
418
+ target_modules=["q_proj", "k_proj", "v_proj", "out_proj"],
419
+ )
420
+ pipe.text_encoder.add_adapter(text_lora_config)
421
+ if hasattr(pipe, "enable_xformers_memory_efficient_attention"):
422
+ try:
423
+ pipe.enable_xformers_memory_efficient_attention()
424
+ LOGGER.info("Enabled xFormers memory-efficient attention.")
425
+ except Exception as exc:
426
+ LOGGER.warning("Could not enable xFormers attention: %s", exc)
427
+ if hasattr(pipe, "enable_attention_slicing"):
428
+ pipe.enable_attention_slicing()
429
+ self.pipe = pipe
430
+ return pipe
431
+ except torch.cuda.OutOfMemoryError as exc:
432
+ clear_vram()
433
+ raise RuntimeError("CUDA VRAM exhausted while loading the diffusion training pipeline.") from exc
434
+ except Exception as exc:
435
+ clear_vram()
436
+ raise RuntimeError(f"Failed to load diffusion training pipeline: {exc}") from exc
437
+
438
+ def train(self, source_images: Sequence[ImageLike], target_object: str) -> Path:
439
+ images = load_source_images(source_images)
440
+ random.seed(self.config.seed)
441
+ np.random.seed(self.config.seed)
442
+ torch.manual_seed(self.config.seed)
443
+ if torch.cuda.is_available():
444
+ torch.cuda.manual_seed_all(self.config.seed)
445
+
446
+ output_dir = Path(self.config.output_dir).expanduser().resolve()
447
+ output_dir.mkdir(parents=True, exist_ok=True)
448
+ prompt = f"a photo of {self.config.instance_token} {target_object}"
449
+ dataset = FewShotImageDataset(images, prompt=prompt, resolution=self.config.resolution)
450
+ loader = DataLoader(dataset, batch_size=self.config.batch_size, shuffle=True, num_workers=0)
451
+ iterator = iter(loader)
452
+ pipe = self._load_pipeline()
453
+
454
+ trainable_params = [p for p in pipe.unet.parameters() if p.requires_grad]
455
+ if self.config.train_text_encoder_lora:
456
+ trainable_params += [p for p in pipe.text_encoder.parameters() if p.requires_grad]
457
+ if not trainable_params:
458
+ raise RuntimeError("No trainable LoRA parameters were registered. Check PEFT and Diffusers versions.")
459
+ optimizer = torch.optim.AdamW(trainable_params, lr=self.config.learning_rate, betas=(0.9, 0.999), weight_decay=0.01)
460
+ scaler_enabled = self.dtype == torch.float32 and self.device == "cuda"
461
+ scaler = torch.cuda.amp.GradScaler(enabled=False)
462
+
463
+ best_structural_loss = float("inf")
464
+ patience_counter = 0
465
+ global_step = 0
466
+ running_loss = 0.0
467
+ start = time.time()
468
+ pipe.unet.train()
469
+ if self.config.train_text_encoder_lora:
470
+ pipe.text_encoder.train()
471
+ LOGGER.info("Starting LoRA optimization with %d source images for target '%s'.", len(images), target_object)
472
+
473
+ while global_step < self.config.train_steps:
474
+ try:
475
+ batch = next(iterator)
476
+ except StopIteration:
477
+ iterator = iter(loader)
478
+ batch = next(iterator)
479
+ pixel_values = batch["pixel_values"].to(device=self.device, dtype=self.dtype)
480
+ prompts = list(batch["prompt"])
481
+
482
+ with torch.no_grad():
483
+ latents = pipe.vae.encode(pixel_values).latent_dist.sample()
484
+ latents = latents * pipe.vae.config.scaling_factor
485
+ noise = torch.randn_like(latents)
486
+ timesteps = torch.randint(
487
+ 0,
488
+ pipe.scheduler.config.num_train_timesteps,
489
+ (latents.shape[0],),
490
+ device=self.device,
491
+ dtype=torch.long,
492
+ )
493
+ noisy_latents = pipe.scheduler.add_noise(latents, noise, timesteps)
494
+ encoder_hidden_states = encode_prompt_for_training(pipe, prompts, self.device).to(dtype=self.dtype)
495
+
496
+ try:
497
+ with torch.amp.autocast('cuda', dtype=torch.float32):
498
+ model_pred = pipe.unet(noisy_latents, timesteps, encoder_hidden_states).sample
499
+ target = noise
500
+ if getattr(pipe.scheduler.config, "prediction_type", None) == "v_prediction":
501
+ target = pipe.scheduler.get_velocity(latents, noise, timesteps)
502
+ loss = F.mse_loss(model_pred.float(), target.float(), reduction="mean")
503
+ loss = loss / self.config.gradient_accumulation_steps
504
+ scaler.scale(loss).backward()
505
+ except torch.cuda.OutOfMemoryError as exc:
506
+ clear_vram()
507
+ raise RuntimeError("CUDA VRAM exhausted during LoRA optimization. Reduce resolution, rank, or batch size.") from exc
508
+
509
+ running_loss += float(loss.detach().cpu().item())
510
+ if (global_step + 1) % self.config.gradient_accumulation_steps == 0:
511
+ scaler.unscale_(optimizer)
512
+ torch.nn.utils.clip_grad_norm_(trainable_params, self.config.max_grad_norm)
513
+ scaler.step(optimizer)
514
+ scaler.update()
515
+ optimizer.zero_grad(set_to_none=True)
516
+
517
+ global_step += 1
518
+ if global_step % max(1, self.config.validation_interval) == 0 or global_step == self.config.train_steps:
519
+ structural_loss = self._validate_structural_loss(pipe, images, target_object)
520
+ avg_train_loss = running_loss / max(1, self.config.validation_interval)
521
+ running_loss = 0.0
522
+ LOGGER.info(
523
+ "step=%d/%d train_loss=%.6f validation_structural_loss=%.6f best=%.6f elapsed=%.1fs",
524
+ global_step,
525
+ self.config.train_steps,
526
+ avg_train_loss,
527
+ structural_loss,
528
+ best_structural_loss,
529
+ time.time() - start,
530
+ )
531
+ if structural_loss + self.config.min_delta < best_structural_loss:
532
+ best_structural_loss = structural_loss
533
+ patience_counter = 0
534
+ self._save_lora(pipe, output_dir)
535
+ LOGGER.info("Saved improved LoRA checkpoint to %s", output_dir)
536
+ else:
537
+ patience_counter += 1
538
+ if patience_counter >= self.config.patience:
539
+ LOGGER.info(
540
+ "Early stopping triggered at step %d after %d stagnant validations.",
541
+ global_step,
542
+ patience_counter,
543
+ )
544
+ break
545
+
546
+ self._save_lora(pipe, output_dir)
547
+ metadata = {
548
+ "target_object": target_object,
549
+ "instance_token": self.config.instance_token,
550
+ "pretrained_model": self.config.pretrained_model,
551
+ "best_structural_loss": best_structural_loss,
552
+ "steps_completed": global_step,
553
+ "training_config": asdict(self.config),
554
+ }
555
+ (output_dir / "conditioning_metadata.json").write_text(json.dumps(metadata, indent=2), encoding="utf-8")
556
+ LOGGER.info("Training complete. LoRA artifacts are available in %s", output_dir)
557
+ return output_dir
558
+
559
+ def _validate_structural_loss(self, pipe: Any, references: Sequence[Image.Image], target_object: str) -> float:
560
+ pipe.unet.eval()
561
+ if self.config.train_text_encoder_lora:
562
+ pipe.text_encoder.eval()
563
+ prompt = f"a centered studio product photo of {self.config.instance_token} {target_object}, neutral background, crisp outline"
564
+ generator = torch.Generator(device=self.device).manual_seed(self.config.seed + 17)
565
+ losses: List[float] = []
566
+ try:
567
+ with torch.no_grad():
568
+ generated = pipe(
569
+ prompt=prompt,
570
+ negative_prompt=self.config.negative_prompt,
571
+ num_images_per_prompt=self.config.num_validation_images,
572
+ num_inference_steps=18,
573
+ guidance_scale=6.0,
574
+ height=self.config.resolution,
575
+ width=self.config.resolution,
576
+ generator=generator,
577
+ ).images
578
+ for image in generated:
579
+ losses.append(compute_structural_loss(image, references, self.config.resolution))
580
+ except torch.cuda.OutOfMemoryError as exc:
581
+ clear_vram()
582
+ raise RuntimeError("CUDA VRAM exhausted during structural validation.") from exc
583
+ finally:
584
+ pipe.unet.train()
585
+ if self.config.train_text_encoder_lora:
586
+ pipe.text_encoder.train()
587
+ if not losses:
588
+ return float("inf")
589
+ return float(np.mean(losses))
590
+
591
+ def _save_lora(self, pipe: Any, output_dir: Path) -> None:
592
+ output_dir.mkdir(parents=True, exist_ok=True)
593
+ if hasattr(pipe, "save_lora_weights") and get_peft_model_state_dict is not None:
594
+ save_kwargs: Dict[str, Any] = {"save_directory": str(output_dir), "unet_lora_layers": get_peft_model_state_dict(pipe.unet)}
595
+ if self.config.train_text_encoder_lora:
596
+ save_kwargs["text_encoder_lora_layers"] = get_peft_model_state_dict(pipe.text_encoder)
597
+ pipe.save_lora_weights(**save_kwargs)
598
+ elif hasattr(pipe.unet, "save_pretrained"):
599
+ pipe.unet.save_pretrained(str(output_dir / "unet_lora"))
600
+ else:
601
+ raise RuntimeError("The loaded pipeline cannot save LoRA weights with the installed Diffusers version.")
602
+
603
+
604
+ class ParametricSynthesizer:
605
+ """Prompt-conditioned synthetic frame generator with systematic visual variation."""
606
+
607
+ def __init__(self, config: SynthesisConfig, instance_token: str = "sksobj") -> None:
608
+ self.config = config
609
+ self.instance_token = instance_token
610
+ self.device = resolve_device()
611
+ self.dtype = torch.float32 if self.device == "cuda" else torch.float32
612
+ self.pipe: Optional[Any] = None
613
+
614
+ def _load_pipeline(self) -> Any:
615
+ require_dependency("diffusers", _DIFFUSERS_IMPORT_ERROR)
616
+ LOGGER.info("Loading synthesis pipeline: %s", self.config.pretrained_model)
617
+ try:
618
+ is_sdxl = "xl" in self.config.pretrained_model.lower() or "sdxl" in self.config.pretrained_model.lower()
619
+ pipeline_cls = StableDiffusionXLPipeline if is_sdxl else StableDiffusionPipeline
620
+ load_kwargs: Dict[str, Any] = {"torch_dtype": self.dtype}
621
+ if not is_sdxl:
622
+ load_kwargs.update({"safety_checker": None, "requires_safety_checker": False})
623
+ pipe = pipeline_cls.from_pretrained(self.config.pretrained_model, **load_kwargs)
624
+ pipe.to(self.device)
625
+ if hasattr(pipe, "enable_attention_slicing"):
626
+ pipe.enable_attention_slicing()
627
+ if hasattr(pipe, "enable_vae_slicing"):
628
+ pipe.enable_vae_slicing()
629
+ if self.config.lora_dir:
630
+ lora_path = Path(self.config.lora_dir).expanduser().resolve()
631
+ if lora_path.exists():
632
+ pipe.load_lora_weights(str(lora_path))
633
+ if hasattr(pipe, "set_adapters"):
634
+ try:
635
+ pipe.set_adapters(["default_0"], adapter_weights=[self.config.lora_scale])
636
+ except Exception:
637
+ LOGGER.debug("Adapter weighting API unavailable or adapter name differs; using loaded LoRA default scale.")
638
+ LOGGER.info("Loaded LoRA weights from %s", lora_path)
639
+ else:
640
+ raise FileNotFoundError(f"Configured LoRA directory does not exist: {lora_path}")
641
+ self.pipe = pipe
642
+ return pipe
643
+ except torch.cuda.OutOfMemoryError as exc:
644
+ clear_vram()
645
+ raise RuntimeError("CUDA VRAM exhausted while loading the synthesis pipeline.") from exc
646
+ except Exception as exc:
647
+ clear_vram()
648
+ raise RuntimeError(f"Failed to load synthesis pipeline: {exc}") from exc
649
+
650
+ def generate(self) -> List[Tuple[Path, Image.Image, str]]:
651
+ pipe = self.pipe or self._load_pipeline()
652
+ output_root = Path(self.config.output_dir).expanduser().resolve()
653
+ image_dir = output_root / "images"
654
+ label_dir = output_root / "labels"
655
+ image_dir.mkdir(parents=True, exist_ok=True)
656
+ label_dir.mkdir(parents=True, exist_ok=True)
657
+ generated_records: List[Tuple[Path, Image.Image, str]] = []
658
+ LOGGER.info("Generating %d synthetic frames into %s", self.config.num_images, image_dir)
659
+ for start_idx in range(0, self.config.num_images, self.config.batch_size):
660
+ current_batch = min(self.config.batch_size, self.config.num_images - start_idx)
661
+ prompts = [
662
+ create_prompt_variation(self.config.prompt, self.config.target_object, self.instance_token, start_idx + i)
663
+ for i in range(current_batch)
664
+ ]
665
+ generators = [torch.Generator(device=self.device).manual_seed(self.config.seed + start_idx + i) for i in range(current_batch)]
666
+ try:
667
+ with torch.no_grad():
668
+ result = pipe(
669
+ prompt=prompts,
670
+ negative_prompt=[self.config.negative_prompt] * current_batch,
671
+ width=self.config.width,
672
+ height=self.config.height,
673
+ num_inference_steps=self.config.inference_steps,
674
+ guidance_scale=self.config.guidance_scale,
675
+ generator=generators,
676
+ )
677
+ except torch.cuda.OutOfMemoryError as exc:
678
+ clear_vram()
679
+ raise RuntimeError("CUDA VRAM exhausted during synthesis. Lower batch size, resolution, or steps.") from exc
680
+ for local_idx, image in enumerate(result.images):
681
+ image_id = start_idx + local_idx
682
+ post_image = self._postprocess_variation(image, image_id)
683
+ image_path = image_dir / f"synthetic_{image_id:06d}.jpg"
684
+ post_image.save(image_path, quality=95)
685
+ generated_records.append((image_path, post_image, prompts[local_idx]))
686
+ LOGGER.debug("Generated %s with prompt: %s", image_path.name, prompts[local_idx])
687
+ LOGGER.info("Generated %d images.", len(generated_records))
688
+ return generated_records
689
+
690
+ def _postprocess_variation(self, image: Image.Image, index: int) -> Image.Image:
691
+ rng = random.Random(self.config.seed + index * 1543)
692
+ image = image.convert("RGB")
693
+ if rng.random() < 0.35:
694
+ overlay = Image.new("RGB", image.size, (255, 255, 255))
695
+ alpha = rng.uniform(0.015, 0.06)
696
+ image = Image.blend(image, overlay, alpha)
697
+ if rng.random() < 0.35:
698
+ image = ImageEnhance.Brightness(image).enhance(rng.uniform(0.88, 1.12))
699
+ if rng.random() < 0.35:
700
+ image = ImageEnhance.Contrast(image).enhance(rng.uniform(0.9, 1.15))
701
+ if rng.random() < 0.25:
702
+ image = image.filter(ImageFilter.GaussianBlur(radius=rng.uniform(0.0, 0.45)))
703
+ return image
704
+
705
+
706
+ class GroundingDINOLabeler:
707
+ """Grounding DINO wrapper for zero-shot bounding-box extraction."""
708
+
709
+ def __init__(self, model_id: str, box_threshold: float, text_threshold: float) -> None:
710
+ require_dependency("transformers", _TRANSFORMERS_IMPORT_ERROR)
711
+ require_dependency("torch", _TORCH_IMPORT_ERROR)
712
+ self.model_id = model_id
713
+ self.box_threshold = box_threshold
714
+ self.text_threshold = text_threshold
715
+ self.device = resolve_device()
716
+ LOGGER.info("Loading Grounding DINO detector: %s", model_id)
717
+ try:
718
+ self.processor = AutoProcessor.from_pretrained(model_id)
719
+ self.model = AutoModelForZeroShotObjectDetection.from_pretrained(model_id).to(self.device)
720
+ self.model.eval()
721
+ except torch.cuda.OutOfMemoryError as exc:
722
+ clear_vram()
723
+ raise RuntimeError("CUDA VRAM exhausted while loading Grounding DINO.") from exc
724
+ except Exception as exc:
725
+ clear_vram()
726
+ raise RuntimeError(f"Failed to load Grounding DINO model '{model_id}': {exc}") from exc
727
+
728
+ def detect(self, image: Image.Image, query: str, image_id: int) -> List[DetectionRecord]:
729
+ text_labels = [[query]]
730
+ width, height = image.size
731
+ try:
732
+ inputs = self.processor(images=image, text=text_labels, return_tensors="pt").to(self.device)
733
+ with torch.no_grad():
734
+ outputs = self.model(**inputs)
735
+ results = self.processor.post_process_grounded_object_detection(
736
+ outputs,
737
+ inputs.input_ids,
738
+ threshold=self.box_threshold,
739
+ text_threshold=self.text_threshold,
740
+ target_sizes=[(height, width)],
741
+ )[0]
742
+ except torch.cuda.OutOfMemoryError as exc:
743
+ clear_vram()
744
+ raise RuntimeError("CUDA VRAM exhausted during Grounding DINO inference.") from exc
745
+ except Exception as exc:
746
+ raise RuntimeError(f"Grounding DINO inference failed for image_id={image_id}: {exc}") from exc
747
+
748
+ records: List[DetectionRecord] = []
749
+ boxes = results.get("boxes", [])
750
+ scores = results.get("scores", [])
751
+ labels = results.get("labels", [])
752
+ for box, score, label in zip(boxes, scores, labels):
753
+ box_tuple = tuple(float(x) for x in box.detach().cpu().tolist())
754
+ score_float = float(score.detach().cpu().item()) if hasattr(score, "detach") else float(score)
755
+ label_text = str(label)
756
+ rec = DetectionRecord(image_id=image_id, label=label_text, score=score_float, box_xyxy=box_tuple, width=width, height=height).clipped()
757
+ xmin, ymin, xmax, ymax = rec.box_xyxy
758
+ if xmax - xmin >= 2 and ymax - ymin >= 2:
759
+ records.append(rec)
760
+ LOGGER.debug("Detector returned %d boxes for image_id=%d", len(records), image_id)
761
+ return records
762
+
763
+
764
+ def convert_detection(record: DetectionRecord, fmt: LabelFormat, class_id: int = 0, category_id: int = 1) -> Union[List[float], Dict[str, Any]]:
765
+ rec = record.clipped()
766
+ xmin, ymin, xmax, ymax = rec.box_xyxy
767
+ box_w = max(0.0, xmax - xmin)
768
+ box_h = max(0.0, ymax - ymin)
769
+ if fmt == "yolo":
770
+ x_center = (xmin + box_w / 2.0) / rec.width
771
+ y_center = (ymin + box_h / 2.0) / rec.height
772
+ return [
773
+ int(class_id),
774
+ round(float(x_center), 6),
775
+ round(float(y_center), 6),
776
+ round(float(box_w / rec.width), 6),
777
+ round(float(box_h / rec.height), 6),
778
+ ]
779
+ if fmt == "coco":
780
+ return {
781
+ "image_id": int(rec.image_id),
782
+ "category_id": int(category_id),
783
+ "bbox": [round(float(xmin), 2), round(float(ymin), 2), round(float(box_w), 2), round(float(box_h), 2)],
784
+ "score": round(float(rec.score), 6),
785
+ "label": rec.label,
786
+ }
787
+ raise ValueError(f"Unsupported label format: {fmt}")
788
+
789
+
790
+ def write_label_file(label_dir: Path, image_path: Path, detections: Sequence[DetectionRecord], config: SynthesisConfig) -> Path:
791
+ label_dir.mkdir(parents=True, exist_ok=True)
792
+ if config.label_format == "yolo":
793
+ label_path = label_dir / f"{image_path.stem}.txt"
794
+ lines = []
795
+ for rec in detections:
796
+ converted = convert_detection(rec, "yolo", class_id=config.class_id)
797
+ lines.append(" ".join(str(x) for x in converted))
798
+ label_path.write_text("\n".join(lines) + ("\n" if lines else ""), encoding="utf-8")
799
+ return label_path
800
+ label_path = label_dir / f"{image_path.stem}.json"
801
+ records = [convert_detection(rec, "coco", category_id=config.category_id) for rec in detections]
802
+ label_path.write_text(json.dumps(records, indent=2), encoding="utf-8")
803
+ return label_path
804
+
805
+
806
+ def draw_debug_previews(
807
+ output_dir: Union[str, Path],
808
+ image_records: Sequence[Tuple[Path, List[DetectionRecord]]],
809
+ label_format: LabelFormat,
810
+ count: int = 5,
811
+ seed: int = 1337,
812
+ ) -> List[Path]:
813
+ require_dependency("opencv-python", _CV2_IMPORT_ERROR)
814
+ output_root = Path(output_dir).expanduser().resolve()
815
+ if not image_records:
816
+ LOGGER.warning("No image records available for debug preview generation.")
817
+ return []
818
+ rng = random.Random(seed)
819
+ sample = list(image_records)
820
+ rng.shuffle(sample)
821
+ selected = sample[: min(count, len(sample))]
822
+ preview_paths: List[Path] = []
823
+ for idx, (image_path, detections) in enumerate(selected):
824
+ img = cv2.imread(str(image_path))
825
+ if img is None:
826
+ LOGGER.warning("OpenCV could not read image for debug preview: %s", image_path)
827
+ continue
828
+ for rec in detections:
829
+ xmin, ymin, xmax, ymax = [int(round(v)) for v in rec.clipped().box_xyxy]
830
+ cv2.rectangle(img, (xmin, ymin), (xmax, ymax), (40, 220, 40), 2)
831
+ label = f"{rec.label} {rec.score:.2f}"
832
+ cv2.putText(img, label, (xmin, max(15, ymin - 6)), cv2.FONT_HERSHEY_SIMPLEX, 0.48, (40, 220, 40), 1, cv2.LINE_AA)
833
+ preview_path = output_root / f"debug_preview_{idx:02d}.jpg"
834
+ cv2.imwrite(str(preview_path), img)
835
+ preview_paths.append(preview_path)
836
+ LOGGER.info("Saved %s debug preview: %s", label_format.upper(), preview_path)
837
+ return preview_paths
838
+
839
+
840
+ def compile_and_label_outputs(generated: Sequence[Tuple[Path, Image.Image, str]], config: SynthesisConfig) -> Dict[str, Any]:
841
+ output_root = Path(config.output_dir).expanduser().resolve()
842
+ image_dir = output_root / "images"
843
+ label_dir = output_root / "labels"
844
+ image_dir.mkdir(parents=True, exist_ok=True)
845
+ label_dir.mkdir(parents=True, exist_ok=True)
846
+ labeler = GroundingDINOLabeler(config.detector_model, config.detection_threshold, config.text_threshold)
847
+ image_records: List[Tuple[Path, List[DetectionRecord]]] = []
848
+ total_boxes = 0
849
+ for image_id, (image_path, image, prompt) in enumerate(generated):
850
+ detections = labeler.detect(image, query=config.target_object, image_id=image_id)
851
+ write_label_file(label_dir, image_path, detections, config)
852
+ image_records.append((image_path, detections))
853
+ total_boxes += len(detections)
854
+ LOGGER.info("Labeled image_id=%d file=%s boxes=%d", image_id, image_path.name, len(detections))
855
+ previews = draw_debug_previews(output_root, image_records, config.label_format, config.debug_preview_count, config.seed)
856
+ manifest = {
857
+ "output_dir": str(output_root),
858
+ "image_dir": str(image_dir),
859
+ "label_dir": str(label_dir),
860
+ "label_format": config.label_format,
861
+ "target_object": config.target_object,
862
+ "num_images": len(generated),
863
+ "total_boxes": total_boxes,
864
+ "debug_previews": [str(path) for path in previews],
865
+ "config": asdict(config),
866
+ }
867
+ (output_root / "manifest.json").write_text(json.dumps(manifest, indent=2), encoding="utf-8")
868
+ LOGGER.info("Compiled labeled dataset in %s with %d boxes.", output_root, total_boxes)
869
+ return manifest
870
+
871
+
872
+ def run_pipeline(source_images: Sequence[ImageLike], train_cfg: TrainingConfig, synth_cfg: SynthesisConfig) -> Dict[str, Any]:
873
+ LOGGER.info("Starting end-to-end synthetic data pipeline.")
874
+ trainer = FewShotLoRATrainer(train_cfg)
875
+ lora_dir = trainer.train(source_images, synth_cfg.target_object)
876
+ synth_cfg.lora_dir = str(lora_dir)
877
+ synthesizer = ParametricSynthesizer(synth_cfg, instance_token=train_cfg.instance_token)
878
+ generated = synthesizer.generate()
879
+ manifest = compile_and_label_outputs(generated, synth_cfg)
880
+ LOGGER.info("Pipeline complete: %s", manifest["output_dir"])
881
+ return manifest
882
+
883
+
884
+ if modal is not None: # pragma: no cover - only active in Modal runtime
885
+ modal_image = (
886
+ modal.Image.debian_slim(python_version="3.11")
887
+ .pip_install(
888
+ "torch",
889
+ "diffusers",
890
+ "transformers",
891
+ "accelerate",
892
+ "opencv-python-headless",
893
+ "pillow",
894
+ "peft",
895
+ "safetensors",
896
+ )
897
+ )
898
+ app = modal.App("standalone-cv-synthetic-data-engine", image=modal_image)
899
+
900
+ @app.function(gpu="A10G", timeout=60 * 60 * 3)
901
+ def modal_run_pipeline(source_dir: str, training_config: Dict[str, Any], synthesis_config: Dict[str, Any]) -> Dict[str, Any]:
902
+ configure_logging("INFO")
903
+ paths = load_image_paths(source_dir)
904
+ train_cfg = TrainingConfig(**training_config)
905
+ synth_cfg = SynthesisConfig(**synthesis_config)
906
+ return run_pipeline(paths, train_cfg, synth_cfg)
907
+ else:
908
+ app = None
909
+
910
+
911
+ def parse_args(argv: Optional[Sequence[str]] = None) -> argparse.Namespace:
912
+ parser = argparse.ArgumentParser(description="Few-shot synthetic CV data generation pipeline.")
913
+ parser.add_argument("--source-dir", required=True, help="Directory containing target object source images.")
914
+ parser.add_argument("--target-object", required=True, help="Exact semantic text string for the target object detector query.")
915
+ parser.add_argument("--prompt", required=True, help="Background/environment prompt, e.g. 'industrial conveyor belt with reflections'.")
916
+ parser.add_argument("--output-dir", default="output_batch", help="Output dataset directory.")
917
+ parser.add_argument("--pretrained-model", default="runwayml/stable-diffusion-v1-5", help="Open diffusion model identifier or local path.")
918
+ parser.add_argument("--detector-model", default="IDEA-Research/grounding-dino-tiny", help="Grounding DINO model identifier.")
919
+ parser.add_argument("--label-format", choices=["yolo", "coco"], default="yolo", help="Annotation output format.")
920
+ parser.add_argument("--num-images", type=int, default=24, help="Number of synthetic images to generate.")
921
+ parser.add_argument("--train-steps", type=int, default=180, help="Maximum LoRA training steps before early stopping.")
922
+ parser.add_argument("--resolution", type=int, default=512, help="Training image resolution.")
923
+ parser.add_argument("--width", type=int, default=512, help="Generated image width.")
924
+ parser.add_argument("--height", type=int, default=512, help="Generated image height.")
925
+ parser.add_argument("--batch-size", type=int, default=1, help="Synthesis batch size.")
926
+ parser.add_argument("--train-batch-size", type=int, default=1, help="LoRA training batch size.")
927
+ parser.add_argument("--rank", type=int, default=8, help="LoRA rank.")
928
+ parser.add_argument("--learning-rate", type=float, default=1e-4, help="LoRA learning rate.")
929
+ parser.add_argument("--seed", type=int, default=1337, help="Random seed.")
930
+ parser.add_argument("--log-level", default="INFO", help="Logging verbosity.")
931
+ return parser.parse_args(argv)
932
+
933
+
934
+ def main(argv: Optional[Sequence[str]] = None) -> int:
935
+ args = parse_args(argv)
936
+ configure_logging(args.log_level)
937
+ try:
938
+ source_paths = load_image_paths(args.source_dir)
939
+ lora_dir = str(Path(args.output_dir).expanduser().resolve() / "conditioned_lora")
940
+ train_cfg = TrainingConfig(
941
+ pretrained_model=args.pretrained_model,
942
+ output_dir=lora_dir,
943
+ resolution=args.resolution,
944
+ train_steps=args.train_steps,
945
+ batch_size=args.train_batch_size,
946
+ rank=args.rank,
947
+ learning_rate=args.learning_rate,
948
+ seed=args.seed,
949
+ )
950
+ synth_cfg = SynthesisConfig(
951
+ prompt=args.prompt,
952
+ target_object=args.target_object,
953
+ output_dir=args.output_dir,
954
+ lora_dir=lora_dir,
955
+ pretrained_model=args.pretrained_model,
956
+ num_images=args.num_images,
957
+ batch_size=args.batch_size,
958
+ width=args.width,
959
+ height=args.height,
960
+ seed=args.seed,
961
+ label_format=args.label_format,
962
+ detector_model=args.detector_model,
963
+ )
964
+ manifest = run_pipeline(source_paths, train_cfg, synth_cfg)
965
+ LOGGER.info("Final manifest: %s", json.dumps(manifest, indent=2))
966
+ return 0
967
+ except KeyboardInterrupt:
968
+ LOGGER.warning("Pipeline interrupted by user.")
969
+ return 130
970
+ except Exception as exc:
971
+ LOGGER.exception("Pipeline failed: %s", exc)
972
+ return 1
973
+ finally:
974
+ clear_vram()
975
+
976
+
977
+ if __name__ == "__main__":
978
+ raise SystemExit(main())