File size: 25,111 Bytes
193e936
bfd60b5
 
95e2d44
bfd60b5
95e2d44
a7a37d5
c6bdae7
95e2d44
aa10ef4
 
a7a37d5
 
bfd60b5
a7a37d5
6cc7643
 
bfd60b5
 
2a37a1f
a7a37d5
e157e7f
95e2d44
a7a37d5
95e2d44
f8bfe50
 
d19eda1
a7a37d5
 
 
 
 
524d875
95e2d44
f8bfe50
 
 
524d875
 
95e2d44
 
 
 
 
aa10ef4
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
f8bfe50
95e2d44
 
 
 
 
 
 
 
 
 
 
 
 
 
4d44787
 
 
 
 
 
 
 
aa76f48
 
4d44787
 
 
aa76f48
 
 
 
 
 
 
 
 
 
 
 
 
95e2d44
 
193e936
 
 
e956cc6
4d44787
055eed4
bfd60b5
055eed4
bfd60b5
 
055eed4
bfd60b5
 
055eed4
bfd60b5
4d44787
 
bfd60b5
f8bfe50
 
95e2d44
 
 
 
 
bfd60b5
193e936
f8bfe50
 
193e936
f8bfe50
193e936
bfd60b5
f8bfe50
 
 
 
193e936
 
 
 
 
 
 
 
 
 
 
 
 
 
f8bfe50
 
193e936
 
 
 
 
 
 
 
 
 
 
 
 
 
f8bfe50
 
193e936
 
 
 
 
 
 
 
95e2d44
bfd60b5
193e936
95e2d44
 
 
 
 
 
 
 
 
 
f8bfe50
 
 
 
 
 
 
 
95e2d44
f8bfe50
 
 
 
193e936
 
 
 
 
 
 
 
 
 
bfd60b5
f8bfe50
 
7762653
f8bfe50
 
 
 
 
7762653
193e936
 
 
95e2d44
 
bfd60b5
95e2d44
 
f8bfe50
 
 
 
 
 
 
193e936
95e2d44
193e936
c2e58d6
bfd60b5
193e936
bfd60b5
95e2d44
bfd60b5
95e2d44
 
 
 
 
 
 
 
 
bfd60b5
95e2d44
 
bfd60b5
193e936
 
95e2d44
 
bfd60b5
95e2d44
 
 
 
 
 
 
 
 
 
bfd60b5
95e2d44
 
 
bfd60b5
95e2d44
 
bfd60b5
f8bfe50
95e2d44
a7a37d5
 
d19eda1
 
95e2d44
f8bfe50
 
95e2d44
 
f8bfe50
 
 
 
 
193e936
bfd60b5
95e2d44
193e936
 
 
 
bfd60b5
95e2d44
 
f8bfe50
 
95e2d44
 
bfd60b5
a7a37d5
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
aa10ef4
a7a37d5
 
 
aa10ef4
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
c6bdae7
 
 
7ab4825
c6bdae7
aa10ef4
 
 
 
d3cfc5a
c6bdae7
d3cfc5a
 
 
 
aa10ef4
 
 
 
 
 
 
 
a7a37d5
 
d3cfc5a
 
 
 
 
 
 
 
c6bdae7
 
 
7ab4825
c6bdae7
 
 
 
 
 
d3cfc5a
c6bdae7
d3cfc5a
 
 
 
 
193e936
e956cc6
f8bfe50
 
 
 
 
 
 
193e936
e956cc6
467f36e
 
 
 
f8bfe50
 
 
 
 
 
 
 
467f36e
 
 
 
60a9385
 
 
 
 
f8bfe50
 
bfd60b5
f8bfe50
e157e7f
f8bfe50
 
 
e157e7f
f8bfe50
 
 
 
 
 
 
 
 
 
d19eda1
4d44787
 
193e936
f8bfe50
 
bfd60b5
a7a37d5
bfd60b5
95e2d44
aa10ef4
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
f8bfe50
95e2d44
 
f8bfe50
 
 
aa10ef4
 
 
 
