PhysiFormer / app.py
yimingc9's picture
Update Space intro copy
8590748
Raw
History Blame Contribute Delete
26.2 kB
from __future__ import annotations
import os
import re
import subprocess
import sys
import time
from pathlib import Path
import gradio as gr
import spaces
from huggingface_hub import hf_hub_download
ROOT = Path(__file__).resolve().parent
SRC_ROOT = ROOT / "src"
if str(SRC_ROOT) not in sys.path:
sys.path.insert(0, str(SRC_ROOT))
CKPT_REPO_ID = os.environ.get("PHYSFORMER_CKPT_REPO_ID", "yslan/physformer")
CKPT_FILENAME = os.environ.get("PHYSFORMER_CKPT_FILENAME", "checkpoint-best.pt")
CKPT_PATH = ROOT / "checkpoints" / "checkpoint-best.pt"
PREVIEW_VERSION = "v4"
RENDER_COLORS = [
(0.86, 0.24, 0.20, 1.0),
(0.20, 0.64, 0.42, 1.0),
(0.20, 0.44, 0.86, 1.0),
(0.92, 0.67, 0.22, 1.0),
(0.62, 0.32, 0.76, 1.0),
]
NAMED_COLORS = {
"cow": (0.00, 0.62, 0.66, 1.0),
"horse": (0.88, 0.30, 0.24, 1.0),
}
MESH_EDGE_COLOR = (0.05, 0.06, 0.07, 0.62)
RIGID_RENDER_ALPHA = 0.96
ELASTIC_RENDER_ALPHA = 0.38
def _natural_path_key(path: Path) -> tuple[object, ...]:
import re
parts: list[object] = []
for part in path.parts:
parts.extend(int(text) if text.isdigit() else text.lower() for text in re.split(r"(\d+)", part) if text)
return tuple(parts)
def _is_input_sample_dir(path: Path) -> bool:
return (
path.name.isdigit()
and (path / "metadata.json").is_file()
and (path / "meshes" / "combined_frame_000.obj").is_file()
and (path / "vertex_velocities" / "combined_frame_000.npy").is_file()
)
def _ood_sample_paths(group: str, fallback_count: int = 10) -> list[str]:
root = ROOT / "ood_examples" / group
samples = [
path.relative_to(ROOT).as_posix()
for path in sorted(root.iterdir(), key=_natural_path_key)
if path.is_dir() and _is_input_sample_dir(path)
] if root.is_dir() else []
if samples:
return samples
return [f"ood_examples/{group}/{i}" for i in range(fallback_count)]
OOD_SAMPLE_CHOICES = {
"2 objects": _ood_sample_paths("2obj_cow_horse"),
"3 objects": _ood_sample_paths("3obj_teapot_fish_bunny"),
}
SAMPLE_CHOICES = {
"OOD mixed materials": OOD_SAMPLE_CHOICES["2 objects"] + OOD_SAMPLE_CHOICES["3 objects"],
"In-distribution rigid": [
"indistri_examples/rigid/sample_000007",
"indistri_examples/rigid/sample_000114",
"indistri_examples/rigid/sample_000969",
"indistri_examples/rigid/sample_001151",
"indistri_examples/rigid/sample_001557",
"indistri_examples/rigid/sample_001874",
"indistri_examples/rigid/sample_002143",
],
"In-distribution elastic": [
"indistri_examples/elastic/sample_000047",
"indistri_examples/elastic/sample_000105",
"indistri_examples/elastic/sample_000121",
"indistri_examples/elastic/sample_000203",
],
}
DEFAULT_MATERIALS = {
"cow": "rigid",
"horse": "elastic",
"teapot": "rigid",
"fish": "elastic",
"bunny": "elastic",
}
def example_ids_for_setting(setting: str, object_count: str = "2 objects") -> list[str]:
if setting == "OOD mixed materials":
samples = OOD_SAMPLE_CHOICES.get(str(object_count), OOD_SAMPLE_CHOICES["2 objects"])
return [str(i) for i in range(len(samples))]
samples = SAMPLE_CHOICES.get(setting, SAMPLE_CHOICES["OOD mixed materials"])
return [str(i) for i in range(len(samples))]
def sample_path_for_example_id(setting: str, object_count: str, example_id: str) -> str:
if setting == "OOD mixed materials":
samples = OOD_SAMPLE_CHOICES.get(str(object_count), OOD_SAMPLE_CHOICES["2 objects"])
else:
samples = SAMPLE_CHOICES.get(setting, SAMPLE_CHOICES["OOD mixed materials"])
try:
idx = int(str(example_id).strip())
except ValueError as exc:
raise ValueError(f"Example must be an integer index, got {example_id!r}") from exc
if idx < 0 or idx >= len(samples):
raise ValueError(f"Example index {idx} is out of range for {setting!r}; valid range is 0..{len(samples) - 1}")
return samples[idx]
def _tail(text: str, max_chars: int = 18000) -> str:
if len(text) <= max_chars:
return text
return "[log truncated]\n" + text[-max_chars:]
def _timing_summary(log: str) -> str:
values = dict(re.findall(r"\[timing\]\s+([A-Za-z0-9_\[\]\.]+)=([0-9.]+)", log))
hardware = re.findall(r"\[hardware\]\s+(.+)", log)
attention = re.findall(r"\[attention\]\[rank=\d+\]\s+(.+)", log)
lines = []
if hardware:
lines.append("Hardware: " + hardware[-1])
if attention:
lines.append("Attention: " + attention[-1])
if "sample[0].gen[0].inference_model_generate_s" in values:
lines.append(f"Model inference: {values['sample[0].gen[0].inference_model_generate_s']} s")
fields = [
("checkpoint_load_s", "Checkpoint load"),
("model_setup_s", "Model setup"),
("sample[0].setup_s", "Input setup"),
("sample[0].gen[0].postprocess_save_npz_s", "Postprocess/save"),
("sample[0].gen[0].render_encode_s", "Render/encode"),
("engine_total_wall_s", "Engine total"),
("gradio_subprocess_wall_s", "Gradio subprocess wall"),
]
for key, label in fields:
if key in values:
lines.append(f"{label}: {values[key]} s")
if lines:
return "\n".join(lines)
return (
"No timing markers were found in the inference output.\n"
"The Space may still be running an older build, or the inference process exited before timing was emitted."
)
def _safe_preview_name(*parts: object) -> str:
text = "_".join([PREVIEW_VERSION, *(str(part) for part in parts)]).lower()
return "".join(ch if ch.isalnum() else "_" for ch in text).strip("_")
def _object_names_from_metadata(metadata_path: Path) -> list[str]:
import json
with metadata_path.open("r", encoding="utf-8") as f:
metadata = json.load(f)
out: list[str] = []
for obj in metadata.get("objects", []):
if isinstance(obj, dict):
name = obj.get("name") or obj.get("mesh_used") or obj.get("mesh_source") or ""
out.append(Path(str(name)).stem)
return out
def _color_for_object(index: int, object_name: str | None) -> tuple[float, float, float, float]:
name = str(object_name or "").lower()
for pattern, color in NAMED_COLORS.items():
if pattern in name:
return color
return RENDER_COLORS[int(index) % len(RENDER_COLORS)]
def _with_alpha(color: tuple[float, float, float, float], alpha: float) -> tuple[float, float, float, float]:
return (float(color[0]), float(color[1]), float(color[2]), float(alpha))
def _is_elastic_material(material: object) -> bool:
if isinstance(material, str):
return material.strip().lower() in {"elastic", "soft"}
if isinstance(material, dict):
kind = str(material.get("kind", "")).strip().lower()
if kind in {"elastic", "soft"}:
return True
if kind in {"rigid", "hard"}:
return False
for key in ("effective_softness", "softness"):
value = material.get(key)
if isinstance(value, (int, float)):
return float(value) >= 0.5
return False
def _default_material_values(sample: str) -> list[str]:
values = [DEFAULT_MATERIALS.get(obj, "elastic") for obj in objects_for_sample(sample)]
while len(values) < 3:
values.append("elastic")
return values[:3]
def _preview_object_alphas(
setting: str,
sample: str,
metadata_path: Path,
material_0: str = "",
material_1: str = "",
material_2: str = "",
) -> list[float]:
import json
if setting == "OOD mixed materials":
defaults = _default_material_values(sample)
materials = [
str(material_0 or defaults[0]),
str(material_1 or defaults[1]),
str(material_2 or defaults[2]),
]
return [ELASTIC_RENDER_ALPHA if _is_elastic_material(material) else RIGID_RENDER_ALPHA for material in materials]
with metadata_path.open("r", encoding="utf-8") as f:
metadata = json.load(f)
alphas: list[float] = []
for obj in metadata.get("objects", []):
material = obj.get("material") if isinstance(obj, dict) else None
alphas.append(ELASTIC_RENDER_ALPHA if _is_elastic_material(material) else RIGID_RENDER_ALPHA)
return alphas
def _shaded_facecolors(vertices, faces, base_color):
import numpy as np
light_direction = np.asarray([0.45, -0.65, 0.75], dtype=np.float32)
tris = vertices[faces]
normals = np.cross(tris[:, 1] - tris[:, 0], tris[:, 2] - tris[:, 0])
normals /= np.maximum(np.linalg.norm(normals, axis=1, keepdims=True), 1e-8)
light = light_direction / np.linalg.norm(light_direction)
intensity = 0.42 + 0.58 * np.clip(normals @ light, 0.0, 1.0)
base = np.asarray(base_color, dtype=np.float32)
facecolors = np.empty((faces.shape[0], 4), dtype=np.float32)
facecolors[:, :3] = np.clip(base[:3][None, :] * intensity[:, None] + 0.10 * (1.0 - intensity[:, None]), 0.0, 1.0)
facecolors[:, 3] = base[3]
return facecolors
def _faces_for_vertex_slice(faces, start: int, end: int):
import numpy as np
in_range = (faces >= int(start)) & (faces < int(end))
keep = np.all(in_range, axis=1)
return faces[keep] - int(start)
def _velocity_indices_for_object(speed, start: int, end: int, max_arrows: int):
import numpy as np
local = np.arange(int(start), int(end), dtype=np.int64)
active = local[speed[local] > 1e-9]
if active.size <= int(max_arrows):
return active
# Deterministic subsample across the object vertices so the preview does not become an arrow cloud.
positions = np.linspace(0, active.size - 1, int(max_arrows)).round().astype(np.int64)
return active[positions]
def _draw_unit_bounds(ax) -> None:
corners = [
(-1.0, -1.0, -1.0),
(-1.0, -1.0, 1.0),
(-1.0, 1.0, -1.0),
(-1.0, 1.0, 1.0),
(1.0, -1.0, -1.0),
(1.0, -1.0, 1.0),
(1.0, 1.0, -1.0),
(1.0, 1.0, 1.0),
]
edges = [
(0, 1), (0, 2), (0, 4), (3, 1), (3, 2), (3, 7),
(5, 1), (5, 4), (5, 7), (6, 2), (6, 4), (6, 7),
]
for start, end in edges:
xs = [corners[start][0], corners[end][0]]
ys = [corners[start][1], corners[end][1]]
zs = [corners[start][2], corners[end][2]]
ax.plot(xs, ys, zs, color=(0.18, 0.22, 0.28, 0.52), linewidth=0.9)
def render_initial_preview(
setting: str,
object_count: str,
example_id: str,
material_0: str = "",
material_1: str = "",
material_2: str = "",
) -> str | None:
sample = sample_path_for_example_id(setting, object_count, example_id)
sample_dir = ROOT / sample
obj_path = sample_dir / "meshes" / "combined_frame_000.obj"
vel_path = sample_dir / "vertex_velocities" / "combined_frame_000.npy"
metadata_path = sample_dir / "metadata.json"
if not obj_path.is_file() or not vel_path.is_file():
return None
out_dir = ROOT / ".inference_work" / "previews"
out_dir.mkdir(parents=True, exist_ok=True)
material_tag = "_".join(str(value or "default") for value in (material_0, material_1, material_2))
out_path = out_dir / f"{_safe_preview_name(setting, object_count, example_id, material_tag)}.png"
if out_path.is_file() and out_path.stat().st_mtime >= max(obj_path.stat().st_mtime, vel_path.stat().st_mtime):
return str(out_path)
os.environ.setdefault("MPLCONFIGDIR", str(ROOT / ".inference_work" / "matplotlib"))
import matplotlib
matplotlib.use("Agg")
import matplotlib.pyplot as plt
import numpy as np
from mpl_toolkits.mplot3d.art3d import Poly3DCollection
from physformer.data.multiobj_utils_multiobj import (
default_vertex_count_json_path,
load_mesh_vertex_counts,
scene_info_from_metadata,
)
from physformer.data.obj_io import load_obj_vertices_faces
vertices, faces = load_obj_vertices_faces(str(obj_path))
velocities = np.load(vel_path).astype(np.float32, copy=False)
if velocities.shape != vertices.shape:
raise ValueError(f"Velocity shape mismatch for {sample}: {velocities.shape} != {vertices.shape}")
scene = scene_info_from_metadata(
str(metadata_path),
vertex_counts=load_mesh_vertex_counts(default_vertex_count_json_path()),
max_num_objects=10,
)
object_names = _object_names_from_metadata(metadata_path)
object_alphas = _preview_object_alphas(setting, sample, metadata_path, material_0, material_1, material_2)
fig = plt.figure(figsize=(6.0, 5.0), dpi=150, facecolor="#f7f8fb")
ax = fig.add_subplot(111, projection="3d")
ax.set_facecolor("#f7f8fb")
for obj_idx, (start, end) in enumerate(scene.vertex_slices):
obj_vertices = vertices[int(start) : int(end)]
obj_faces = _faces_for_vertex_slice(faces, int(start), int(end))
if obj_vertices.size == 0 or obj_faces.size == 0:
continue
object_name = object_names[obj_idx] if obj_idx < len(object_names) else None
alpha = object_alphas[obj_idx] if obj_idx < len(object_alphas) else RIGID_RENDER_ALPHA
color = _with_alpha(_color_for_object(obj_idx, object_name), alpha)
poly = Poly3DCollection(
obj_vertices[obj_faces],
facecolors=_shaded_facecolors(obj_vertices, obj_faces, color),
edgecolors=MESH_EDGE_COLOR,
linewidths=0.28,
alpha=alpha,
antialiased=True,
)
ax.add_collection3d(poly)
speed = np.linalg.norm(velocities, axis=1)
arrow_cap = 20 if setting == "OOD mixed materials" else 10
arrow_indices = [
_velocity_indices_for_object(speed, int(start), int(end), arrow_cap)
for start, end in scene.vertex_slices
]
active = np.concatenate([idx for idx in arrow_indices if idx.size]) if any(idx.size for idx in arrow_indices) else np.empty((0,), dtype=np.int64)
bbox_diag = float(np.linalg.norm(np.asarray([2.0, 2.0, 2.0], dtype=np.float32)))
speed_ref = float(np.percentile(speed[active], 95)) if active.size else 0.0
scale = (0.20 * bbox_diag / speed_ref) if speed_ref > 0 and bbox_diag > 0 else 1.0
for obj_idx, active_obj in enumerate(arrow_indices):
if not active_obj.size:
continue
object_name = object_names[obj_idx] if obj_idx < len(object_names) else None
color = _color_for_object(obj_idx, object_name)
v = velocities[active_obj] * scale
ax.quiver(
vertices[active_obj, 0],
vertices[active_obj, 1],
vertices[active_obj, 2],
v[:, 0],
v[:, 1],
v[:, 2],
color=color[:3],
linewidth=0.85,
arrow_length_ratio=0.22,
normalize=False,
)
_draw_unit_bounds(ax)
ax.set_xlim(-1.0, 1.0)
ax.set_ylim(-1.0, 1.0)
ax.set_zlim(-1.0, 1.0)
ax.set_box_aspect([1, 1, 1])
ax.view_init(elev=24, azim=-56)
ax.set_title("Initial mesh per-vertex position and velocity", fontsize=10)
ax.set_xlabel("x")
ax.set_ylabel("y")
ax.set_zlabel("z")
ax.grid(True, linewidth=0.35, alpha=0.35)
for axis in (ax.xaxis, ax.yaxis, ax.zaxis):
axis.pane.set_facecolor((0.95, 0.96, 0.98, 0.72))
axis.pane.set_edgecolor((0.72, 0.75, 0.80, 0.50))
fig.tight_layout()
fig.savefig(out_path, bbox_inches="tight")
plt.close(fig)
return str(out_path)
def ensure_checkpoint() -> str:
if CKPT_PATH.is_file():
return f"Checkpoint found: {CKPT_PATH}"
CKPT_PATH.parent.mkdir(parents=True, exist_ok=True)
token = os.environ.get("HF_TOKEN") or None
downloaded = hf_hub_download(
repo_id=CKPT_REPO_ID,
filename=CKPT_FILENAME,
local_dir=str(CKPT_PATH.parent),
token=token,
)
downloaded_path = Path(downloaded)
if downloaded_path.resolve() != CKPT_PATH.resolve():
downloaded_path.replace(CKPT_PATH)
return f"Downloaded checkpoint from {CKPT_REPO_ID}/{CKPT_FILENAME}"
def objects_for_sample(sample: str) -> list[str]:
sample = str(sample)
if "/2obj_cow_horse/" in sample:
return ["cow", "horse"]
if "/3obj_teapot_fish_bunny/" in sample:
return ["teapot", "fish", "bunny"]
return []
def material_controls_for_sample(setting: str, object_count: str, example_id: str) -> tuple[dict, dict, dict]:
sample = sample_path_for_example_id(setting, object_count, example_id)
objects = objects_for_sample(sample) if setting == "OOD mixed materials" else []
updates: list[dict] = []
for idx in range(3):
if idx < len(objects):
obj = objects[idx]
updates.append(
gr.update(
label=f"{obj} material",
value=DEFAULT_MATERIALS[obj],
visible=True,
)
)
else:
updates.append(gr.update(visible=False))
return tuple(updates) # type: ignore[return-value]
def material_controls_and_preview(setting: str, object_count: str, example_id: str) -> tuple[dict, dict, dict, str | None]:
sample = sample_path_for_example_id(setting, object_count, example_id)
material_0, material_1, material_2 = _default_material_values(sample)
return (
*material_controls_for_sample(setting, object_count, example_id),
render_initial_preview(setting, object_count, example_id, material_0, material_1, material_2),
)
def preview_for_materials(
setting: str,
object_count: str,
example_id: str,
material_0: str,
material_1: str,
material_2: str,
) -> str | None:
return render_initial_preview(setting, object_count, example_id, material_0, material_1, material_2)
def update_setting_controls(setting: str) -> tuple[dict, dict, dict, dict, dict, str | None]:
if setting == "OOD mixed materials":
object_count = "2 objects"
choices = example_ids_for_setting(setting, object_count)
example_id = choices[0]
return (
gr.update(visible=True, value=object_count),
gr.update(choices=choices, value=example_id),
*material_controls_and_preview(setting, object_count, example_id),
)
choices = example_ids_for_setting(setting)
example_id = choices[0]
return (
gr.update(visible=False, value="2 objects"),
gr.update(choices=choices, value=example_id),
*material_controls_and_preview(setting, "2 objects", example_id),
)
def update_ood_object_count_controls(setting: str, object_count: str) -> tuple[dict, dict, dict, dict, str | None]:
choices = example_ids_for_setting(setting, object_count)
example_id = choices[0]
return (gr.update(choices=choices, value=example_id), *material_controls_and_preview(setting, object_count, example_id))
def _ood_material_args(sample: str, material_0: str, material_1: str, material_2: str) -> list[str]:
objects = objects_for_sample(sample)
materials = [material_0, material_1, material_2]
args: list[str] = []
for obj, material in zip(objects, materials):
material = str(material).strip().lower()
if material not in {"elastic", "rigid"}:
raise ValueError(f"Invalid material for {obj}: {material!r}")
args.extend([f"--{material}", obj])
return args
def _command_for_example(
setting: str,
object_count: str,
example_id: str,
sampling_steps: int,
material_0: str,
material_1: str,
material_2: str,
) -> list[str]:
sample = sample_path_for_example_id(setting, object_count, example_id)
common = [
sys.executable,
"run_official_demo_inference.py",
"--demo-root",
sample,
"--include",
"all",
"--generations",
"1",
"--num-sampling-steps",
str(int(sampling_steps)),
"--checkpoint",
str(CKPT_PATH),
"--device",
"cuda",
"--amp",
os.environ.get("PHYSFORMER_AMP", "bf16"),
"--overwrite",
"--save-mp4",
"--verbose",
"--attention-debug",
]
if setting == "OOD mixed materials":
return common + _ood_material_args(sample, material_0, material_1, material_2)
if setting == "In-distribution rigid":
return common + ["--rigid", "all"]
if setting == "In-distribution elastic":
return common + ["--elastic", "all"]
raise ValueError(f"Unknown setting: {setting}")
def _latest_mp4_since(start_time: float) -> Path | None:
candidates = sorted(
[
path
for path in ROOT.glob("**/inference.mp4")
if ".inference_work" not in path.parts and path.stat().st_mtime >= start_time - 1.0
],
key=lambda path: path.stat().st_mtime,
reverse=True,
)
return candidates[0] if candidates else None
DEMO_CSS = """
#generated-rollout {
width: min(100%, 840px) !important;
max-width: 840px !important;
}
#generated-rollout video {
width: 100% !important;
max-height: 480px !important;
object-fit: contain !important;
}
"""
@spaces.GPU(duration=120)
def run_inference(
setting: str,
object_count: str,
example_id: str,
sampling_steps: int,
material_0: str,
material_1: str,
material_2: str,
setup_log: str,
) -> tuple[str | None, str, str]:
if not CKPT_PATH.is_file():
log = setup_log + "\nCheckpoint is missing; click Run again after the download finishes."
return None, "Checkpoint missing.", log
start_time = time.time()
subprocess_t0 = time.perf_counter()
cmd = _command_for_example(setting, str(object_count), str(example_id), int(sampling_steps), material_0, material_1, material_2)
env = os.environ.copy()
env.setdefault("PYTHONUNBUFFERED", "1")
env.setdefault("MPLCONFIGDIR", str(ROOT / ".inference_work" / "matplotlib"))
proc = subprocess.run(
cmd,
cwd=ROOT,
env=env,
text=True,
stdout=subprocess.PIPE,
stderr=subprocess.STDOUT,
check=False,
timeout=900,
)
subprocess_s = time.perf_counter() - subprocess_t0
log = setup_log + "\n\n$ " + " ".join(cmd) + "\n" + proc.stdout + f"\n[timing] gradio_subprocess_wall_s={subprocess_s:.3f}"
mp4 = _latest_mp4_since(start_time)
if proc.returncode != 0:
fail_log = log + f"\nInference failed with exit code {proc.returncode}."
return None, _timing_summary(fail_log), _tail(fail_log)
if mp4 is None:
missing_log = log + "\nInference finished, but no inference.mp4 was found."
return None, _timing_summary(missing_log), _tail(missing_log)
final_log = log + f"\nGenerated video: {mp4.relative_to(ROOT)}"
return str(mp4), _timing_summary(final_log), _tail(final_log)
with gr.Blocks(title="PhysFormer", css=DEMO_CSS) as demo:
gr.Markdown(
"""
# PhysiFormer Minimal ZeroGPU Demo
Select the example, material conditions, and denoising step numbers to run PhysiFormer inference.
This demo runs on Hugging Face ZeroGPU, dynamically allocating a 48GB NVIDIA RTX Pro 6000 Blackwell GPU for each generation.
"""
)
with gr.Row():
setting = gr.Dropdown(
choices=["OOD mixed materials", "In-distribution rigid", "In-distribution elastic"],
value="OOD mixed materials",
label="Setting",
)
object_count = gr.Dropdown(
choices=["2 objects", "3 objects"],
value="2 objects",
label="Object Count",
visible=True,
)
example_id = gr.Dropdown(
choices=example_ids_for_setting("OOD mixed materials"),
value="0",
label="Example",
)
sampling_steps = gr.Slider(5, 50, value=10, step=1, label="Denoising steps")
with gr.Row():
material_0 = gr.Dropdown(
choices=["elastic", "rigid"],
value="rigid",
label="cow material",
visible=True,
)
material_1 = gr.Dropdown(
choices=["elastic", "rigid"],
value="elastic",
label="horse material",
visible=True,
)
material_2 = gr.Dropdown(
choices=["elastic", "rigid"],
value="elastic",
label="material",
visible=False,
)
preview = gr.Image(
value=render_initial_preview("OOD mixed materials", "2 objects", "0", "rigid", "elastic", "elastic"),
label="Initial mesh per-vertex position and velocity",
type="filepath",
height=420,
)
run_button = gr.Button("Generate", variant="primary")
video = gr.Video(label="Generated rollout", height=480, width=840, elem_id="generated-rollout")
timing = gr.Textbox(label="Timing summary", lines=8, value="Run a rollout to see timing.")
log = gr.Textbox(label="Log", lines=18)
setting.change(
update_setting_controls,
inputs=setting,
outputs=[object_count, example_id, material_0, material_1, material_2, preview],
)
object_count.change(
update_ood_object_count_controls,
inputs=[setting, object_count],
outputs=[example_id, material_0, material_1, material_2, preview],
)
example_id.change(
material_controls_and_preview,
inputs=[setting, object_count, example_id],
outputs=[material_0, material_1, material_2, preview],
)
for material_control in (material_0, material_1, material_2):
material_control.change(
preview_for_materials,
inputs=[setting, object_count, example_id, material_0, material_1, material_2],
outputs=preview,
)
run_button.click(ensure_checkpoint, outputs=log).then(
run_inference,
inputs=[setting, object_count, example_id, sampling_steps, material_0, material_1, material_2, log],
outputs=[video, timing, log],
)
if __name__ == "__main__":
demo.queue(default_concurrency_limit=1, max_size=8).launch()