File size: 17,732 Bytes
fd0b9df
 
66c36b4
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
fd0b9df
 
 
 
 
53940d1
 
 
 
fd0b9df
53940d1
 
fd0b9df
 
 
0a993a9
fd0b9df
 
 
1bcd3f8
291f11c
66c36b4
291f11c
 
 
 
 
 
 
0a993a9
 
 
68b5356
 
 
 
 
 
a108fbc
 
1f6e488
68b5356
 
 
fb5e5c4
68b5356
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
a108fbc
68b5356
 
 
 
 
a108fbc
68b5356
 
 
 
 
 
 
 
 
 
 
 
 
a108fbc
68b5356
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
0a993a9
68b5356
fd0b9df
 
53940d1
fd0b9df
 
e5752bc
53940d1
e5752bc
 
 
 
d5b9763
e5752bc
 
 
53940d1
 
e5752bc
53940d1
 
 
 
 
 
 
 
 
 
 
e5752bc
 
 
 
 
 
7e55dd0
 
53940d1
fd0b9df
a2af9f2
 
 
 
 
 
 
 
fd0b9df
53940d1
 
 
 
01c4239
 
47ef589
 
01c4239
 
47ef589
 
 
 
01c4239
 
 
18bda20
 
 
 
 
 
 
 
 
 
 
47ef589
4b81168
01c4239
 
18bda20
 
53940d1
35bc0e8
 
 
 
 
 
 
18bda20
35bc0e8
 
9d5c943
 
 
 
fd0b9df
 
66c36b4
 
 
 
 
 
 
 
 
 
 
 
 
 
 
fd0b9df
 
 
36cd8ec
 
ec4a4b3
 
 
 
66c36b4
 
fd0b9df
66c36b4
 
36cd8ec
 
 
fd0b9df
 
66c36b4
36cd8ec
 
 
 
 
 
 
 
 
 
 
35bc0e8
36cd8ec
 
 
 
 
 
 
 
 
 
 
 
35bc0e8
36cd8ec
 
fd0b9df
 
66c36b4
 
 
 
 
 
 
 
 
 
fd0b9df
 
53940d1
fd0b9df
 
ec4a4b3
53940d1
0a993a9
4b81168
36cd8ec
 
 
 
ec4a4b3
 
 
 
 
 
36cd8ec
 
 
 
 
 
 
9d5c943
 
36cd8ec
 
ec4a4b3
 
 
 
53940d1
 
940cec7
 
 
 
fd0b9df
 
53940d1
66c36b4
 
 
 
53940d1
 
885e6f8
66c36b4
 
 
 
 
 
 
ec4a4b3
66c36b4
ec4a4b3
 
 
1
2
3
4
5
6
7
8
9
10
11
12
13
14
15
16
17
18
19
20
21
22
23
24
25
26
27
28
29
30
31
32
33
34
35
36
37
38
39
40
41
42
43
44
45
46
47
48
49
50
51
52
53
54
55
56
57
58
59
60
61
62
63
64
65
66
67
68
69
70
71
72
73
74
75
76
77
78
79
80
81
82
83
84
85
86
87
88
89
90
91
92
93
94
95
96
97
98
99
100
101
102
103
104
105
106
107
108
109
110
111
112
113
114
115
116
117
118
119
120
121
122
123
124
125
126
127
128
129
130
131
132
133
134
135
136
137
138
139
140
141
142
143
144
145
146
147
148
149
150
151
152
153
154
155
156
157
158
159
160
161
162
163
164
165
166
167
168
169
170
171
172
173
174
175
176
177
178
179
180
181
182
183
184
185
186
187
188
189
190
191
192
193
194
195
196
197
198
199
200
201
202
203
204
205
206
207
208
209
210
211
212
213
214
215
216
217
218
219
220
221
222
223
224
225
226
227
228
229
230
231
232
233
234
235
236
237
238
239
240
241
242
243
244
245
246
247
248
249
250
251
252
253
254
255
256
257
258
259
260
261
262
263
264
265
266
267
268
269
270
271
272
273
274
275
276
277
278
279
280
281
282
283
284
285
286
287
288
289
290
291
292
293
294
295
296
297
298
299
300
301
302
303
304
305
306
307
308
309
310
311
312
313
314
315
316
317
318
319
320
321
322
323
324
325
326
327
328
329
330
331
332
333
334
335
336
337
338
339
340
341
342
343
344
345
346
347
348
349
350
351
352
353
354
355
356
357
358
359
360
361
362
363
364
365
366
367
368
369
370
371
372
373
374
375
376
377
378
379
380
381
382
383
384
385
386
387
388
389
390
391
392
393
394
395
396
import random
import os
import sys
import threading
import asyncio