7864cbe
5f8ca6f
193e936
7762653
513aa45
2bba772
513aa45
2bba772
513aa45
f8bfe50
1291bc6
5f8ca6f
bfd60b5
95e2d44
4d44787
f8bfe50
 
 
 
 
 
95e2d44
 
f8bfe50
 
 
 
 
bfd60b5
95e2d44
6cc7643
f8bfe50
 
 
6cc7643
bfd60b5
 
2bba772
bfd60b5
 
95e2d44
f8bfe50
 
 
 
 
 
 
 
 
 
 
 
 
 
bfd60b5
f8bfe50
bfd60b5
4d44787
 
bfd60b5
 
4d44787
bfd60b5
6cc7643
 
 
f8bfe50
6cc7643
 
 
bfd60b5
a7a37d5
 
 
fcce03b
 
 
 
a7a37d5
d3cfc5a
a7a37d5
d3cfc5a
a7a37d5
 
 
7ab4825
a7a37d5
 
 
d3cfc5a
 
 
 
 
a7a37d5
 
 
 
 
 
513aa45
a7a37d5
 
c6bdae7
a7a37d5
 
c6bdae7
 
 
 
d3cfc5a
 
c6bdae7
 
 
 
a7a37d5
 
 
 
 
 
055eed4
a7a37d5
 
 
 
 
 
aa10ef4
 
a7a37d5
 
 
 
 
 
 
 
055eed4
a7a37d5
 
 
 
055eed4
a7a37d5
 
 
 
 
 
d3cfc5a
 
 
 
 
 
 
 
a7a37d5
95e2d44
f8bfe50
 
193e936
 
 
 
95e2d44
 
f8bfe50
 
 
 
 
4d44787
95e2d44
 
5fbbb9c
a7a37d5
95e2d44
 
 
055eed4
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
397
398
399
400
401
402
403
404
405
406
407
408
409
410
411
412
413
414
415
416
417
418
419
420
421
422
423
424
425
426
427
428
429
430
431
432
433
434
435
436
437
438
439
440
441
442
443
444
445
446
447
448
449
450
451
452
453
454
455
456
457
458
459
460
461
462
463
464
465
466
467
468
469
470
471
472
473
474
475
476
477
478
479
480
481
482
483
484
485
486
487
488
489
490
491
492
493
494
495
496
497
498
499
500
501
502
503
504
505
506
507
508
509
510
511
512
513
514
515
516
517
518
519
520
521
522
523
524
525
526
527
528
529
530
531
532
533
534
535
536
537
538
539
540
541
542
543
544
545
546
547
548
549
550
551
552
553
554
555
556
557
558
559
560
561
562
563
564
565
566
567
568
569
570
571
572
573
574
575
576
577
578
579
580
581
582
583
584
585
586
587
588
589
590
591
592
593
594
595
596
597
598
599
600
601
602
603
604
605
606
607
608
609
610
611
612
613
614
615
616
617
618
619
620
621
622
623
624
625
626
627
628
629
630
631
632
633
634
635
636
637
638
639
640
641
642
643
644
645
646
647
648
649
650
651
652
653
654
655
656
657
658
659
660
661
662
663
664
665
666
667
668
669
670
671
672
673
674
675
676
677
678
679
680
681
682
683
684
685
686
687
688
689
690
691
692
693
694
695
696
697
698
699
700
701
702
703
704
705
706
707
708
709
710
711
712
713
714
715
716
717
718
719
720
721
722
723
724
725
726
727
728
729
730
731
732
733
734
735
736
737
738
739
740
741
742
743
744
745
746
747
748
749
750
751
752
753
754
755
756
757
758
759
760
761
762
763
764
765
766
767
768
769
770
771
772
773
774
775
776
777
778
779
import spaces
import logging
import os
import random
import re
import sys
import tempfile
import uuid
import warnings
import atexit
import threading
from pathlib import Path
from io import BytesIO

import numpy as np
from PIL import Image
from diffusers import AutoencoderKL, FlowMatchEulerDiscreteScheduler
import gradio as gr
import torch
from transformers import AutoModelForCausalLM, AutoTokenizer
from plyfile import PlyData

