Hecmac1/CalcGen-fx-CG50
0
1from __future__ import annotations2 3import secrets4import tempfile5from pathlib import Path6 7import gradio as gr8import numpy as np9from PIL import Image10import torch11 12try:13 import spaces14 ZERO_GPU_AVAILABLE = True15except ImportError: # Local smoke tests do not run inside a ZeroGPU Space.16 ZERO_GPU_AVAILABLE = False17 class _SpacesFallback:18 @staticmethod19 def GPU(*_args, **_kwargs):20 def decorate(function):21 return function22 23 return decorate24 25 spaces = _SpacesFallback()26 27from calcgen_model import GuidedCalculatorGenerator28from calcgen_renderer import build_inputs29 30 31ROOT = Path(__file__).resolve().parent32MODEL_PATH = ROOT / "model" / "calcgen-v7-fp32.pt"33PLAN_COUNT = 40034ANIMATION_FRAMES = 1235ANIMATION_DURATION_MS = 11036 37 38def load_model() -> GuidedCalculatorGenerator:39 checkpoint = torch.load(MODEL_PATH, map_location="cpu", weights_only=True)40 model = GuidedCalculatorGenerator(condition_dim=int(checkpoint["condition_dim"]))41 model.load_state_dict(checkpoint["model"])42 model.eval()43 return model44 45 46MODEL = load_model()47 48 49@torch.inference_mode()50def _generate_impl(plan: int, yaw: float, pitch: float, roll: float, use_gpu: bool = False):51 plan = max(0, min(PLAN_COUNT - 1, int(plan)))52 condition, guide, details = build_inputs(plan, float(yaw), float(pitch), float(roll))53 device = torch.device("cuda" if use_gpu else "cpu")54 condition_tensor = torch.tensor(condition, dtype=torch.float32, device=device).unsqueeze(0)55 guide_array = np.asarray(guide, dtype=np.float32).transpose(2, 0, 1) / 255.056 guide_tensor = torch.from_numpy(guide_array.copy()).unsqueeze(0).to(device)57 try:58 MODEL.to(device)59 prediction = MODEL(condition_tensor, guide_tensor)[0].cpu()60 finally:61 if device.type == "cuda":62 MODEL.to("cpu")63 torch.cuda.empty_cache()64 output = _prediction_to_image(prediction)65 caption = (66 f"Plan {plan + 1:03d}/400 · LCD {details['lcd']} · "67 f"{details['columns']}×{details['rows']} keypad · {details['shape']} keys"68 )69 return output, guide, caption70 71 72def _prediction_to_image(prediction: torch.Tensor) -> Image.Image:73 pixels = ((prediction.clamp(-1, 1) + 1.0) * 127.5).round().byte()74 return Image.fromarray(pixels.permute(1, 2, 0).cpu().numpy(), mode="RGB")75 76 77def _animation_angles(yaw: float, pitch: float, roll: float):78 """Return a bounded, looping orbit around the selected pose."""79 phase = np.linspace(0.0, 2.0 * np.pi, ANIMATION_FRAMES, endpoint=False)80 for t in phase:81 yield (82 float(np.clip(float(yaw) + 32.0 * np.cos(t), -58.0, 58.0)),83 float(np.clip(float(pitch) + 18.0 * np.sin(t), -44.0, 44.0)),84 float(roll),85 )86 87 88@torch.inference_mode()89def _animation_impl(plan: int, yaw: float, pitch: float, roll: float, use_gpu: bool = False):90 plan = max(0, min(PLAN_COUNT - 1, int(plan)))91 angles = list(_animation_angles(yaw, pitch, roll))92 conditions = []93 guides = []94 details = None95 for frame_yaw, frame_pitch, frame_roll in angles:96 condition, guide, details = build_inputs(plan, frame_yaw, frame_pitch, frame_roll)97 conditions.append(condition)98 guide_array = np.asarray(guide, dtype=np.float32).transpose(2, 0, 1) / 255.099 guides.append(guide_array)100 101 device = torch.device("cuda" if use_gpu else "cpu")102 condition_tensor = torch.tensor(np.asarray(conditions), dtype=torch.float32, device=device)103 guide_tensor = torch.from_numpy(np.stack(guides).copy()).to(device)104 try:105 MODEL.to(device)106 predictions = MODEL(condition_tensor, guide_tensor).cpu()107 finally:108 if device.type == "cuda":109 MODEL.to("cpu")110 torch.cuda.empty_cache()111 112 output_frames = [_prediction_to_image(frame) for frame in predictions]113 guide_frames = [114 Image.fromarray((guide.transpose(1, 2, 0) * 255.0).round().astype(np.uint8), mode="RGB")115 for guide in guides116 ]117 output_path = tempfile.NamedTemporaryFile(suffix=".gif", delete=False).name118 guide_path = tempfile.NamedTemporaryFile(suffix=".gif", delete=False).name119 output_frames[0].save(120 output_path, save_all=True, append_images=output_frames[1:], duration=ANIMATION_DURATION_MS, loop=0121 )122 guide_frames[0].save(123 guide_path, save_all=True, append_images=guide_frames[1:], duration=ANIMATION_DURATION_MS, loop=0124 )125 caption = (126 f"Plan {plan + 1:03d}/400 · {ANIMATION_FRAMES}-frame orbit · LCD {details['lcd']} · "127 f"{details['columns']}×{details['rows']} keypad · {details['shape']} keys"128 )129 return output_path, guide_path, caption130 131 132@spaces.GPU(duration=30)133def generate(plan: int, yaw: float, pitch: float, roll: float):134 return _generate_impl(plan, yaw, pitch, roll, use_gpu=ZERO_GPU_AVAILABLE)135 136 137@spaces.GPU(duration=30)138def randomize(yaw: float, pitch: float, roll: float):139 plan = secrets.randbelow(PLAN_COUNT)140 output, guide, caption = _generate_impl(plan, yaw, pitch, roll, use_gpu=ZERO_GPU_AVAILABLE)141 return plan, output, guide, caption142 143 144@spaces.GPU(duration=30)145def generate_animation(plan: int, yaw: float, pitch: float, roll: float):146 return _animation_impl(plan, yaw, pitch, roll, use_gpu=ZERO_GPU_AVAILABLE)147 148 149CSS = """150.gradio-container { max-width: 1040px !important; }151.pixel-output img { image-rendering: pixelated; object-fit: contain !important; }152.pixel-output { min-height: 430px; }153.guide-output img { image-rendering: pixelated; object-fit: contain !important; }154.guide-output { min-height: 430px; }155#details { text-align: center; color: var(--body-text-color-subdued); }156"""157 158 159with gr.Blocks(title="CalcGen fx-CG50") as demo:160 gr.Markdown(161 """162 # CalcGen fx-CG50163 A **277,891-parameter conditional convolutional decoder** trained to generate164 ordinary rectangular calculators.165 """166 )167 168 with gr.Row():169 with gr.Column(scale=1, min_width=270):170 plan = gr.Slider(0, PLAN_COUNT - 1, value=0, step=1, label="Calculator design")171 yaw = gr.Slider(-58, 58, value=0, step=1, label="Yaw")172 pitch = gr.Slider(-44, 44, value=0, step=1, label="Pitch")173 roll = gr.Slider(-27, 27, value=0, step=1, label="Roll")174 with gr.Row():175 random_button = gr.Button("Randomize", variant="secondary")176 generate_button = gr.Button("Generate", variant="primary")177 animation_button = gr.Button("Generate Animation", variant="secondary")178 with gr.Column(scale=2):179 with gr.Row():180 output = gr.Image(181 label="Generated image",182 type="pil",183 format="png",184 height=430,185 width=240,186 interactive=False,187 buttons=["download", "fullscreen"],188 elem_classes="pixel-output",189 )190 guide = gr.Image(191 label="Semantic guide",192 type="pil",193 format="png",194 height=430,195 width=240,196 interactive=False,197 buttons=["download", "fullscreen"],198 elem_classes="guide-output",199 )200 with gr.Row():201 animation = gr.Image(202 label="Generated animation",203 type="filepath",204 format="gif",205 height=430,206 width=240,207 interactive=False,208 buttons=["download", "fullscreen"],209 elem_classes="pixel-output",210 )211 animation_guide = gr.Image(212 label="Animation guide",213 type="filepath",214 format="gif",215 height=430,216 width=240,217 interactive=False,218 buttons=["download", "fullscreen"],219 elem_classes="guide-output",220 )221 details = gr.Markdown(elem_id="details")222 223 inputs = [plan, yaw, pitch, roll]224 outputs = [output, guide, details]225 generate_button.click(226 generate,227 inputs=inputs,228 outputs=outputs,229 api_name="generate",230 concurrency_limit=1,231 concurrency_id="calcgen-gpu",232 )233 random_button.click(234 randomize,235 inputs=[yaw, pitch, roll],236 outputs=[plan, output, guide, details],237 api_name="randomize",238 concurrency_limit=1,239 concurrency_id="calcgen-gpu",240 )241 animation_button.click(242 generate_animation,243 inputs=inputs,244 outputs=[animation, animation_guide, details],245 api_name="generate_animation",246 concurrency_limit=1,247 concurrency_id="calcgen-gpu",248 )249 # The initial example is cheap enough on CPU; reserve ZeroGPU allocation for250 # explicit Generate and Randomize actions.251 demo.load(_generate_impl, inputs=inputs, outputs=outputs)252 253 254if __name__ == "__main__":255 demo.queue(default_concurrency_limit=1).launch(css=CSS)256 