# ---------------------------------------------------------------------------
# ZeroGPU shim โ€” spaces MUST be imported BEFORE torch.
# On HF ZeroGPU, spaces patches CUDA initialisation; that patch must land
# before torch is imported or GPU calls can silently mis-behave.
# Locally spaces isn't installed, so we fall back to a no-op @GPU decorator
# that makes the same code run unchanged on MPS / CPU.
# ---------------------------------------------------------------------------

try:
    import spaces  # type: ignore
    GPU = spaces.GPU
    ON_ZEROGPU = True
except Exception:
    def GPU(*dargs, **dkwargs):  # noqa: N802 โ€” mirror the spaces.GPU API
        def wrap(fn):
            return fn
        if len(dargs) == 1 and callable(dargs[0]) and not dkwargs:
            return dargs[0]
        return wrap
    ON_ZEROGPU = False

import torch

# Patch for torch < 2.4 which lacks torch.xpu (required by diffusers >= 0.30)
if not hasattr(torch, "xpu"):
    class _MockXPU:
        is_available    = staticmethod(lambda: False)
        device_count    = staticmethod(lambda: 0)
        empty_cache     = staticmethod(lambda: None)
        manual_seed     = staticmethod(lambda seed: None)
        reset_peak_memory_stats = staticmethod(lambda: None)
        max_memory_allocated    = staticmethod(lambda: 0)
        synchronize     = staticmethod(lambda: None)
    torch.xpu = _MockXPU()

# ---------------------------------------------------------------------------
# Detect HuggingFace Spaces
# ---------------------------------------------------------------------------

IS_HF_SPACE = os.environ.get("SPACE_ID") is not None

if IS_HF_SPACE:
    # Suppress Python 3.13 asyncio GC bug (Invalid file descriptor: -1)
    _orig_unraisable = sys.unraisablehook
    def _unraisable_hook(args):
        if args.exc_type is ValueError and "Invalid file descriptor" in str(args.exc_value):
            return
        _orig_unraisable(args)
    sys.unraisablehook = _unraisable_hook

# ---------------------------------------------------------------------------
# Persistent storage cache โ€” survives sleep/restart on HF Spaces
# ---------------------------------------------------------------------------
# The /data mount repeatedly gets poisoned: partial downloads leave files where
# huggingface_hub/xet expect directories ([Errno 20] Not a directory), and the
# cache root itself can end up as a file with I/O errors. This block is
# bulletproof: any failure setting up /data falls back to the ephemeral
# container cache (slower cold start, but always works) and NEVER crashes the
# import โ€” a dead import takes the whole app down.
_MODEL_ID  = "black-forest-labs/FLUX.2-klein-4B"
_MODEL_DIR = "models--black-forest-labs--FLUX.2-klein-4B"


def _force_remove(path):
    """Remove path whether it's a file, dir, or broken โ€” never raises."""
    import shutil
    try:
        if os.path.isdir(path) and not os.path.islink(path):
            shutil.rmtree(path)
        elif os.path.exists(path) or os.path.islink(path):
            os.remove(path)
    except Exception as e:
        print(f"โš ๏ธ  Could not remove {path}: {e}")


def _setup_persistent_cache():
    """Point HF cache at /data if usable. Returns True if persistence is active."""
    cache_dir = "/data/hf_cache_v4"   # bump path to escape any poisoned older cache
    hub_dir   = os.path.join(cache_dir, "hub")

    # Probe: can we create a clean directory tree at cache_dir? If the path is a
    # poisoned file or the mount errors, bail out to ephemeral caching.
    try:
        if os.path.exists(cache_dir) or os.path.islink(cache_dir):
            # Validate existing cache: a real model snapshot with model_index.json
            snaps = os.path.join(hub_dir, _MODEL_DIR, "snapshots")
            valid = os.path.isdir(snaps) and any(
                os.path.isfile(os.path.join(snaps, s, "model_index.json"))
                for s in os.listdir(snaps)
            ) if os.path.isdir(snaps) else False
            if not valid:
                print(f"Cache at {cache_dir} is incomplete/poisoned โ€” wiping")
                _force_remove(cache_dir)
        os.makedirs(hub_dir, exist_ok=True)
    except Exception as e:
        print(f"โš ๏ธ  /data cache unusable ({e}) โ€” using ephemeral container cache")
        return False

    os.environ["HF_HOME"]      = cache_dir
    os.environ["HF_HUB_CACHE"] = hub_dir
    print(f"Persistent cache active โ†’ {cache_dir}")
    return True