sys.path.append(os.path.dirname(os.path.abspath(__file__)))
sys.path.append(os.path.join(os.path.dirname(os.path.abspath(__file__)), "ml-sharp", "src"))

from diffusers import ZImagePipeline
from diffusers.models.transformers.transformer_z_image import ZImageTransformer2DModel

# SHARP imports for 3D Gaussian splat generation
from sharp.models import create_predictor, PredictorParams
from sharp.utils.gaussians import save_ply
from sharp.cli.predict import predict_image, DEFAULT_MODEL_URL

# ==================== Environment Variables ==================================
MODEL_PATH = os.environ.get("MODEL_PATH", "Tongyi-MAI/Z-Image-Turbo")
ENABLE_COMPILE = os.environ.get("ENABLE_COMPILE", "false").lower() == "true"
ENABLE_WARMUP = os.environ.get("ENABLE_WARMUP", "false").lower() == "true"
ATTENTION_BACKEND = os.environ.get("ATTENTION_BACKEND", "native")
HF_TOKEN = os.environ.get("HF_TOKEN")
# =============================================================================

os.environ["TOKENIZERS_PARALLELISM"] = "false"
warnings.filterwarnings("ignore")
logging.getLogger("transformers").setLevel(logging.ERROR)

# Temporary file cleanup system
_temp_files_lock = threading.Lock()
_temp_files = []

def register_temp_file(path: str):
    """Register a temporary file for cleanup."""
    with _temp_files_lock:
        _temp_files.append(path)

def cleanup_temp_files():
    """Clean up all registered temporary files."""
    with _temp_files_lock:
        for path in _temp_files:
            try:
                if os.path.exists(path):
                    os.unlink(path)
            except Exception as e:
                print(f"Failed to delete temp file {path}: {e}")
        _temp_files.clear()

atexit.register(cleanup_temp_files)

def optimize_memory():
    """Clear CUDA cache and run garbage collection."""
    import gc
    gc.collect()
    if torch.cuda.is_available():
        torch.cuda.empty_cache()
        torch.cuda.synchronize()

RES_CHOICES = {
    "1024": [
        "1024x1024 ( 1:1 )",
        "1152x896 ( 9:7 )",
        "896x1152 ( 7:9 )",
        "1152x864 ( 4:3 )",
        "864x1152 ( 3:4 )",
        "1248x832 ( 3:2 )",
        "832x1248 ( 2:3 )",
        "1280x720 ( 16:9 )",
        "720x1280 ( 9:16 )",
        "1344x576 ( 21:9 )",
        "576x1344 ( 9:21 )",
    ],
    "1280": [
        "1280x1280 ( 1:1 )",
        "1440x1120 ( 9:7 )",
        "1120x1440 ( 7:9 )",
        "1472x1104 ( 4:3 )",
        "1104x1472 ( 3:4 )",
        "1536x1024 ( 3:2 )",
        "1024x1536 ( 2:3 )",
        "1536x864 ( 16:9 )",
        "864x1536 ( 9:16 )",
        "1680x720 ( 21:9 )",
        "720x1680 ( 9:21 )",
    ],
    "1536": [
        "1536x1536 ( 1:1 )",
        "1728x1344 ( 9:7 )",
        "1344x1728 ( 7:9 )",
        "1728x1296 ( 4:3 )",
        "1296x1728 ( 3:4 )",
        "1872x1248 ( 3:2 )",
        "1248x1872 ( 2:3 )",
        "2048x1152 ( 16:9 )",
        "1152x2048 ( 9:16 )",
        "2016x864 ( 21:9 )",
        "864x2016 ( 9:21 )",
    ],
}

RESOLUTION_SET = []
for resolutions in RES_CHOICES.values():
    RESOLUTION_SET.extend(resolutions)

EXAMPLE_PROMPTS = [
    ["Alien UFO landing in a dark forest with a starry sky"],
    [
        "Underwater city with futuristic buildings and colorful coral reefs, vibrant marine life swimming around, sunlight filtering through the water, digital art"
    ],
    [
        "A serene mountain landscape during autumn, with a clear blue lake reflecting the colorful foliage, high-resolution photograph"
    ],
    [
        "A bustling cyberpunk city street at night, neon signs in various languages,style of Syd Mead and Katsuhiro Otomo"
    ],
]


