Kyle Pearson commited on
Commit
d3cfc5a
·
1 Parent(s): 0f3547d

fixes to interface

Browse files
Files changed (1) hide show
  1. app.py +46 -8
app.py CHANGED
@@ -441,7 +441,12 @@ def generate_splat(selected_image, progress=gr.Progress(track_tqdm=True)):
441
  register_temp_file(str(ply_path))
442
 
443
  status_msg = f"✅ Generated {num_gaussians:,} gaussians | PLY: {ply_path.stat().st_size/1024:.1f}KB"
444
- return str(ply_path), status_msg
 
 
 
 
 
445
 
446
  except gr.Error:
447
  raise
@@ -452,6 +457,26 @@ def generate_splat(selected_image, progress=gr.Progress(track_tqdm=True)):
452
  raise gr.Error(f"Failed to generate 3D splat: {str(e)}")
453
 
454
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
455
  @spaces.GPU
456
  def generate(
457
  prompt: str,
@@ -639,8 +664,9 @@ with gr.Blocks(title="Z-Image Demo") as demo:
639
  a 3D Gaussian splat model. You can then download PLY files or SPLAT in an additional step."""
640
  )
641
 
642
- # State to hold selected image
643
  selected_image_state = gr.State(value=None)
 
644
 
645
  with gr.Row():
646
  generate_splat_btn = gr.Button(
@@ -648,6 +674,11 @@ a 3D Gaussian splat model. You can then download PLY files or SPLAT in an additi
648
  variant="secondary",
649
  interactive=False,
650
  )
 
 
 
 
 
651
 
652
  splat_status = gr.Textbox(
653
  label="Status",
@@ -659,7 +690,11 @@ a 3D Gaussian splat model. You can then download PLY files or SPLAT in an additi
659
  # Download buttons
660
  with gr.Row():
661
  ply_download = gr.File(
662
- label="Download PLY File",
 
 
 
 
663
  visible=False,
664
  )
665
 
@@ -696,11 +731,14 @@ a 3D Gaussian splat model. You can then download PLY files or SPLAT in an additi
696
  generate_splat_btn.click(
697
  generate_splat,
698
  inputs=[selected_image_state],
699
- outputs=[ply_download, splat_status],
700
- ).success(
701
- lambda ply_path, status: gr.update(visible=True, value=ply_path) if ply_path and os.path.exists(ply_path) else gr.update(visible=False),
702
- inputs=[ply_download, splat_status],
703
- outputs=[ply_download],
 
 
 
704
  )
705
 
706
  def update_res_choices(res_cat_value):
 
441
  register_temp_file(str(ply_path))
442
 
443
  status_msg = f"✅ Generated {num_gaussians:,} gaussians | PLY: {ply_path.stat().st_size/1024:.1f}KB"
444
+ return (
445
+ gr.update(visible=True, value=str(ply_path)), # ply_download
446
+ str(ply_path), # ply_path_state
447
+ gr.update(visible=True), # convert_splat_btn
448
+ status_msg, # splat_status
449
+ )
450
 
451
  except gr.Error:
452
  raise
 
457
  raise gr.Error(f"Failed to generate 3D splat: {str(e)}")
458
 
459
 
460
+ def convert_and_save_splat(ply_path):
461
+ """Convert PLY to SPLAT format and save to temp file."""
462
+ if not ply_path or not os.path.exists(ply_path):
463
+ raise gr.Error("PLY file not found. Please generate a 3D splat first.")
464
+
465
+ try:
466
+ splat_data = convert_ply_to_splat(ply_path)
467
+ splat_file = tempfile.NamedTemporaryFile(suffix=".splat", delete=False)
468
+ splat_file.write(splat_data)
469
+ splat_file.close()
470
+ register_temp_file(splat_file.name)
471
+
472
+ size_kb = os.path.getsize(splat_file.name) / 1024
473
+ status_msg = f"✅ SPLAT file created | Size: {size_kb:.1f}KB"
474
+ return status_msg, gr.update(visible=True, value=splat_file.name)
475
+ except Exception as e:
476
+ print(f"Error converting to SPLAT: {e}")
477
+ raise gr.Error(f"Failed to convert to SPLAT: {str(e)}")
478
+
479
+
480
  @spaces.GPU
481
  def generate(
482
  prompt: str,
 
664
  a 3D Gaussian splat model. You can then download PLY files or SPLAT in an additional step."""
665
  )
666
 
667
+ # State to hold selected image and PLY path
668
  selected_image_state = gr.State(value=None)
669
+ ply_path_state = gr.State(value=None)
670
 
671
  with gr.Row():
672
  generate_splat_btn = gr.Button(
 
674
  variant="secondary",
675
  interactive=False,
676
  )
677
+ convert_splat_btn = gr.Button(
678
+ "Convert to SPLAT",
679
+ variant="secondary",
680
+ visible=False,
681
+ )
682
 
683
  splat_status = gr.Textbox(
684
  label="Status",
 
690
  # Download buttons
691
  with gr.Row():
692
  ply_download = gr.File(
693
+ label="Download PLY",
694
+ visible=False,
695
+ )
696
+ splat_download = gr.File(
697
+ label="Download SPLAT",
698
  visible=False,
699
  )
700
 
 
731
  generate_splat_btn.click(
732
  generate_splat,
733
  inputs=[selected_image_state],
734
+ outputs=[ply_download, ply_path_state, convert_splat_btn, splat_status],
735
+ )
736
+
737
+ # SPLAT conversion handler
738
+ convert_splat_btn.click(
739
+ convert_and_save_splat,
740
+ inputs=[ply_path_state],
741
+ outputs=[splat_status, splat_download],
742
  )
743
 
744
  def update_res_choices(res_cat_value):