def _is_model_cached():
    snaps = os.path.join(os.environ.get("HF_HUB_CACHE", ""), _MODEL_DIR, "snapshots")
    if not os.path.isdir(snaps):
        return False
    return any(
        os.path.isfile(os.path.join(snaps, s, "model_index.json"))
        for s in os.listdir(snaps)
    )


def _predownload_model():
    """Blocking pre-download so the first @spaces.GPU call is fast. Never raises."""
    if _is_model_cached():
        print(f"โœ… FLUX.2-klein-4B already cached โ€” skipping download")
        return
    print(f"Downloading FLUX.2-klein-4Bโ€ฆ")
    try:
        from huggingface_hub import snapshot_download
        hf_token = os.environ.get("HF_TOKEN") or os.environ.get("HUGGING_FACE_HUB_TOKEN")
        result = {}
        def _download():
            loop = asyncio.new_event_loop()
            asyncio.set_event_loop(loop)
            try:
                snapshot_download(
                    _MODEL_ID,
                    token=hf_token,
                    ignore_patterns=["*.msgpack", "*.h5", "flax_model*"],
                )
                result["ok"] = True
            except Exception as e:
                result["error"] = e
            finally:
                try: loop.close()
                except: pass
        t = threading.Thread(target=_download, daemon=True)
        t.start(); t.join()
        if "error" in result:
            raise result["error"]
        print("โœ… FLUX.2-klein-4B cached successfully")
    except Exception as e:
        print(f"โš ๏ธ  Pre-cache warning (will retry at generation time): {e}")


if IS_HF_SPACE:
    if os.path.isdir("/data"):
        _setup_persistent_cache()   # falls back to ephemeral on any failure
    _predownload_model()
else:
    print("Local mode.")

# ---------------------------------------------------------------------------
# Emoji โ†’ descriptive text maps (fed into FLUX prompt)
# ---------------------------------------------------------------------------

# Keep keys + order in sync with ANIMAL_EMOJIS in app.py (grouped by lookalike).
ANIMAL_MAP: dict[str, str] = {
    "๐Ÿถ": "puppy",         "๐Ÿฑ": "kitten",        "๐Ÿฐ": "bunny",
    "๐Ÿญ": "baby mouse",    "๐ŸฆŠ": "baby fox",      "๐Ÿบ": "baby wolf",
    "๐Ÿป": "baby bear",     "๐Ÿผ": "baby panda",    "๐Ÿจ": "baby koala",
    "๐Ÿฆ": "baby lion",     "๐Ÿฏ": "baby tiger",    "๐Ÿท": "baby pig",
    "๐Ÿฎ": "baby cow",      "๐Ÿด": "pony",          "๐Ÿ‘": "baby lamb",
    "๐Ÿธ": "baby frog",     "๐Ÿฆฆ": "baby otter",    "๐Ÿง": "baby penguin",
    "๐Ÿ™": "baby octopus",  "๐Ÿฆ‹": "butterfly",     "๐Ÿ": "bee",
    "๐Ÿž": "ladybug",       "๐Ÿฆ„": "unicorn",       "๐Ÿ‰": "baby dragon",
}

# Keep keys + order in sync with PLACE_EMOJIS in app.py.
PLACE_MAP: dict[str, str] = {
    "๐ŸŒŠ": "on a sunny beach with ocean waves",
    "๐Ÿ”๏ธ": "on a snowy mountain top",
    "๐ŸŒธ": "in a cherry blossom garden",
    "๐ŸŒˆ": "under a rainbow",
    "๐ŸŒ™": "on a glowing crescent moon",
    "โญ": "surrounded by sparkling stars",
    "๐ŸŒด": "on a tropical island",
    "๐Ÿก": "in a cosy cottage garden",
    "๐ŸŒบ": "in a field of tropical flowers",
    "๐Ÿ„": "in an enchanted mushroom forest",
    "๐Ÿœ๏ธ": "in a red rock canyon",
    "๐Ÿฐ": "in a fairytale castle",
    "๐ŸŽข": "at a fun theme park",
    "โ›บ": "in a cosy camping tent",
    "๐Ÿชฉ": "under a sparkling disco ball",
    "๐ŸŽช": "at a colourful circus",
    "โ›ต": "in a little boat on calm water",
    "๐Ÿš€": "aboard a space rocket among the stars",
}

# Shared style suffix โ€” must stay in sync with STYLE in scripts/generate_dataset.py.
# This exact string appears in every training caption so the LoRA learns it as a trigger.
NUMZOO_STYLE = (
    "kawaii children's book illustration, pastel anime art style, "
    "soft painterly lighting, detailed rich background with warm fairy lights, "
    "cozy magical atmosphere, cute chibi character with big sparkling eyes, "
    "soft pastel color palette, highly detailed scene, no text"
)