def get_resolution(resolution: str) -> tuple[int, int]:
    """Parse resolution string to width and height tuple."""
    match = re.search(r"(\d+)\s*[×x]\s*(\d+)", resolution)
    if match:
        return int(match.group(1)), int(match.group(2))
    return 1024, 1024


def load_models(model_path: str, enable_compile: bool = False, attention_backend: str = "native"):
    """
    Load all models required for Z-Image generation.
    Uses device_map="cuda" for ZeroGPU compatibility.
    """
    print(f"Loading models from {model_path}...")

    use_auth_token = HF_TOKEN if HF_TOKEN else True
    is_local = os.path.exists(model_path)

    # Load VAE
    if is_local:
        vae = AutoencoderKL.from_pretrained(
            os.path.join(model_path, "vae"),
            torch_dtype=torch.bfloat16,
            device_map="cuda",
        )
    else:
        vae = AutoencoderKL.from_pretrained(
            model_path,
            subfolder="vae",
            torch_dtype=torch.bfloat16,
            device_map="cuda",
            use_auth_token=use_auth_token,
        )

    # Load Text Encoder
    if is_local:
        text_encoder = AutoModelForCausalLM.from_pretrained(
            os.path.join(model_path, "text_encoder"),
            torch_dtype=torch.bfloat16,
            device_map="cuda",
        ).eval()
    else:
        text_encoder = AutoModelForCausalLM.from_pretrained(
            model_path,
            subfolder="text_encoder",
            torch_dtype=torch.bfloat16,
            device_map="cuda",
            use_auth_token=use_auth_token,
        ).eval()

    # Load Tokenizer
    if is_local:
        tokenizer = AutoTokenizer.from_pretrained(os.path.join(model_path, "tokenizer"))
    else:
        tokenizer = AutoTokenizer.from_pretrained(
            model_path,
            subfolder="tokenizer",
            use_auth_token=use_auth_token,
        )
    tokenizer.padding_side = "left"

    # Configure torch.compile optimizations
    if enable_compile:
        print("Enabling torch.compile optimizations...")
        torch._inductor.config.conv_1x1_as_mm = True
        torch._inductor.config.coordinate_descent_tuning = True
        torch._inductor.config.epilogue_fusion = False
        torch._inductor.config.coordinate_descent_check_all_directions = True
        torch._inductor.config.max_autotune_gemm = True
        torch._inductor.config.max_autotune_gemm_backends = "TRITON,ATEN"
        torch._inductor.config.triton.cudagraphs = False

    # Create pipeline
    pipe = ZImagePipeline(
        scheduler=None,
        vae=vae,
        text_encoder=text_encoder,
        tokenizer=tokenizer,
        transformer=None,
    )

    if enable_compile:
        pipe.vae.disable_tiling()

    # Load Transformer
    if is_local:
        transformer = ZImageTransformer2DModel.from_pretrained(
            os.path.join(model_path, "transformer")
        ).to("cuda", torch.bfloat16)
    else:
        transformer = ZImageTransformer2DModel.from_pretrained(
            model_path,
            subfolder="transformer",
            use_auth_token=use_auth_token,
        ).to("cuda", torch.bfloat16)

    pipe.transformer = transformer
    pipe.transformer.set_attention_backend(attention_backend)

    if enable_compile:
        print("Compiling transformer...")
        pipe.transformer = torch.compile(
            pipe.transformer, mode="max-autotune-no-cudagraphs", fullgraph=False
        )

    pipe.to("cuda", torch.bfloat16)

    print("Models loaded successfully")
    return pipe


