Spaces:
Running on Zero
Running on Zero
Download app.py from AlayaLab/FloodDiffusion2-Live: direct link, hf CLI and curl.
- Browser
- Download file 18.2 kB
-
https://huggingface.co/spaces/AlayaLab/FloodDiffusion2-Live/resolve/main/app.py
- Command line
-
hf download hf://spaces/AlayaLab/FloodDiffusion2-Live/app.py
-
curl -L -o app.py https://huggingface.co/spaces/AlayaLab/FloodDiffusion2-Live/resolve/main/app.py
18.2 kB
| import spaces | |
| from spaces.config import Config as SpacesConfig | |
| from collections import deque | |
| import base64 | |
| import os | |
| from pathlib import Path | |
| import sys | |
| import tempfile | |
| import threading | |
| import time | |
| import uuid | |
| import gradio as gr | |
| from huggingface_hub import hf_hub_download | |
| import numpy as np | |
| import torch | |
| ROOT = Path(__file__).resolve().parent | |
| sys.path.insert(0, str(ROOT / 'space')) | |
| from space.assets_loader import resolve_t5, REVISION, REPO | |
| from space.inference import load_model | |
| from space.text_encoder import load_text_encoder, load_prompt_features, save_prompt_features | |
| from fast_window_runtime import FastWindowRuntime | |
| from hml263_runtime import load_hml263_model | |
| from window_control import WindowPath as SeedPath | |
| from hml263_control import WindowPath as HumanML3DPath | |
| from window_viewer import create_viewer | |
| from cloud_sessions import SessionStore | |
| MODELS = ('SEED', 'HumanML3D') | |
| DEFAULTS = [('A person walks forward.', .8), | |
| ('A person runs forward.', 2.5), | |
| ('A person walks forward in a crouched position.', .45), | |
| ('A person is dancing.', .8)] | |
| GPU_CONCURRENCY = None if SpacesConfig.zero_gpu else 1 | |
| LOCAL_GPU_LOCK = threading.Lock() | |
| STORE = SessionStore(Path(tempfile.gettempdir()) / 'flood2_live_window_controls') | |
| FEATURES = ROOT / 'space/assets/prompts.npz' | |
| DEFAULT_BANK = load_prompt_features(ROOT / 'space/assets/default_prompts.npz') | |
| BASE_BANK = load_prompt_features(FEATURES) | |
| seed_checkpoint = os.getenv('FLOOD2_SEED_CHECKPOINT') or hf_hub_download( | |
| REPO, 'checkpoints/seed_path_fk_300k/model.ckpt', revision=REVISION) | |
| hml_checkpoint = os.getenv('FLOOD2_HML_CHECKPOINT') or hf_hub_download( | |
| REPO, 'checkpoints/humanml3d_babel_path_200k/model.ckpt', revision=REVISION) | |
| hml_config = os.getenv('FLOOD2_HML_CONFIG') or hf_hub_download( | |
| REPO, 'checkpoints/humanml3d_babel_path_200k/config.yaml', revision=REVISION) | |
| RUNTIMES = { | |
| 'SEED': FastWindowRuntime.from_runtime(load_model(ROOT/'space/vendor', seed_checkpoint, FEATURES)), | |
| 'HumanML3D': load_hml263_model(ROOT/'space/vendor', hml_checkpoint, hml_config, FEATURES), | |
| } | |
| # Module-scope loads are virtualized by ZeroGPU. Encoding and graph capture | |
| # happen only inside decorated allocations. | |
| encoder_path, tokenizer_path = resolve_t5() | |
| ENCODER = load_text_encoder(ROOT/'space/vendor', encoder_path, tokenizer_path) | |
| def owner(request): | |
| value = getattr(request, 'session_hash', None) | |
| if not value: | |
| raise gr.Error('Reload the page to start a session.') | |
| return value | |
| def checked_model(value): | |
| if value not in MODELS: | |
| raise gr.Error('Choose SEED or HumanML3D.') | |
| return value | |
| def read(session, request): | |
| try: | |
| return STORE.read(session, owner(request)) | |
| except (ValueError, FileNotFoundError) as exc: | |
| raise gr.Error(str(exc)) from exc | |
| def prepare(previous_session, model, p1, p2, p3, p4, v1, v2, v3, v4, request: gr.Request): | |
| model = checked_model(model) | |
| browser = owner(request) | |
| prompts = [str(prompt).strip() for prompt in (p1, p2, p3, p4)] | |
| if any(not prompt or len(prompt) > 500 for prompt in prompts): | |
| raise gr.Error('Enter 1 to 500 characters in every prompt slot.') | |
| try: | |
| speeds = [float(value) for value in (v1, v2, v3, v4)] | |
| except (TypeError, ValueError): | |
| raise gr.Error('Every speed must be between 0 and 5 m/s.') | |
| if any(not np.isfinite(value) or not 0 <= value <= 5 for value in speeds): | |
| raise gr.Error('Every speed must be between 0 and 5 m/s.') | |
| if previous_session: | |
| previous = read(previous_session, request) | |
| if previous['state'] not in ('paused', 'error') or previous.get('stream_active'): | |
| raise gr.Error('Stop before starting a new motion.') | |
| bank = {text: DEFAULT_BANK[text] for text in set(prompts) if text in DEFAULT_BANK} | |
| missing = [text for text in dict.fromkeys(prompts) if text not in bank] | |
| if missing: | |
| bank.update(ENCODER.encode(missing)) | |
| session = uuid.uuid4().hex | |
| save_prompt_features(STORE.path(session).with_suffix('.npz'), bank) | |
| STORE.create(session, dict(owner=browser, model=model, slots=prompts, speeds=speeds, | |
| slot=0, x=0., z=0., client_seq=-1, keys_updated_at=0., | |
| version=0, state='starting', stop=False, stream_active=False, | |
| frames=0, runtime={}, error=None)) | |
| return session | |
| def control(session, seq, x, z, shift=False, slot=0, request: gr.Request=None): | |
| if not session: | |
| return 'Ready' | |
| try: | |
| data = STORE.control(session, owner(request), seq, x, z, slot) | |
| except (ValueError, FileNotFoundError) as exc: | |
| raise gr.Error(str(exc)) from exc | |
| return 'Running' if data['state'] == 'running' else 'Ready' | |
| def keyboard_control(event: gr.EventData, request: gr.Request): | |
| data = event._data | |
| if isinstance(data, dict): | |
| control(data.get('session_id', ''), data.get('seq', -1), data.get('x', 0.), | |
| data.get('z', 0.), False, data.get('slot', 0), request) | |
| def stop_session(session, request: gr.Request): | |
| if not session: | |
| return 'Ready' | |
| try: | |
| data = STORE.stop(session, owner(request)) | |
| except (ValueError, FileNotFoundError) as exc: | |
| raise gr.Error(str(exc)) from exc | |
| return 'Stopping…' if data['stream_active'] else 'Stopped' | |
| def session_status(session, request: gr.Request): | |
| if not session: | |
| return dict(state='ready', frames=0, stream_active=False) | |
| data = read(session, request) | |
| return {key: value for key, value in data.items() if key != 'owner'} | |
| def stream(session, request: gr.Request): | |
| browser = owner(request) | |
| initial = STORE.claim(session, browser) | |
| model = initial['model'] | |
| runtime = RUNTIMES[model] | |
| planner = SeedPath() if model == 'SEED' else HumanML3DPath() | |
| fps = 30 if model == 'SEED' else 20 | |
| completed = sequence = 0 | |
| started = time.perf_counter() | |
| deadline = started + 55. | |
| playback_origin = first_frame_seconds = None | |
| compute_seconds = 0. | |
| frame_seconds = deque(maxlen=60) | |
| text_metas = [] | |
| latest = initial | |
| last_meta = dict(slot=0, frame=-1, prompt=initial['slots'][0]) | |
| gpu_name = None | |
| acquired = False | |
| runtime_started = False | |
| def packet(status, frames=None, metas=None, rotations=None, error=None): | |
| nonlocal sequence | |
| frames, metas = frames or [], metas or [] | |
| info = runtime.status() if runtime_started else {} | |
| active_slot = last_meta['slot'] if completed else latest['slot'] | |
| result = dict(session_id=session, sequence=sequence, run_revision=1, | |
| model=model, representation='mesh' if model == 'SEED' else 'skeleton', | |
| start_frame=completed-len(frames), frames=frames, frame_meta=metas, | |
| control_points=[dict(meta) for meta in metas]+[dict(point) for point in planner.points], | |
| control_points_delta=True, control_replace_from_frame=completed-len(frames), | |
| spring_point=dict(planner.points[-1]) if planner.points else None, | |
| target_point=planner.target_point(active_slot), | |
| slots=list(latest['slots']), speeds=list(latest['speeds']), | |
| active_slot=active_slot, base_slot=latest['slot'], | |
| prompt=latest['slots'][active_slot], action=latest['slots'][active_slot], | |
| speed=latest['speeds'][active_slot], completed_frame=completed-1, | |
| input_frame=info.get('steps', 0)-1, fps=fps, generated_frames=completed, | |
| elapsed_seconds=time.perf_counter()-started, first_frame_seconds=first_frame_seconds, | |
| inference_fps=len(frame_seconds)/max(sum(frame_seconds), 1e-9), | |
| session_inference_fps=completed/max(compute_seconds, 1e-9), | |
| graph_enabled=bool(info.get('graph', {}).get('replays')), | |
| graph=info.get('graph'), text_bank=info.get('fixed_text'), | |
| control_version=last_meta.get('version', 0), path_mode='window', | |
| prompt_history=dict(enabled=True, max_history=30, | |
| resets=info.get('history_resets', 0), | |
| visible_history=info.get('visible_history', 0), | |
| cutoff=info.get('history_cutoff_absolute')), | |
| gpu=gpu_name, status=status, error=error) | |
| if rotations: | |
| result['mesh_pose'] = dict( | |
| rotations=base64.b64encode(np.asarray(rotations, dtype='<f4').tobytes()).decode('ascii'), | |
| origins=[frame[0] for frame in frames], frame_count=len(rotations)) | |
| STORE.patch(session, browser, frames=completed, runtime=info, error=error) | |
| sequence += 1 | |
| return result | |
| try: | |
| if not SpacesConfig.zero_gpu: | |
| LOCAL_GPU_LOCK.acquire() | |
| acquired = True | |
| gpu_name = torch.cuda.get_device_name() | |
| # A reused worker always starts with a fresh latent, planner, decoder, | |
| # text bank and captured graph; no mutable state crosses sessions. | |
| runtime.close() | |
| runtime.model.text_module.text_cache = {} | |
| runtime.add_prompt_features(BASE_BANK) | |
| runtime.add_prompt_features(load_prompt_features(STORE.path(session).with_suffix('.npz'))) | |
| yield packet('start') | |
| latest = STORE.read(session, browser) | |
| if latest['stop'] or time.perf_counter() >= deadline: | |
| STORE.patch(session, browser, state='paused', stop=True, stream_active=False) | |
| yield packet('paused') | |
| return | |
| tick = time.perf_counter() | |
| runtime.start(seed=0, history=60, use_graph=True, | |
| reset_history_on_prompt_change=True, cfg_scale=4.) | |
| runtime_started = True | |
| compute_seconds += time.perf_counter()-tick | |
| latest = STORE.read(session, browser) | |
| if not latest['stop']: | |
| latest = STORE.patch(session, browser, state='running') | |
| while time.perf_counter() < deadline: | |
| latest = STORE.read(session, browser) | |
| if latest['stop']: | |
| break | |
| tick = time.perf_counter() | |
| alive = time.time()-latest['keys_updated_at'] < 1.2 | |
| x, z = (latest['x'], latest['z']) if alive else (0., 0.) | |
| slot = latest['slot'] | |
| prompt = latest['slots'][slot] | |
| for meta in text_metas: | |
| meta.update(slot=slot, base_slot=slot, run=False, | |
| prompt=prompt, text_version=latest['version']) | |
| text_metas.append(dict(frame=runtime.status()['steps'], slot=slot, base_slot=slot, | |
| run=False, prompt=prompt, text_version=latest['version'])) | |
| rows, points = planner.plan(text_metas, x=x, z=z, shift=False, slot=slot, | |
| speeds=latest['speeds'], version=latest['version']) | |
| raw = runtime.step_replanned(prompt, rows, replan_text=True) | |
| frames, metas, rotations = [], [], [] | |
| if raw is not None: | |
| native, meta = planner.commit_first() | |
| text_metas.pop(0) | |
| if meta['frame'] != completed or not np.allclose(raw[:3], native, rtol=1e-5, atol=1e-6): | |
| raise RuntimeError('Motion and path frame alignment was lost.') | |
| last_meta = meta | |
| frames = [np.round(runtime.recover(raw)[0], 5).tolist()] | |
| metas = [dict(meta)] | |
| if model == 'SEED': | |
| rotations = [runtime.render_rotations(raw)] | |
| completed += 1 | |
| if first_frame_seconds is None: | |
| first_frame_seconds = time.perf_counter()-started | |
| playback_origin = time.perf_counter() | |
| seconds = time.perf_counter()-tick | |
| compute_seconds += seconds | |
| if raw is not None: | |
| frame_seconds.append(seconds) | |
| if frames or runtime.status()['steps'] % 6 == 0: | |
| yield packet('running' if completed else 'start', frames, metas, rotations) | |
| if playback_origin is not None: | |
| delay = playback_origin+completed/fps-time.perf_counter() | |
| if delay > 0: | |
| time.sleep(min(delay, .04)) | |
| STORE.patch(session, browser, state='paused', stop=True, stream_active=False, | |
| x=0., z=0., keys_updated_at=0.) | |
| yield packet('paused') | |
| except GeneratorExit: | |
| raise | |
| except Exception as exc: | |
| import traceback | |
| traceback.print_exc() | |
| message = f'{type(exc).__name__}: {exc}' | |
| STORE.patch(session, browser, state='error', stop=True, stream_active=False, error=message) | |
| yield packet('error', error=message) | |
| finally: | |
| try: | |
| runtime.close() | |
| finally: | |
| try: | |
| state = STORE.read(session, browser) | |
| STORE.patch(session, browser, state='error' if state['error'] else 'paused', | |
| stop=True, stream_active=False, x=0., z=0., keys_updated_at=0.) | |
| finally: | |
| if acquired: | |
| LOCAL_GPU_LOCK.release() | |
| with gr.Blocks(title='FloodDiffusion 2') as demo: | |
| gr.Markdown('# FloodDiffusion 2') | |
| session = gr.Textbox(value='', visible=False) | |
| body_loaded = gr.State(False) | |
| preparing = gr.State(False) | |
| model = gr.Radio(MODELS, value='SEED', show_label=False, container=False) | |
| with gr.Row(): | |
| start = gr.Button('Start', variant='primary', interactive=False) | |
| stop = gr.Button('Stop', interactive=False) | |
| prompts, speeds = [], [] | |
| with gr.Accordion('Prompts & speeds', open=False): | |
| for index, (text, value) in enumerate(DEFAULTS, 1): | |
| with gr.Row(): | |
| prompts.append(gr.Textbox(value=text, label=str(index), lines=1, max_length=500, scale=5)) | |
| speeds.append(gr.Number(value=value, label='Speed (m/s)', minimum=0, | |
| maximum=5, step=.05, scale=1, min_width=120)) | |
| status = gr.Markdown('Loading…') | |
| viewer = create_viewer(require_body_ready=True, action_slots=True) | |
| def sync_controls(session_id, loaded, preparing_now, request: gr.Request): | |
| info = session_status(session_id, request) | |
| phase = info['state'] | |
| busy = bool(preparing_now) or phase in ('starting', 'running', 'pausing') | |
| message = {'ready':'', 'paused':'Stopped', 'starting':'Preparing…', | |
| 'running':'', 'pausing':'Stopping…', 'error':info.get('error', '')}.get(phase, '') | |
| if not loaded: | |
| message = 'Loading…' | |
| elif preparing_now: | |
| message = 'Preparing…' | |
| return (gr.Button(interactive=bool(loaded) and not busy and phase in ('ready', 'paused', 'error')), | |
| gr.Button(interactive=busy and not preparing_now and phase != 'pausing'), | |
| gr.Radio(interactive=not busy), | |
| *(gr.Textbox(interactive=not busy) for _ in prompts), | |
| *(gr.Number(interactive=not busy) for _ in speeds), message) | |
| ui_outputs = [start, stop, model, *prompts, *speeds, status] | |
| loaded = viewer.body_ready(lambda: True, outputs=body_loaded, queue=False, api_name=False) | |
| ui_inputs = [session, body_loaded, preparing] | |
| loaded.then(sync_controls, ui_inputs, ui_outputs, queue=False, api_name=False) | |
| gr.Timer(.5).tick(sync_controls, ui_inputs, ui_outputs, | |
| queue=False, trigger_mode='always_last', api_name=False) | |
| def lock_controls(): | |
| return (gr.Button(interactive=False), gr.Button(interactive=False), | |
| gr.Radio(interactive=False), | |
| *(gr.Textbox(interactive=False) for _ in prompts), | |
| *(gr.Number(interactive=False) for _ in speeds), 'Preparing…', True) | |
| locking = start.click(lock_controls, outputs=[*ui_outputs, preparing], queue=False, api_name=False) | |
| beginning = locking.success(prepare, [session, model, *prompts, *speeds], session, | |
| api_name='prepare', concurrency_limit=GPU_CONCURRENCY, | |
| concurrency_id='gpu' if GPU_CONCURRENCY else None) | |
| prepared = beginning.success(lambda: False, outputs=preparing, queue=False, api_name=False) | |
| prepared.success(sync_controls, ui_inputs, ui_outputs, queue=False, api_name=False) | |
| running = prepared.success(stream, session, viewer, api_name='stream', | |
| concurrency_limit=GPU_CONCURRENCY, stream_every=.05, | |
| concurrency_id='gpu' if GPU_CONCURRENCY else None) | |
| running.then(sync_controls, ui_inputs, ui_outputs, queue=False, api_name=False) | |
| failed = beginning.failure(lambda: False, outputs=preparing, queue=False, api_name=False) | |
| failed.then(sync_controls, ui_inputs, ui_outputs, queue=False, api_name=False) | |
| stop.click(stop_session, session, status, queue=False, api_name='stop') | |
| viewer.control(keyboard_control, inputs=[], outputs=[], queue=False, trigger_mode='multiple', api_name=False) | |
| with gr.Row(visible=False): | |
| control_sequence = gr.Number(value=0) | |
| control_x, control_z = gr.Number(value=0), gr.Number(value=0) | |
| control_shift, control_slot = gr.Checkbox(value=False), gr.Number(value=0) | |
| control_status = gr.Textbox() | |
| gr.Button('Control API').click(control, [session, control_sequence, control_x, control_z, | |
| control_shift, control_slot], control_status, | |
| api_name='control', queue=False) | |
| inspect_status = gr.JSON() | |
| gr.Button('Session status').click(session_status, session, inspect_status, | |
| queue=False, api_name='session_status') | |
| demo.queue(max_size=32) | |
| if __name__ == '__main__': | |
| demo.launch(server_name='0.0.0.0', server_port=int(os.getenv('PORT', '7860')), | |
| show_error=True, footer_links=[]) | |