# ---------------------------------------------------------------------------
# Prompt builder
# ---------------------------------------------------------------------------


def _join_animals(parts: list[str]) -> str:
    """Join 1โ€“5 animals; each one after the first is prefixed 'a cute' so the model
    renders them as distinct subjects, e.g. 'puppy, a cute kitten and a cute bunny'."""
    if len(parts) == 1:
        return parts[0]
    head, rest = parts[0], [f"a cute {p}" for p in parts[1:]]
    if len(rest) == 1:
        return f"{head} and {rest[0]}"
    return f"{head}, " + ", ".join(rest[:-1]) + f" and {rest[-1]}"


def _join_places(parts: list[str]) -> str:
    """Join 1โ€“3 phrases: 'a', 'a and b', 'a, b and c'."""
    if len(parts) == 1:
        return parts[0]
    if len(parts) == 2:
        return f"{parts[0]} and {parts[1]}"
    return f"{parts[0]}, {parts[1]} and {parts[2]}"


def build_subject(animals: list[str], places: list[str]) -> str:
    """The 'A cute {animals} {places}' core โ€” SHARED with scripts/generate_dataset.py
    so training captions and live prompts use the exact same structure & vocabulary."""
    animal_list = animals[:5] if animals else [random.choice(list(ANIMAL_MAP))]
    place_list  = places[:3]  if places  else [random.choice(list(PLACE_MAP))]
    animal_text = _join_animals([ANIMAL_MAP.get(a, "bunny") for a in animal_list])
    place_text  = _join_places([PLACE_MAP.get(p, "in a magical garden") for p in place_list])
    return f"A cute {animal_text} {place_text}"


# Trigger word the NumZoo LoRA learned to associate with the whole cozy aesthetic.
# Training captions were "NUMZOO. A cute {animals} {places}, <detail>", so the live
# prompt mirrors that prefix exactly โ€” the trigger now carries the style, making the
# verbose NUMZOO_STYLE suffix redundant (kept as a constant for the dataset generator).
NUMZOO_TRIGGER = "NUMZOO"


def build_prompt(animals: list[str], places: list[str]) -> str:
    # "no text" suppresses the model rendering the NUMZOO trigger word as a literal
    # sign/title (the distilled 4-step klein is prone to this without it).
    # With multiple animals, a layout hint reduces subject-merging (the model
    # otherwise tends to fuse adjacent animals into one).
    layout = ", side by side, each a separate full-body character" if len(animals) > 1 else ""
    return f"{NUMZOO_TRIGGER}. {build_subject(animals, places)}{layout}, no text"

# ---------------------------------------------------------------------------
# Pipeline loader โ€” two-phase, following the reference ZeroGPU pattern:
#
#   Phase 1 ยท _load_pipeline_cpu()  โ€” from_pretrained to CPU RAM.
#     Called at module scope (outside any @GPU function) so the weights are
#     already resident when the first ZeroGPU call arrives.  This keeps the
#     model-download cost out of the 60 s GPU-runtime budget and prevents
#     cold-start timeouts on the very first generation.
#
#   Phase 2 ยท get_pipeline()  โ€” .to(device).
#     Must be called INSIDE a @GPU-decorated function (where a GPU slice is
#     guaranteed). Moving already-loaded CPU tensors to CUDA is fast (~1 s)
#     and comfortably within the budget.
#
# Locally (MPS / CPU) both phases happen inside generate_reward_image because
# the no-op @GPU decorator doesn't impose any budget constraint.
# ---------------------------------------------------------------------------

_pipe = None
_pipe_lock = threading.Lock()   # serialize loading โ€” pregenerate + on-demand can race
_gen_lock  = threading.Lock()   # serialize inference โ€” one device can't run two forwards at once
# Single-slot cache so the on-demand fallback reuses the in-flight pre-generation
# for the SAME reward instead of generating it twice. Keyed by reward_id (level);
# a different reward_id always regenerates, so rewards stay varied across levels.
_recent_reward = {"id": None, "image": None}
_use_mps = (not IS_HF_SPACE) and (not torch.cuda.is_available()) and torch.backends.mps.is_available()
_dtype   = torch.float16 if _use_mps else torch.bfloat16   # float16 on MPS (bfloat16 unsupported)