def generate_image(
    pipe,
    prompt: str,
    resolution: str = "1024x1024",
    seed: int = 42,
    guidance_scale: float = 0.0,
    num_inference_steps: int = 9,
    shift: float = 3.0,
    max_sequence_length: int = 512,
    progress=gr.Progress(track_tqdm=True),
):
    """Generate a single image using the Z-Image pipeline."""
    width, height = get_resolution(resolution)

    generator = torch.Generator("cuda").manual_seed(seed)
    scheduler = FlowMatchEulerDiscreteScheduler(num_train_timesteps=1000, shift=shift)
    pipe.scheduler = scheduler

    image = pipe(
        prompt=prompt,
        height=height,
        width=width,
        guidance_scale=guidance_scale,
        num_inference_steps=num_inference_steps,
        generator=generator,
        max_sequence_length=max_sequence_length,
    ).images[0]

    return image


def warmup_model(pipe, resolutions: list[str]):
    """Warm up the model by running inference on dummy prompts."""
    print("Starting warmup phase...")
    dummy_prompt = "warmup"

    for res_str in resolutions:
        print(f"Warming up for resolution: {res_str}")
        try:
            for i in range(3):
                generate_image(
                    pipe,
                    prompt=dummy_prompt,
                    resolution=res_str,
                    num_inference_steps=9,
                    guidance_scale=0.0,
                    seed=42 + i,
                )
        except Exception as e:
            print(f"Warmup failed for {res_str}: {e}")

    print("Warmup completed.")


# Global pipeline instance
pipe = None
# Global SHARP splat predictor (lazy loaded)
splat_predictor = None


def init_app():
    """Initialize the application by loading models."""
    global pipe

    try:
        pipe = load_models(
            MODEL_PATH,
            enable_compile=ENABLE_COMPILE,
            attention_backend=ATTENTION_BACKEND,
        )
        print(f"Model loaded. Compile: {ENABLE_COMPILE}, Backend: {ATTENTION_BACKEND}")

        if ENABLE_WARMUP:
            all_resolutions = []
            for cat in RES_CHOICES.values():
                all_resolutions.extend(cat)
            warmup_model(pipe, all_resolutions)

    except Exception as e:
        print(f"Error loading model: {e}")
        import traceback
        traceback.print_exc()
        pipe = None


def load_splat_predictor(device: str = "cuda"):
    """Load the SHARP Gaussian splat predictor model (lazy loading)."""
    global splat_predictor

    if splat_predictor is not None:
        return splat_predictor

    print(f"Loading SHARP splat predictor from {DEFAULT_MODEL_URL}...")

    # Download and load the model
    state_dict = torch.hub.load_state_dict_from_url(
        DEFAULT_MODEL_URL,
        progress=True,
    )

    predictor = create_predictor(PredictorParams())
    predictor.load_state_dict(state_dict)
    predictor.eval()
    predictor.to(device)

    splat_predictor = predictor
    print("SHARP predictor loaded successfully")
    return splat_predictor


def convert_ply_to_splat(ply_file_path: str) -> bytes:
    """
    Convert a PLY file to SPLAT format for the antimatter15 viewer.
    Returns the splat data as bytes.
    """
    plydata = PlyData.read(ply_file_path)
    vert = plydata["vertex"]
    
    sorted_indices = np.argsort(
        -np.exp(vert["scale_0"] + vert["scale_1"] + vert["scale_2"])
        / (1 + np.exp(-vert["opacity"]))
    )
    
    buffer = BytesIO()
    for idx in sorted_indices:
        v = plydata["vertex"][idx]
        position = np.array([v["x"], v["y"], v["z"]], dtype=np.float32)
        scales = np.exp(
            np.array([v["scale_0"], v["scale_1"], v["scale_2"]], dtype=np.float32)
        )
        color = np.array([
            0.5 + 0.28209479177387814 * v["f_dc_0"],
            0.5 + 0.28209479177387814 * v["f_dc_1"],
            0.5 + 0.28209479177387814 * v["f_dc_2"],
            1 / (1 + np.exp(-v["opacity"])),
        ])
        rot = np.array([v["rot_0"], v["rot_1"], v["rot_2"], v["rot_3"]], dtype=np.float32)
        buffer.write(position.tobytes())
        buffer.write(scales.tobytes())
        buffer.write((color * 255).clip(0, 255).astype(np.uint8).tobytes())
        buffer.write(
            ((rot / np.linalg.norm(rot)) * 128 + 128).clip(0, 255).astype(np.uint8).tobytes()
        )
    
    return buffer.getvalue()



@spaces.GPU
def generate_splat(selected_image, progress=gr.Progress(track_tqdm=True)):
    """Generate a 3D Gaussian splat from the selected image."""
    if selected_image is None:
        raise gr.Error("Please select an image from the gallery first")

    try:
        if isinstance(selected_image, str):
            if not os.path.exists(selected_image):
                raise gr.Error(f"Image file not found: {selected_image}")
            pil_image = Image.open(selected_image).convert("RGB")
            image_np = np.array(pil_image)
        elif hasattr(selected_image, "convert"):
            image_np = np.array(selected_image.convert("RGB"))
        elif isinstance(selected_image, np.ndarray):
            image_np = selected_image.copy()
        else:
            image_np = np.array(selected_image)

        if image_np is None or image_np.size == 0:
            raise gr.Error("Invalid image data")

        if image_np.ndim == 2:
            image_np = np.stack([image_np] * 3, axis=-1)
        elif image_np.ndim == 3 and image_np.shape[-1] == 4:
            image_np = image_np[:, :, :3]
        elif image_np.ndim != 3 or image_np.shape[-1] != 3:
            raise gr.Error(f"Unexpected image shape: {image_np.shape}")

        height, width = image_np.shape[:2]
        if height < 64 or width < 64:
            raise gr.Error(f"Image too small: {width}x{height}. Minimum is 64x64.")

        f_mm = 30.0
        f_px = f_mm * np.sqrt(width**2 + height**2) / np.sqrt(36**2 + 24**2)
        device = torch.device("cuda" if torch.cuda.is_available() else "cpu")
        predictor = load_splat_predictor(str(device))

        print(f"Generating 3D Gaussian splat from {width}x{height} image...")
        gaussians = predict_image(predictor, image_np, f_px, device)
        num_gaussians = gaussians.mean_vectors.shape[1]
        print(f"Generated {num_gaussians} gaussians")

        # Use gradio's temp directory for better compatibility with Spaces
        cache_dir = Path(tempfile.gettempdir()) / "gradio_cache"
        cache_dir.mkdir(exist_ok=True)
        ply_path = cache_dir / f"scene_{uuid.uuid4().hex[:8]}.ply"

        save_ply(gaussians, f_px, (height, width), ply_path)
        register_temp_file(str(ply_path))

        status_msg = f"✅ Generated {num_gaussians:,} gaussians | PLY: {ply_path.stat().st_size/1024:.1f}KB"
        return (
            str(ply_path),  # ply_download - return path directly
            str(ply_path),  # ply_path_state
            gr.update(visible=True),  # convert_splat_btn
            status_msg,  # splat_status
        )

    except gr.Error:
        raise
    except Exception as e:
        print(f"Error generating splat: {e}")
        import traceback
        traceback.print_exc()
        raise gr.Error(f"Failed to generate 3D splat: {str(e)}")


def convert_and_save_splat(ply_path):
    """Convert PLY to SPLAT format and save to temp file."""
    if not ply_path or not os.path.exists(ply_path):
        raise gr.Error("PLY file not found. Please generate a 3D splat first.")

    try:
        splat_data = convert_ply_to_splat(ply_path)

        # Use same cache directory as PLY files
        cache_dir = Path(tempfile.gettempdir()) / "gradio_cache"
        cache_dir.mkdir(exist_ok=True)
        splat_path = cache_dir / f"scene_{uuid.uuid4().hex[:8]}.splat"

        with open(splat_path, "wb") as f:
            f.write(splat_data)
        register_temp_file(str(splat_path))

        size_kb = splat_path.stat().st_size / 1024
        status_msg = f"✅ SPLAT file created | Size: {size_kb:.1f}KB"
        return status_msg, str(splat_path)  # Return path directly
    except Exception as e:
        print(f"Error converting to SPLAT: {e}")
        raise gr.Error(f"Failed to convert to SPLAT: {str(e)}")