def _load_pipeline_cpu() -> None:
    """Phase 1: load model weights into CPU RAM. Thread-safe โ€” pregenerate and
    on-demand handlers can call this concurrently; the lock prevents a double load
    (which previously applied the LoRA twice and crashed with a meta-tensor error)."""
    global _pipe
    if _pipe is not None:
        return
    with _pipe_lock:
        if _pipe is not None:   # re-check inside the lock
            return
        from diffusers import Flux2KleinPipeline
        hf_token = os.environ.get("HF_TOKEN") or os.environ.get("HUGGING_FACE_HUB_TOKEN")
        print(f"Loading FLUX.2-klein-4B on CPUโ€ฆ (dtype={_dtype}, token={'set' if hf_token else 'NOT SET'})")
        pipe = Flux2KleinPipeline.from_pretrained(
            _MODEL_ID,
            torch_dtype=_dtype,
            token=hf_token,
        )

        # Apply the NumZoo style LoRA (trained on FLUX.2-klein-base-4B, loads on distilled).
        # Bundled in the repo via Git LFS. Loaded here at module/CPU scope so it stays out
        # of the ZeroGPU 60 s budget. Never crash the import โ€” fall back to the base model.
        lora_path = os.path.join(os.path.dirname(os.path.abspath(__file__)), "lora", "numzoo_klein.safetensors")
        if os.path.isfile(lora_path):
            try:
                pipe.load_lora_weights(lora_path)
                print(f"โœ… NumZoo LoRA loaded from {lora_path}")
            except Exception as e:
                print(f"โš ๏ธ  Could not load NumZoo LoRA ({e}) โ€” using base model")
        else:
            print(f"โš ๏ธ  NumZoo LoRA not found at {lora_path} โ€” using base model")

        _pipe = pipe   # publish only once fully built (incl. LoRA)
        print("Pipeline loaded on CPU โ€” ready for device placement.")


def get_pipeline():
    """Phase 2: move pipeline to the target device. Call inside @GPU on HF Spaces."""
    global _pipe
    if _pipe is None:
        _load_pipeline_cpu()   # fallback for local / first-call safety
    if torch.cuda.is_available():
        return _pipe.to("cuda")
    if _use_mps:
        return _pipe.to("mps")
    return _pipe.to("cpu")

# ---------------------------------------------------------------------------
# Core generation (always wrapped in try/except)
# ---------------------------------------------------------------------------

def _generate(animals: list[str], places: list[str], reward_id=None):
    try:
        import time
        prompt = build_prompt(animals, places)
        # Serialize device placement + inference: pregenerate and on-demand handlers
        # can fire concurrently, but one device (MPS/GPU slice) can't run two forward
        # passes at once. ZeroGPU serializes @GPU calls for us; locally we must.
        with _gen_lock:
            # Reuse the pre-generated image for this reward instead of generating twice.
            if reward_id is not None and _recent_reward["id"] == reward_id \
                    and _recent_reward["image"] is not None:
                print(f"โ™ป๏ธ  Reusing pre-generated image for reward {reward_id}")
                return _recent_reward["image"], prompt

            pipe = get_pipeline()
            print(f"Generating | prompt: {prompt}")
            t0 = time.time()
            result = pipe(
                prompt=prompt,
                num_inference_steps=4,
                guidance_scale=1.0,   # klein-4B distilled: always 1.0 (not schnell's 0.0)
                height=768,           # 768 (vs 512) gives multiple subjects room โ†’ less merging
                width=768,
            )
            print(f"โœ… Generated in {time.time() - t0:.1f}s")
            image = result.images[0]
            if reward_id is not None:
                _recent_reward["id"], _recent_reward["image"] = reward_id, image
        return image, prompt

    except Exception as e:
        import traceback
        print(f"[image_generator] โŒ generation failed: {e}")
        print(traceback.format_exc())
        return None, str(e)


# ---------------------------------------------------------------------------
# Module-scope CPU pre-load (HF Spaces only).
# Runs after the persistent cache is set up and the snapshot is downloaded,
# so from_pretrained finds the weights locally and completes quickly.
# Locally this is skipped โ€” the pipeline loads lazily on first generate call.
# ---------------------------------------------------------------------------

if IS_HF_SPACE:
    _load_pipeline_cpu()

# ---------------------------------------------------------------------------
# Public API โ€” single function, @GPU decorator is a no-op locally.
# ---------------------------------------------------------------------------

@GPU(duration=60)
def generate_reward_image(animals: list[str], places: list[str], reward_id=None):
    """Generate a reward image.  On HF Spaces runs inside a ZeroGPU slice;
    locally the @GPU decorator is a no-op and MPS/CPU is used instead.
    Pass reward_id so the pre-generate + on-demand paths dedupe the same reward."""
    return _generate(animals, places, reward_id=reward_id)