@spaces.GPU
def generate(
    prompt: str,
    resolution: str = "1024x1024 ( 1:1 )",
    seed: int = 42,
    steps: int = 9,
    shift: float = 3.0,
    random_seed: bool = True,
    gallery_images: list = None,
    progress=gr.Progress(track_tqdm=True),
):
    """
    Generate an image using the Z-Image model based on the provided prompt and settings.

    Args:
        prompt: Text prompt describing the desired image content
        resolution: Output resolution in format "WIDTHxHEIGHT ( RATIO )"
        seed: Seed for reproducible generation
        steps: Number of inference steps for the diffusion process
        shift: Time shift parameter for the flow matching scheduler
        random_seed: Whether to generate a new random seed
        gallery_images: List of previously generated images to append to
        progress: Gradio progress tracker

    Returns:
        tuple: (gallery_images, seed_str, seed_int)
    """
    if random_seed:
        new_seed = random.randint(1, 1000000)
    else:
        new_seed = seed if seed != -1 else random.randint(1, 1000000)

    if pipe is None:
        raise gr.Error("Model not loaded. Please check the console for errors.")

    # Parse resolution
    try:
        resolution_str = resolution.split(" ")[0]
    except:
        resolution_str = "1024x1024"

    # Generate image
    image = generate_image(
        pipe=pipe,
        prompt=prompt,
        resolution=resolution_str,
        seed=new_seed,
        guidance_scale=0.0,
        num_inference_steps=int(steps + 1),
        shift=shift,
    )

    if gallery_images is None:
        gallery_images = []

    # Add latest output to the top of the list
    gallery_images = [image] + gallery_images

    return gallery_images, new_seed


@spaces.GPU
def generate_batch(
    prompt: str,
    resolution: str,
    seed: int,
    steps: int,
    shift: float,
    batch_size: int = 2,
    gallery_images: list = None,
    progress=gr.Progress(track_tqdm=True),
):
    """Generate multiple images in a batch for efficiency."""
    if pipe is None:
        raise gr.Error("Model not loaded.")
    
    if gallery_images is None:
        gallery_images = []
    
    new_images = []
    for i in range(batch_size):
        current_seed = seed + i
        image = generate_image(
            pipe=pipe,
            prompt=prompt,
            resolution=resolution.split(" ")[0],
            seed=current_seed,
            guidance_scale=0.0,
            num_inference_steps=int(steps + 1),
            shift=shift,
        )
        new_images.append(image)
    
    optimize_memory()
    return new_images + gallery_images, seed


# Initialize the app
init_app()


# ==================== Gradio UI ====================

css = """
.fillable{max-width: 1230px !important}
"""

with gr.Blocks(title="Z-Image Demo") as demo:
    gr.Markdown(
        """<div align="center">

# Generative 3D Gaussian Splat

* Generate images from text prompts using [![GitHub](https://img.shields.io/badge/GitHub-Z--Image-181717?logo=github&logoColor=white)](https://github.com/Tongyi-MAI/Z-Image) 

* Create 3D Gaussian splat models from generated images using [![GitHub](https://img.shields.io/badge/GitHub-SHARP-181717?logo=github&logoColor=white)](https://github.com/apple/ml-sharp)

</div>"""
    )

    with gr.Row():
        with gr.Column(scale=1):
            prompt_input = gr.Textbox(
                label="Prompt",
                lines=3,
                placeholder="Enter your prompt here...",
            )

            with gr.Row():
                choices = [int(k) for k in RES_CHOICES.keys()]
                res_cat = gr.Dropdown(
                    value=1024,
                    choices=choices,
                    label="Resolution Category",
                )

                initial_res_choices = RES_CHOICES["1024"]
                resolution = gr.Dropdown(
                    value=initial_res_choices[0],
                    choices=RESOLUTION_SET,
                    label="Width x Height (Ratio)",
                )

            with gr.Row():
                seed = gr.Number(label="Seed", value=42, precision=0)
                random_seed = gr.Checkbox(label="Random Seed", value=True)

            with gr.Row():
                steps = gr.Slider(
                    label="Steps",
                    minimum=1,
                    maximum=100,
                    value=8,
                    step=1,
                )
                shift = gr.Slider(
                    label="Time Shift",
                    minimum=1.0,
                    maximum=10.0,
                    value=3.0,
                    step=0.1,
                )

            generate_btn = gr.Button("Generate", variant="primary")

            # Example prompts
            gr.Markdown("### 📝 Example Prompts")
            gr.Examples(examples=EXAMPLE_PROMPTS, inputs=prompt_input, label=None)

        with gr.Column(scale=1):
            output_gallery = gr.Gallery(
                label="Generated Images",
                columns=2,
                rows=2,
                height=600,
                object_fit="contain",
                format="png",
                interactive=False,
            )

            # 3D Gaussian Splat Generation Section
            with gr.Accordion("3D Gaussian Splat Generation", open=False):
#                 gr.Markdown(
#                     """Click on an image in the gallery above to select it, then click "Generate 3D Splat" to create
# a 3D Gaussian splat model. You can then download PLY files or SPLAT in an additional step."""
#                 )

                # State to hold selected image and PLY path
                selected_image_state = gr.State(value=None)
                ply_path_state = gr.State(value=None)

                with gr.Row():
                    generate_splat_btn = gr.Button(
                        "Generate 3D Scene",
                        variant="secondary",
                        interactive=False,
                    )
                    convert_splat_btn = gr.Button(
                        "Convert to SPLAT",
                        variant="secondary",
                        visible=False,
                    )

                splat_status = gr.Textbox(
                    label="Status",
                    interactive=False,
                    visible=True,
                    value="Click an image in the gallery to select it",
                    lines=2
                )

                # Download section - always visible, empty until files are generated
                with gr.Row():
                    ply_download = gr.File(
                        label="PLY File (click to download)",
                        visible=True,
                        interactive=False,
                        value=None,
                    )
                    splat_download = gr.File(
                        label="SPLAT File (click to download)",
                        visible=True,
                        interactive=False,
                        value=None,
                    )

    # Gallery image selection handler
    def on_gallery_select(evt: gr.SelectData, gallery_images):
        """Handle gallery image selection."""
        if gallery_images is None or len(gallery_images) == 0:
            return None, gr.update(interactive=False), "No image selected"

        selected_idx = evt.index
        if selected_idx < len(gallery_images):
            selected_item = gallery_images[selected_idx]
            if isinstance(selected_item, tuple):
                selected_img = selected_item[0]
            elif isinstance(selected_item, dict):
                selected_img = selected_item.get('image') or selected_item.get('name')
            else:
                selected_img = selected_item
            return (
                selected_img,
                gr.update(interactive=True),
                f"Image {selected_idx + 1} selected - ready to generate 3D splat",
            )

        return None, gr.update(interactive=False), "Selection error"

    output_gallery.select(
        on_gallery_select,
        inputs=[output_gallery],
        outputs=[selected_image_state, generate_splat_btn, splat_status],
    )

    # Splat generation handler
    generate_splat_btn.click(
        generate_splat,
        inputs=[selected_image_state],
        outputs=[ply_download, ply_path_state, convert_splat_btn, splat_status],
    )

    # SPLAT conversion handler
    convert_splat_btn.click(
        convert_and_save_splat,
        inputs=[ply_path_state],
        outputs=[splat_status, splat_download],
    )

    def update_res_choices(res_cat_value):
        """Update resolution choices based on selected category."""
        if str(res_cat_value) in RES_CHOICES:
            res_choices = RES_CHOICES[str(res_cat_value)]
        else:
            res_choices = RES_CHOICES["1024"]
        return gr.update(value=res_choices[0], choices=res_choices)

    res_cat.change(
        update_res_choices,
        inputs=res_cat,
        outputs=resolution,
    )

    generate_btn.click(
        generate,
        inputs=[prompt_input, resolution, seed, steps, shift, random_seed, output_gallery],
        outputs=[output_gallery, seed],
    )

if __name__ == "__main__":
    demo.launch(css=css)