Anthony Liang commited on
Commit
a2a78e8
·
1 Parent(s): 61da1da

Add LFS tracking metadata

Browse files
__pycache__/app.cpython-310.pyc DELETED
Binary file (24.6 kB)
 
__pycache__/app.cpython-311.pyc DELETED
Binary file (63.4 kB)
 
app_internal.py ADDED
@@ -0,0 +1,1495 @@
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
1
+ #!/usr/bin/env python3
2
+ """
3
+ Gradio app for RBM (Reward Foundation Model) inference visualization.
4
+ Supports single video (progress/success) and dual video (preference/progress) predictions.
5
+ Uses eval server for inference instead of loading models locally.
6
+ """
7
+
8
+ import os
9
+ import tempfile
10
+ from pathlib import Path
11
+ from typing import Optional, Tuple
12
+ import logging
13
+
14
+ import gradio as gr
15
+
16
+ try:
17
+ import spaces # Required for ZeroGPU on Hugging Face Spaces
18
+ except ImportError:
19
+ spaces = None # Not available when running locally
20
+ import matplotlib
21
+
22
+ matplotlib.use("Agg") # Use non-interactive backend
23
+ import matplotlib.pyplot as plt
24
+ import numpy as np
25
+ import requests
26
+ from typing import Any, List, Optional, Tuple
27
+
28
+ from dataset_types import Trajectory, ProgressSample, PreferenceSample
29
+ from eval_utils import build_payload, post_batch_npy
30
+ from eval_viz_utils import create_combined_progress_success_plot, extract_frames
31
+ from datasets import load_dataset as load_dataset_hf, get_dataset_config_names
32
+
33
+ logger = logging.getLogger(__name__)
34
+
35
+ # Predefined dataset names (same as visualizer)
36
+ PREDEFINED_DATASETS = [
37
+ "abraranwar/agibotworld_alpha_rfm",
38
+ "abraranwar/libero_rfm",
39
+ "abraranwar/usc_koch_rewind_rfm",
40
+ "aliangdw/metaworld",
41
+ "anqil/rh20t_rfm",
42
+ "anqil/rh20t_subset_rfm",
43
+ "jesbu1/auto_eval_rfm",
44
+ "jesbu1/egodex_rfm",
45
+ "jesbu1/epic_rfm",
46
+ "jesbu1/fino_net_rfm",
47
+ "jesbu1/failsafe_rfm",
48
+ "jesbu1/hand_paired_rfm",
49
+ "jesbu1/galaxea_rfm",
50
+ "jesbu1/h2r_rfm",
51
+ "jesbu1/humanoid_everyday_rfm",
52
+ "jesbu1/molmoact_rfm",
53
+ "jesbu1/motif_rfm",
54
+ "jesbu1/oxe_rfm",
55
+ "jesbu1/oxe_rfm_eval",
56
+ "jesbu1/ph2d_rfm",
57
+ "jesbu1/racer_rfm",
58
+ "jesbu1/roboarena_0825_rfm",
59
+ "jesbu1/soar_rfm",
60
+ "ykorkmaz/libero_failure_rfm",
61
+ "aliangdw/usc_xarm_policy_ranking",
62
+ "aliangdw/usc_franka_policy_ranking",
63
+ "aliangdw/utd_so101_policy_ranking",
64
+ "aliangdw/utd_so101_human",
65
+ "jesbu1/utd_so101_clean_policy_ranking_top",
66
+ "jesbu1/utd_so101_clean_policy_ranking_wrist",
67
+ "jesbu1/mit_franka_p-rank_rfm",
68
+ "jesbu1/usc_koch_p_ranking_rfm",
69
+ ]
70
+
71
+ # Global server state
72
+ _server_state = {
73
+ "server_url": None,
74
+ "base_url": "https://robometer.a.pinggy.link", # Default: Pinggy tunnel or use http://HOST for port scan
75
+ }
76
+
77
+
78
+ def discover_available_models(
79
+ base_url: str = "http://40.119.56.66", port_range: tuple = (8000, 8010)
80
+ ) -> List[Tuple[str, str]]:
81
+ """Discover available models by pinging the base URL as-is, or ports in the specified range.
82
+
83
+ If base_url is a full URL (e.g. https://robometer.a.pinggy.link), it is tried as-is first.
84
+ Otherwise we try base_url:8000, base_url:8001, ... up to end_port.
85
+
86
+ Returns:
87
+ List of tuples: [(server_url, model_name), ...]
88
+ """
89
+ base_url = base_url.strip().rstrip("/")
90
+ if not base_url:
91
+ return []
92
+
93
+ available_models = []
94
+ # Try base_url as-is first (for Pinggy/tunnel URLs like https://robometer.a.pinggy.link)
95
+ try:
96
+ health_url = f"{base_url}/health"
97
+ health_response = requests.get(health_url, timeout=5.0)
98
+ if health_response.status_code == 200:
99
+ try:
100
+ model_info_url = f"{base_url}/model_info"
101
+ model_info_response = requests.get(model_info_url, timeout=5.0)
102
+ if model_info_response.status_code == 200:
103
+ model_info_data = model_info_response.json()
104
+ model_name = model_info_data.get("model_path", base_url)
105
+ available_models.append((base_url, model_name))
106
+ else:
107
+ available_models.append((base_url, base_url))
108
+ except Exception:
109
+ available_models.append((base_url, base_url))
110
+ return available_models
111
+ except requests.exceptions.RequestException:
112
+ pass
113
+
114
+ # Port scan: base_url is a host (e.g. http://40.119.56.66), try ports in range
115
+ start_port, end_port = port_range
116
+ for port in range(start_port, end_port + 1):
117
+ server_url = f"{base_url}:{port}"
118
+ try:
119
+ health_url = f"{server_url}/health"
120
+ health_response = requests.get(health_url, timeout=2.0)
121
+ if health_response.status_code == 200:
122
+ try:
123
+ model_info_url = f"{server_url}/model_info"
124
+ model_info_response = requests.get(model_info_url, timeout=2.0)
125
+ if model_info_response.status_code == 200:
126
+ model_info_data = model_info_response.json()
127
+ model_name = model_info_data.get("model_path", f"Model on port {port}")
128
+ available_models.append((server_url, model_name))
129
+ else:
130
+ available_models.append((server_url, f"Model on port {port}"))
131
+ except Exception:
132
+ available_models.append((server_url, f"Model on port {port}"))
133
+ except requests.exceptions.RequestException:
134
+ continue
135
+
136
+ return available_models
137
+
138
+
139
+ def get_model_info_for_url(server_url: str) -> Optional[str]:
140
+ """Get formatted model info for a given server URL."""
141
+ if not server_url:
142
+ return None
143
+
144
+ try:
145
+ model_info_url = server_url.rstrip("/") + "/model_info"
146
+ model_info_response = requests.get(model_info_url, timeout=5.0)
147
+ if model_info_response.status_code == 200:
148
+ model_info_data = model_info_response.json()
149
+ return format_model_info(model_info_data)
150
+ except Exception as e:
151
+ logger.warning(f"Could not fetch model info: {e}")
152
+ return None
153
+
154
+
155
+ def check_server_health(server_url: str) -> Tuple[str, Optional[dict], Optional[str]]:
156
+ """Check server health and get model info."""
157
+ if not server_url:
158
+ return "Please provide a server URL.", None, None
159
+
160
+ try:
161
+ url = server_url.rstrip("/") + "/health"
162
+ response = requests.get(url, timeout=5.0)
163
+ response.raise_for_status()
164
+ health_data = response.json()
165
+
166
+ # Also try to get GPU status for more info
167
+ try:
168
+ status_url = server_url.rstrip("/") + "/gpu_status"
169
+ status_response = requests.get(status_url, timeout=5.0)
170
+ if status_response.status_code == 200:
171
+ status_data = status_response.json()
172
+ health_data.update(status_data)
173
+ except:
174
+ pass
175
+
176
+ # Try to get model info
177
+ model_info_text = get_model_info_for_url(server_url)
178
+
179
+ _server_state["server_url"] = server_url
180
+ return (
181
+ f"Server connected: {health_data.get('available_gpus', 0)}/{health_data.get('total_gpus', 0)} GPUs available",
182
+ health_data,
183
+ model_info_text,
184
+ )
185
+ except requests.exceptions.RequestException as e:
186
+ return f"Error connecting to server: {str(e)}", None, None
187
+
188
+
189
+ def format_model_info(model_info: dict) -> str:
190
+ """Format model info and experiment config as markdown."""
191
+ lines = ["## Model Information\n"]
192
+
193
+ # Model path
194
+ model_path = model_info.get("model_path", "Unknown")
195
+ lines.append(f"**Model Path:** `{model_path}`\n")
196
+
197
+ # Number of GPUs
198
+ num_gpus = model_info.get("num_gpus", "Unknown")
199
+ lines.append(f"**Number of GPUs:** {num_gpus}\n")
200
+
201
+ # Model architecture
202
+ model_arch = model_info.get("model_architecture", {})
203
+ if model_arch and "error" not in model_arch:
204
+ lines.append("\n## Model Architecture\n")
205
+
206
+ model_class = model_arch.get("model_class", "Unknown")
207
+ model_module = model_arch.get("model_module", "Unknown")
208
+ lines.append(f"- **Model Class:** `{model_class}`\n")
209
+ lines.append(f"- **Module:** `{model_module}`\n")
210
+
211
+ # Parameter counts
212
+ total_params = model_arch.get("total_parameters")
213
+ trainable_params = model_arch.get("trainable_parameters")
214
+ frozen_params = model_arch.get("frozen_parameters")
215
+ trainable_pct = model_arch.get("trainable_percentage")
216
+
217
+ if total_params is not None:
218
+ lines.append(f"\n### Parameter Statistics\n")
219
+ lines.append(f"- **Total Parameters:** {total_params:,}\n")
220
+ if trainable_params is not None:
221
+ lines.append(f"- **Trainable Parameters:** {trainable_params:,}\n")
222
+ if frozen_params is not None:
223
+ lines.append(f"- **Frozen Parameters:** {frozen_params:,}\n")
224
+ if trainable_pct is not None:
225
+ lines.append(f"- **Trainable Percentage:** {trainable_pct:.2f}%\n")
226
+
227
+ # Architecture summary
228
+ arch_summary = model_arch.get("architecture_summary", [])
229
+ if arch_summary:
230
+ lines.append(f"\n### Architecture Summary (Top-Level Modules)\n")
231
+ for module_info in arch_summary[:10]: # Show first 10 modules
232
+ name = module_info.get("name", "Unknown")
233
+ module_type = module_info.get("type", "Unknown")
234
+ params = module_info.get("parameters", 0)
235
+ lines.append(f"- **{name}** (`{module_type}`): {params:,} parameters\n")
236
+
237
+ # Experiment config
238
+ exp_config = model_info.get("experiment_config", {})
239
+ if exp_config:
240
+ lines.append("\n## Experiment Configuration\n")
241
+
242
+ # Model config
243
+ model_cfg = exp_config.get("model", {})
244
+ if model_cfg:
245
+ lines.append("### Model Configuration\n")
246
+ lines.append(f"- **Base Model:** `{model_cfg.get('base_model_id', 'N/A')}`\n")
247
+ lines.append(f"- **Model Type:** `{model_cfg.get('model_type', 'N/A')}`\n")
248
+ lines.append(f"- **Train Progress Head:** {model_cfg.get('train_progress_head', False)}\n")
249
+ lines.append(f"- **Train Preference Head:** {model_cfg.get('train_preference_head', False)}\n")
250
+ lines.append(f"- **Train Success Head:** {model_cfg.get('train_success_head', False)}\n")
251
+ lines.append(f"- **Use PEFT:** {model_cfg.get('use_peft', False)}\n")
252
+ lines.append(f"- **Use Unsloth:** {model_cfg.get('use_unsloth', False)}\n")
253
+
254
+ # Data config
255
+ data_cfg = exp_config.get("data", {})
256
+ if data_cfg:
257
+ lines.append("\n### Data Configuration\n")
258
+ lines.append(f"- **Max Frames:** {data_cfg.get('max_frames', 'N/A')}\n")
259
+ lines.append(
260
+ f"- **Resized Dimensions:** {data_cfg.get('resized_height', 'N/A')}x{data_cfg.get('resized_width', 'N/A')}\n"
261
+ )
262
+ train_datasets = data_cfg.get("train_datasets", [])
263
+ if train_datasets:
264
+ lines.append(f"- **Train Datasets:** {', '.join(train_datasets)}\n")
265
+ eval_datasets = data_cfg.get("eval_datasets", [])
266
+ if eval_datasets:
267
+ lines.append(f"- **Eval Datasets:** {', '.join(eval_datasets)}\n")
268
+
269
+ # Training config
270
+ training_cfg = exp_config.get("training", {})
271
+ if training_cfg:
272
+ lines.append("\n### Training Configuration\n")
273
+ lines.append(f"- **Learning Rate:** {training_cfg.get('learning_rate', 'N/A')}\n")
274
+ lines.append(f"- **Batch Size:** {training_cfg.get('per_device_train_batch_size', 'N/A')}\n")
275
+ lines.append(
276
+ f"- **Gradient Accumulation Steps:** {training_cfg.get('gradient_accumulation_steps', 'N/A')}\n"
277
+ )
278
+ lines.append(f"- **Max Steps:** {training_cfg.get('max_steps', 'N/A')}\n")
279
+
280
+ return "".join(lines)
281
+
282
+
283
+ def load_rbm_dataset(dataset_name, config_name):
284
+ """Load an RBM-format dataset from HuggingFace Hub."""
285
+ try:
286
+ if not dataset_name or not config_name:
287
+ return None, "Please provide both dataset name and configuration"
288
+
289
+ dataset = load_dataset_hf(dataset_name, name=config_name, split="train")
290
+
291
+ if len(dataset) == 0:
292
+ return None, f"Dataset {dataset_name}/{config_name} is empty"
293
+
294
+ return dataset, f"Loaded {len(dataset)} trajectories from {dataset_name}/{config_name}"
295
+ except Exception as e:
296
+ error_msg = str(e)
297
+ if "not found" in error_msg.lower():
298
+ return None, f"Dataset or configuration not found: {dataset_name}/{config_name}"
299
+ elif "authentication" in error_msg.lower():
300
+ return None, f"Authentication required for {dataset_name}"
301
+ else:
302
+ return None, f"Error loading dataset: {error_msg}"
303
+
304
+
305
+ def get_available_configs(dataset_name):
306
+ """Get available configurations for a dataset."""
307
+ try:
308
+ configs = get_dataset_config_names(dataset_name)
309
+ return configs
310
+ except Exception as e:
311
+ logger.warning(f"Error getting configs for {dataset_name}: {e}")
312
+ return []
313
+
314
+
315
+ def get_trajectory_video_path(dataset, index, dataset_name):
316
+ """Get video path and metadata from a trajectory in the dataset."""
317
+ try:
318
+ item = dataset[int(index)]
319
+ frames_data = item["frames"]
320
+
321
+ if isinstance(frames_data, str):
322
+ # Construct HuggingFace Hub URL
323
+ if dataset_name:
324
+ video_path = f"https://huggingface.co/datasets/{dataset_name}/resolve/main/{frames_data}"
325
+ else:
326
+ video_path = f"https://huggingface.co/datasets/rewardfm/rbm-1m/resolve/main/{frames_data}"
327
+
328
+ task = item.get("task", "Complete the task")
329
+ quality_label = item.get("quality_label", None)
330
+ partial_success = item.get("partial_success", None)
331
+
332
+ return video_path, task, quality_label, partial_success
333
+ else:
334
+ return None, None, None, None
335
+ except Exception as e:
336
+ logger.error(f"Error getting trajectory video path: {e}")
337
+ return None, None, None, None
338
+
339
+
340
+ def process_single_video(
341
+ video_path: str,
342
+ task_text: str = "Complete the task",
343
+ server_url: str = "",
344
+ fps: float = 1.0,
345
+ use_frame_steps: bool = False,
346
+ ) -> Tuple[Optional[str], Optional[str]]:
347
+ """Process single video for progress and success predictions using eval server."""
348
+ # Get server URL from state if not provided
349
+ if not server_url:
350
+ server_url = _server_state.get("server_url")
351
+
352
+ if not server_url:
353
+ return None, "Please select a model from the dropdown above and ensure it's connected."
354
+
355
+ if video_path is None:
356
+ return None, "Please provide a video."
357
+
358
+ try:
359
+ frames_array = extract_frames(video_path, fps=fps)
360
+ if frames_array is None or frames_array.size == 0:
361
+ return None, "Could not extract frames from video."
362
+
363
+ # Convert frames to (T, H, W, C) numpy array with uint8 values
364
+ if frames_array.dtype != np.uint8:
365
+ frames_array = np.clip(frames_array, 0, 255).astype(np.uint8)
366
+
367
+ num_frames = frames_array.shape[0]
368
+ frames_shape = frames_array.shape # (T, H, W, C)
369
+
370
+ # Create target progress (placeholder - would be None in real use)
371
+ target_progress = np.linspace(0.0, 1.0, num=num_frames).tolist()
372
+ success_label = [1.0 if prog > 0.5 else 0.0 for prog in target_progress]
373
+
374
+ # predict_last_frame_mask: server collator requires a list (1.0 per frame = no masking for inference)
375
+ predict_last_frame_mask = [1.0] * num_frames
376
+
377
+ # Create Trajectory
378
+ trajectory = Trajectory(
379
+ task=task_text,
380
+ frames=frames_array,
381
+ frames_shape=frames_shape,
382
+ target_progress=target_progress,
383
+ success_label=success_label,
384
+ predict_last_frame_mask=predict_last_frame_mask,
385
+ metadata={"source": "gradio_app"},
386
+ )
387
+
388
+ # Create ProgressSample
389
+ progress_sample = ProgressSample(
390
+ trajectory=trajectory,
391
+ data_gen_strategy="demo",
392
+ )
393
+
394
+ # Build payload and send to server
395
+ files, sample_data = build_payload([progress_sample])
396
+ # Add use_frame_steps flag as extra form data
397
+ extra_data = {"use_frame_steps": use_frame_steps} if use_frame_steps else None
398
+ response = post_batch_npy(server_url, files, sample_data, timeout_s=120.0, extra_form_data=extra_data)
399
+
400
+ # Process response
401
+ outputs_progress = response.get("outputs_progress", {})
402
+ progress_pred = outputs_progress.get("progress_pred", [])
403
+ outputs_success = response.get("outputs_success", {})
404
+ success_probs = outputs_success.get("success_probs", []) if outputs_success else None
405
+
406
+ # Extract progress predictions
407
+ if progress_pred and len(progress_pred) > 0:
408
+ progress_array = np.array(progress_pred[0]) # First sample
409
+ else:
410
+ progress_array = np.array([])
411
+
412
+ # Extract success predictions if available
413
+ success_array = None
414
+ if success_probs and len(success_probs) > 0:
415
+ success_array = np.array(success_probs[0])
416
+
417
+ # Convert success_array to binary if available
418
+ success_binary = None
419
+ if success_array is not None:
420
+ success_binary = (success_array > 0.5).astype(float)
421
+
422
+ # Create combined plot using shared helper function
423
+ fig = create_combined_progress_success_plot(
424
+ progress_pred=progress_array if len(progress_array) > 0 else np.array([0.0]),
425
+ num_frames=num_frames,
426
+ success_binary=success_binary,
427
+ success_probs=success_array,
428
+ success_labels=None, # No ground truth labels available
429
+ is_discrete_mode=False,
430
+ title=f"Progress & Success - {task_text}",
431
+ )
432
+
433
+ # Save to temporary file
434
+ tmp_file = tempfile.NamedTemporaryFile(delete=False, suffix=".png")
435
+ fig.savefig(tmp_file.name, dpi=150, bbox_inches="tight")
436
+ plt.close(fig)
437
+ progress_plot = tmp_file.name
438
+
439
+ info_text = f"**Frames processed:** {num_frames}\n"
440
+ if len(progress_array) > 0:
441
+ info_text += f"**Final progress:** {progress_array[-1]:.3f}\n"
442
+ if success_array is not None and len(success_array) > 0:
443
+ info_text += f"**Final success probability:** {success_array[-1]:.3f}\n"
444
+
445
+ # Return combined plot (which includes success if available)
446
+ return progress_plot, info_text
447
+
448
+ except Exception as e:
449
+ return None, f"Error processing video: {str(e)}"
450
+
451
+
452
+ def process_two_videos(
453
+ video_a_path: str,
454
+ video_b_path: str,
455
+ task_text: str = "Complete the task",
456
+ prediction_type: str = "preference",
457
+ server_url: str = "",
458
+ fps: float = 1.0,
459
+ ) -> Tuple[Optional[str], Optional[str], Optional[str]]:
460
+ """Process two videos for preference or progress prediction using eval server."""
461
+ # Get server URL from state if not provided
462
+ if not server_url:
463
+ server_url = _server_state.get("server_url")
464
+
465
+ if not server_url:
466
+ return "Please select a model from the dropdown above and ensure it's connected.", None, None
467
+
468
+ if video_a_path is None or video_b_path is None:
469
+ return "Please provide both videos.", None, None
470
+
471
+ try:
472
+ frames_array_a = extract_frames(video_a_path, fps=fps)
473
+ frames_array_b = extract_frames(video_b_path, fps=fps)
474
+
475
+ if frames_array_a is None or frames_array_a.size == 0:
476
+ return "Could not extract frames from video A.", None, None
477
+ if frames_array_b is None or frames_array_b.size == 0:
478
+ return "Could not extract frames from video B.", None, None
479
+
480
+ # Convert frames to uint8
481
+ if frames_array_a.dtype != np.uint8:
482
+ frames_array_a = np.clip(frames_array_a, 0, 255).astype(np.uint8)
483
+ if frames_array_b.dtype != np.uint8:
484
+ frames_array_b = np.clip(frames_array_b, 0, 255).astype(np.uint8)
485
+
486
+ num_frames_a = frames_array_a.shape[0]
487
+ num_frames_b = frames_array_b.shape[0]
488
+ frames_shape_a = frames_array_a.shape
489
+ frames_shape_b = frames_array_b.shape
490
+
491
+ # Create target progress for both trajectories
492
+ target_progress_a = np.linspace(0.0, 1.0, num=num_frames_a).tolist()
493
+ target_progress_b = np.linspace(0.0, 1.0, num=num_frames_b).tolist()
494
+ success_label_a = [1.0 if prog > 0.5 else 0.0 for prog in target_progress_a]
495
+ success_label_b = [1.0 if prog > 0.5 else 0.0 for prog in target_progress_b]
496
+
497
+ # predict_last_frame_mask: server collator requires a list per trajectory (1.0 = no masking)
498
+ mask_a = [1.0] * num_frames_a
499
+ mask_b = [1.0] * num_frames_b
500
+
501
+ # Create trajectories
502
+ trajectory_a = Trajectory(
503
+ task=task_text,
504
+ frames=frames_array_a,
505
+ frames_shape=frames_shape_a,
506
+ target_progress=target_progress_a,
507
+ success_label=success_label_a,
508
+ predict_last_frame_mask=mask_a,
509
+ metadata={"source": "gradio_app", "trajectory": "A"},
510
+ )
511
+
512
+ trajectory_b = Trajectory(
513
+ task=task_text,
514
+ frames=frames_array_b,
515
+ frames_shape=frames_shape_b,
516
+ target_progress=target_progress_b,
517
+ success_label=success_label_b,
518
+ predict_last_frame_mask=mask_b,
519
+ metadata={"source": "gradio_app", "trajectory": "B"},
520
+ )
521
+
522
+ if prediction_type == "preference":
523
+ # Create PreferenceSample (A = chosen, B = rejected)
524
+ preference_sample = PreferenceSample(
525
+ chosen_trajectory=trajectory_a,
526
+ rejected_trajectory=trajectory_b,
527
+ data_gen_strategy="demo",
528
+ )
529
+
530
+ # Build payload and send to server
531
+ files, sample_data = build_payload([preference_sample])
532
+ response = post_batch_npy(server_url, files, sample_data, timeout_s=120.0)
533
+
534
+ # Process response
535
+ outputs_preference = response.get("outputs_preference", {})
536
+ predictions = outputs_preference.get("predictions", [])
537
+ prediction_probs = outputs_preference.get("prediction_probs", [])
538
+
539
+ result_text = f"**Preference Prediction:**\n"
540
+ if prediction_probs and len(prediction_probs) > 0:
541
+ prob = prediction_probs[0]
542
+ result_text += f"- Probability (A preferred): {prob:.3f}\n"
543
+ result_text += f"- Interpretation: {'Video A is preferred' if prob > 0.5 else 'Video B is preferred'}\n"
544
+ else:
545
+ result_text += "Could not extract preference prediction from server response.\n"
546
+
547
+ elif prediction_type == "progress":
548
+ # Create ProgressSamples for both videos
549
+ progress_sample_a = ProgressSample(
550
+ trajectory=trajectory_a,
551
+ data_gen_strategy="demo",
552
+ )
553
+ progress_sample_b = ProgressSample(
554
+ trajectory=trajectory_b,
555
+ data_gen_strategy="demo",
556
+ )
557
+
558
+ # Build payload and send to server
559
+ files, sample_data = build_payload([progress_sample_a, progress_sample_b])
560
+ response = post_batch_npy(server_url, files, sample_data, timeout_s=120.0)
561
+
562
+ # Process response
563
+ outputs_progress = response.get("outputs_progress", {})
564
+ progress_pred = outputs_progress.get("progress_pred", [])
565
+
566
+ result_text = f"**Progress Comparison:**\n"
567
+ if progress_pred and len(progress_pred) >= 2:
568
+ progress_a = np.array(progress_pred[0])
569
+ progress_b = np.array(progress_pred[1])
570
+
571
+ final_progress_a = float(progress_a[-1]) if len(progress_a) > 0 else 0.0
572
+ final_progress_b = float(progress_b[-1]) if len(progress_b) > 0 else 0.0
573
+
574
+ result_text += f"- Video A final progress: {final_progress_a:.3f}\n"
575
+ result_text += f"- Video B final progress: {final_progress_b:.3f}\n"
576
+ result_text += f"- Difference: {abs(final_progress_a - final_progress_b):.3f}\n"
577
+ if final_progress_a > final_progress_b:
578
+ result_text += f"- Video A has higher progress\n"
579
+ elif final_progress_b > final_progress_a:
580
+ result_text += f"- Video B has higher progress\n"
581
+ else:
582
+ result_text += f"- Both videos have equal progress\n"
583
+ else:
584
+ result_text += "Could not extract progress predictions from server response.\n"
585
+
586
+ # Return result text and both video paths
587
+ return result_text, video_a_path, video_b_path
588
+
589
+ except Exception as e:
590
+ return f"Error processing videos: {str(e)}", None, None
591
+
592
+
593
+ # Create Gradio interface
594
+ try:
595
+ # Try with theme (Gradio 4.0+)
596
+ demo = gr.Blocks(title="Robometer Evaluation Server", theme=gr.themes.Soft())
597
+ except TypeError:
598
+ # Fallback for older Gradio versions without theme support
599
+ demo = gr.Blocks(title="Robometer Evaluation Server")
600
+
601
+ with demo:
602
+ gr.Markdown(
603
+ """
604
+ # Robometer Evaluation Server
605
+ """
606
+ )
607
+
608
+ # Hidden state to store server URL and model mapping (define before use)
609
+ server_url_state = gr.State(value=None)
610
+ model_url_mapping_state = gr.State(value={}) # Maps model_name -> server_url
611
+
612
+ # Function definitions for event handlers
613
+ def discover_and_select_models(base_url: str):
614
+ """Discover models and update dropdown."""
615
+ if not base_url:
616
+ return (
617
+ gr.update(choices=[], value=None),
618
+ gr.update(value="Please provide a base URL", visible=True),
619
+ gr.update(value="", visible=True),
620
+ None,
621
+ {}, # Empty mapping
622
+ )
623
+
624
+ _server_state["base_url"] = base_url
625
+ models = discover_available_models(base_url, port_range=(8000, 8010))
626
+
627
+ if not models:
628
+ return (
629
+ gr.update(choices=[], value=None),
630
+ gr.update(value="❌ No models found on ports 8000-8010. Make sure servers are running.", visible=True),
631
+ gr.update(value="", visible=True),
632
+ None,
633
+ {}, # Empty mapping
634
+ )
635
+
636
+ # Format choices: show model_name in dropdown
637
+ # Store mapping of model_name to URL in state
638
+ choices = []
639
+ url_map = {}
640
+ for url, name in models:
641
+ choices.append(name)
642
+ url_map[name] = url
643
+
644
+ # Auto-select first model
645
+ selected_choice = choices[0] if choices else None
646
+ selected_url = url_map.get(selected_choice) if selected_choice else None
647
+
648
+ # Get model info for selected model
649
+ model_info_text = get_model_info_for_url(selected_url) if selected_url else ""
650
+ status_text = f"✅ Found {len(models)} model(s). Auto-selected first model."
651
+
652
+ _server_state["server_url"] = selected_url
653
+
654
+ return (
655
+ gr.update(choices=choices, value=selected_choice),
656
+ gr.update(value=status_text, visible=True),
657
+ gr.update(value=model_info_text, visible=True),
658
+ selected_url,
659
+ url_map, # Return mapping for state
660
+ )
661
+
662
+ def on_model_selected(model_choice: str, url_mapping: dict):
663
+ """Handle model selection change."""
664
+ if not model_choice:
665
+ return (
666
+ gr.update(value="No model selected", visible=True),
667
+ gr.update(value="", visible=True),
668
+ None,
669
+ )
670
+
671
+ # Get URL from mapping
672
+ server_url = url_mapping.get(model_choice) if url_mapping else None
673
+
674
+ if not server_url:
675
+ return (
676
+ gr.update(
677
+ value="Could not find server URL for selected model. Please rediscover models.", visible=True
678
+ ),
679
+ gr.update(value="", visible=True),
680
+ None,
681
+ )
682
+
683
+ # Get model info
684
+ model_info_text = get_model_info_for_url(server_url) or ""
685
+ status, health_data, _ = check_server_health(server_url)
686
+
687
+ _server_state["server_url"] = server_url
688
+
689
+ return (
690
+ gr.update(value=status, visible=True),
691
+ gr.update(value=model_info_text, visible=True),
692
+ server_url,
693
+ )
694
+
695
+ # Use Gradio's built-in Sidebar component (collapsible by default)
696
+ with gr.Sidebar():
697
+ gr.Markdown("### 🔧 Model Configuration")
698
+
699
+ base_url_input = gr.Textbox(
700
+ label="Base Server URL",
701
+ placeholder="https://robometer.a.pinggy.link or http://40.119.56.66",
702
+ value="https://robometer.a.pinggy.link",
703
+ interactive=True,
704
+ info="Full URL (e.g. Pinggy tunnel) or host; discovery tries URL as-is first, then ports 8000-8010",
705
+ )
706
+
707
+ discover_btn = gr.Button("🔍 Discover Models", variant="primary", size="lg")
708
+
709
+ model_dropdown = gr.Dropdown(
710
+ label="Select Model",
711
+ choices=[],
712
+ value=None,
713
+ interactive=True,
714
+ info="Click Discover to find the eval server (single URL or ports 8000-8010)",
715
+ )
716
+
717
+ server_status = gr.Markdown("Click 'Discover Models' to find available models")
718
+
719
+ gr.Markdown("---")
720
+ gr.Markdown("### 📋 Model Information")
721
+ model_info_display = gr.Markdown("")
722
+
723
+ # Event handlers for sidebar
724
+ discover_btn.click(
725
+ fn=discover_and_select_models,
726
+ inputs=[base_url_input],
727
+ outputs=[model_dropdown, server_status, model_info_display, server_url_state, model_url_mapping_state],
728
+ )
729
+
730
+ model_dropdown.change(
731
+ fn=on_model_selected,
732
+ inputs=[model_dropdown, model_url_mapping_state],
733
+ outputs=[server_status, model_info_display, server_url_state],
734
+ )
735
+
736
+ # Main content area with tabs
737
+ with gr.Tabs():
738
+ with gr.Tab("Progress Prediction"):
739
+ with gr.Row():
740
+ with gr.Column():
741
+ single_video_input = gr.Video(label="Upload Video", height=300)
742
+ task_text_input = gr.Textbox(
743
+ label="Task Description",
744
+ placeholder="Describe the task (e.g., 'Pick up the red block')",
745
+ value="Complete the task",
746
+ )
747
+ fps_input_single = gr.Slider(
748
+ label="FPS (Frames Per Second)",
749
+ minimum=0.1,
750
+ maximum=10.0,
751
+ value=1.0,
752
+ step=0.1,
753
+ info="Frames per second to extract from video (higher = more frames)",
754
+ )
755
+ use_frame_steps_single = gr.Checkbox(
756
+ label="Per Frame Progress Prediction",
757
+ value=False,
758
+ info="If enabled, predict progress per frame rather than feeding the entire video at once",
759
+ )
760
+ analyze_single_btn = gr.Button("Compute Progress", variant="primary")
761
+
762
+ gr.Markdown("---")
763
+ gr.Markdown("**OR Select from Dataset**")
764
+ gr.Markdown("---")
765
+
766
+ with gr.Accordion("📁 Select from Dataset", open=False):
767
+ dataset_name_single = gr.Dropdown(
768
+ choices=PREDEFINED_DATASETS,
769
+ value="jesbu1/oxe_rfm",
770
+ label="Dataset Name",
771
+ allow_custom_value=True,
772
+ )
773
+ config_name_single = gr.Dropdown(
774
+ choices=[], value="", label="Configuration Name", allow_custom_value=True
775
+ )
776
+ with gr.Row():
777
+ refresh_configs_btn = gr.Button("🔄 Refresh Configs", variant="secondary", size="sm")
778
+ load_dataset_btn = gr.Button("Load Dataset", variant="secondary", size="sm")
779
+
780
+ dataset_status_single = gr.Markdown("", visible=False)
781
+ with gr.Row():
782
+ prev_traj_btn = gr.Button("⬅️ Prev", variant="secondary", size="sm")
783
+ trajectory_slider = gr.Slider(
784
+ minimum=0, maximum=0, step=1, value=0, label="Trajectory Index", interactive=True
785
+ )
786
+ next_traj_btn = gr.Button("Next ➡️", variant="secondary", size="sm")
787
+ trajectory_metadata = gr.Markdown("", visible=False)
788
+ use_dataset_video_btn = gr.Button("Use Selected Video", variant="secondary")
789
+
790
+ with gr.Column():
791
+ progress_plot = gr.Image(label="Progress & Success Prediction", height=400)
792
+ info_output = gr.Markdown("")
793
+
794
+ # State variables for dataset
795
+ current_dataset_single = gr.State(None)
796
+
797
+ def update_config_choices_single(dataset_name):
798
+ """Update config choices when dataset changes."""
799
+ if not dataset_name:
800
+ return gr.update(choices=[], value="")
801
+ try:
802
+ configs = get_available_configs(dataset_name)
803
+ if configs:
804
+ return gr.update(choices=configs, value=configs[0])
805
+ else:
806
+ return gr.update(choices=[], value="")
807
+ except Exception as e:
808
+ logger.warning(f"Could not fetch configs: {e}")
809
+ return gr.update(choices=[], value="")
810
+
811
+ def load_dataset_single(dataset_name, config_name):
812
+ """Load dataset and update slider."""
813
+ dataset, status = load_rbm_dataset(dataset_name, config_name)
814
+ if dataset is not None:
815
+ max_index = len(dataset) - 1
816
+ return (
817
+ dataset,
818
+ gr.update(value=status, visible=True),
819
+ gr.update(
820
+ maximum=max_index, value=0, interactive=True, label=f"Trajectory Index (0 to {max_index})"
821
+ ),
822
+ )
823
+ else:
824
+ return None, gr.update(value=status, visible=True), gr.update(maximum=0, value=0, interactive=False)
825
+
826
+ def use_dataset_video(dataset, index, dataset_name):
827
+ """Load video from dataset and update inputs."""
828
+ if dataset is None:
829
+ return (
830
+ None,
831
+ "Complete the task",
832
+ gr.update(value="No dataset loaded", visible=True),
833
+ gr.update(visible=False),
834
+ )
835
+
836
+ video_path, task, quality_label, partial_success = get_trajectory_video_path(
837
+ dataset, index, dataset_name
838
+ )
839
+ if video_path:
840
+ # Build metadata text
841
+ metadata_lines = []
842
+ if quality_label:
843
+ metadata_lines.append(f"**Quality Label:** {quality_label}")
844
+ if partial_success is not None:
845
+ metadata_lines.append(f"**Partial Success:** {partial_success:.3f}")
846
+
847
+ metadata_text = "\n".join(metadata_lines) if metadata_lines else ""
848
+ status_text = f"✅ Loaded trajectory {index} from dataset"
849
+ if metadata_text:
850
+ status_text += f"\n\n{metadata_text}"
851
+
852
+ return (
853
+ video_path,
854
+ task,
855
+ gr.update(value=status_text, visible=True),
856
+ gr.update(value=metadata_text, visible=bool(metadata_text)),
857
+ )
858
+ else:
859
+ return (
860
+ None,
861
+ "Complete the task",
862
+ gr.update(value="❌ Error loading trajectory", visible=True),
863
+ gr.update(visible=False),
864
+ )
865
+
866
+ def next_trajectory(dataset, current_idx, dataset_name):
867
+ """Go to next trajectory."""
868
+ if dataset is None:
869
+ return 0, None, "Complete the task", gr.update(visible=False), gr.update(visible=False)
870
+ next_idx = min(current_idx + 1, len(dataset) - 1)
871
+ video_path, task, quality_label, partial_success = get_trajectory_video_path(
872
+ dataset, next_idx, dataset_name
873
+ )
874
+
875
+ if video_path:
876
+ # Build metadata text
877
+ metadata_lines = []
878
+ if quality_label:
879
+ metadata_lines.append(f"**Quality Label:** {quality_label}")
880
+ if partial_success is not None:
881
+ metadata_lines.append(f"**Partial Success:** {partial_success:.3f}")
882
+
883
+ metadata_text = "\n".join(metadata_lines) if metadata_lines else ""
884
+ return (
885
+ next_idx,
886
+ video_path,
887
+ task,
888
+ gr.update(value=metadata_text, visible=bool(metadata_text)),
889
+ gr.update(value=f"✅ Trajectory {next_idx}/{len(dataset) - 1}", visible=True),
890
+ )
891
+ else:
892
+ return current_idx, None, "Complete the task", gr.update(visible=False), gr.update(visible=False)
893
+
894
+ def prev_trajectory(dataset, current_idx, dataset_name):
895
+ """Go to previous trajectory."""
896
+ if dataset is None:
897
+ return 0, None, "Complete the task", gr.update(visible=False), gr.update(visible=False)
898
+ prev_idx = max(current_idx - 1, 0)
899
+ video_path, task, quality_label, partial_success = get_trajectory_video_path(
900
+ dataset, prev_idx, dataset_name
901
+ )
902
+
903
+ if video_path:
904
+ # Build metadata text
905
+ metadata_lines = []
906
+ if quality_label:
907
+ metadata_lines.append(f"**Quality Label:** {quality_label}")
908
+ if partial_success is not None:
909
+ metadata_lines.append(f"**Partial Success:** {partial_success:.3f}")
910
+
911
+ metadata_text = "\n".join(metadata_lines) if metadata_lines else ""
912
+ return (
913
+ prev_idx,
914
+ video_path,
915
+ task,
916
+ gr.update(value=metadata_text, visible=bool(metadata_text)),
917
+ gr.update(value=f"✅ Trajectory {prev_idx}/{len(dataset) - 1}", visible=True),
918
+ )
919
+ else:
920
+ return current_idx, None, "Complete the task", gr.update(visible=False), gr.update(visible=False)
921
+
922
+ def update_trajectory_on_slider_change(dataset, index, dataset_name):
923
+ """Update trajectory metadata when slider changes."""
924
+ if dataset is None:
925
+ return gr.update(visible=False), gr.update(visible=False)
926
+
927
+ video_path, task, quality_label, partial_success = get_trajectory_video_path(
928
+ dataset, index, dataset_name
929
+ )
930
+ if video_path:
931
+ # Build metadata text
932
+ metadata_lines = []
933
+ if quality_label:
934
+ metadata_lines.append(f"**Quality Label:** {quality_label}")
935
+ if partial_success is not None:
936
+ metadata_lines.append(f"**Partial Success:** {partial_success:.3f}")
937
+
938
+ metadata_text = "\n".join(metadata_lines) if metadata_lines else ""
939
+ return (
940
+ gr.update(value=metadata_text, visible=bool(metadata_text)),
941
+ gr.update(value=f"Trajectory {index}/{len(dataset) - 1}", visible=True),
942
+ )
943
+ else:
944
+ return gr.update(visible=False), gr.update(visible=False)
945
+
946
+ # Dataset selection handlers
947
+ dataset_name_single.change(
948
+ fn=update_config_choices_single, inputs=[dataset_name_single], outputs=[config_name_single]
949
+ )
950
+
951
+ refresh_configs_btn.click(
952
+ fn=update_config_choices_single, inputs=[dataset_name_single], outputs=[config_name_single]
953
+ )
954
+
955
+ load_dataset_btn.click(
956
+ fn=load_dataset_single,
957
+ inputs=[dataset_name_single, config_name_single],
958
+ outputs=[current_dataset_single, dataset_status_single, trajectory_slider],
959
+ )
960
+
961
+ use_dataset_video_btn.click(
962
+ fn=use_dataset_video,
963
+ inputs=[current_dataset_single, trajectory_slider, dataset_name_single],
964
+ outputs=[single_video_input, task_text_input, dataset_status_single, trajectory_metadata],
965
+ )
966
+
967
+ # Navigation buttons
968
+ next_traj_btn.click(
969
+ fn=next_trajectory,
970
+ inputs=[current_dataset_single, trajectory_slider, dataset_name_single],
971
+ outputs=[
972
+ trajectory_slider,
973
+ single_video_input,
974
+ task_text_input,
975
+ trajectory_metadata,
976
+ dataset_status_single,
977
+ ],
978
+ )
979
+
980
+ prev_traj_btn.click(
981
+ fn=prev_trajectory,
982
+ inputs=[current_dataset_single, trajectory_slider, dataset_name_single],
983
+ outputs=[
984
+ trajectory_slider,
985
+ single_video_input,
986
+ task_text_input,
987
+ trajectory_metadata,
988
+ dataset_status_single,
989
+ ],
990
+ )
991
+
992
+ # Update metadata when slider changes
993
+ trajectory_slider.change(
994
+ fn=update_trajectory_on_slider_change,
995
+ inputs=[current_dataset_single, trajectory_slider, dataset_name_single],
996
+ outputs=[trajectory_metadata, dataset_status_single],
997
+ )
998
+
999
+ analyze_single_btn.click(
1000
+ fn=process_single_video,
1001
+ inputs=[
1002
+ single_video_input,
1003
+ task_text_input,
1004
+ server_url_state,
1005
+ fps_input_single,
1006
+ use_frame_steps_single,
1007
+ ],
1008
+ outputs=[progress_plot, info_output],
1009
+ api_name="process_single_video",
1010
+ )
1011
+
1012
+ with gr.Tab("Preference Analysis"):
1013
+ # Full-width row: two videos side by side
1014
+ with gr.Row():
1015
+ video_a_input = gr.Video(label="Video A", height=320)
1016
+ video_b_input = gr.Video(label="Video B", height=320)
1017
+
1018
+ task_text_dual = gr.Textbox(
1019
+ label="Task Description",
1020
+ placeholder="Describe the task",
1021
+ value="Complete the task",
1022
+ )
1023
+ analyze_dual_btn = gr.Button("Compute Preference", variant="primary")
1024
+
1025
+ gr.Markdown("---")
1026
+ gr.Markdown("**OR Select from Dataset**")
1027
+ gr.Markdown("---")
1028
+
1029
+ with gr.Accordion("📁 Video A - Select from Dataset", open=False):
1030
+ dataset_name_a = gr.Dropdown(
1031
+ choices=PREDEFINED_DATASETS,
1032
+ value="jesbu1/oxe_rfm",
1033
+ label="Dataset Name",
1034
+ allow_custom_value=True,
1035
+ )
1036
+ config_name_a = gr.Dropdown(
1037
+ choices=[], value="", label="Configuration Name", allow_custom_value=True
1038
+ )
1039
+ with gr.Row():
1040
+ refresh_configs_btn_a = gr.Button("🔄 Refresh Configs", variant="secondary", size="sm")
1041
+ load_dataset_btn_a = gr.Button("Load Dataset", variant="secondary", size="sm")
1042
+
1043
+ dataset_status_a = gr.Markdown("", visible=False)
1044
+ with gr.Row():
1045
+ prev_traj_btn_a = gr.Button("⬅️ Prev", variant="secondary", size="sm")
1046
+ trajectory_slider_a = gr.Slider(
1047
+ minimum=0, maximum=0, step=1, value=0, label="Trajectory Index", interactive=True
1048
+ )
1049
+ next_traj_btn_a = gr.Button("Next ➡️", variant="secondary", size="sm")
1050
+ trajectory_metadata_a = gr.Markdown("", visible=False)
1051
+ use_dataset_video_btn_a = gr.Button("Use Selected Video for A", variant="secondary")
1052
+
1053
+ with gr.Accordion("📁 Video B - Select from Dataset", open=False):
1054
+ dataset_name_b = gr.Dropdown(
1055
+ choices=PREDEFINED_DATASETS,
1056
+ value="jesbu1/oxe_rfm",
1057
+ label="Dataset Name",
1058
+ allow_custom_value=True,
1059
+ )
1060
+ config_name_b = gr.Dropdown(
1061
+ choices=[], value="", label="Configuration Name", allow_custom_value=True
1062
+ )
1063
+ with gr.Row():
1064
+ refresh_configs_btn_b = gr.Button("🔄 Refresh Configs", variant="secondary", size="sm")
1065
+ load_dataset_btn_b = gr.Button("Load Dataset", variant="secondary", size="sm")
1066
+
1067
+ dataset_status_b = gr.Markdown("", visible=False)
1068
+ with gr.Row():
1069
+ prev_traj_btn_b = gr.Button("⬅️ Prev", variant="secondary", size="sm")
1070
+ trajectory_slider_b = gr.Slider(
1071
+ minimum=0, maximum=0, step=1, value=0, label="Trajectory Index", interactive=True
1072
+ )
1073
+ next_traj_btn_b = gr.Button("Next ➡️", variant="secondary", size="sm")
1074
+ trajectory_metadata_b = gr.Markdown("", visible=False)
1075
+ use_dataset_video_btn_b = gr.Button("Use Selected Video for B", variant="secondary")
1076
+
1077
+ gr.Markdown("---")
1078
+ gr.Markdown("### Preference result")
1079
+ result_text = gr.Markdown("")
1080
+
1081
+ # State variables for datasets
1082
+ current_dataset_a = gr.State(None)
1083
+ current_dataset_b = gr.State(None)
1084
+
1085
+ # Helper functions for Video A
1086
+ def update_config_choices_a(dataset_name):
1087
+ """Update config choices for Video A when dataset changes."""
1088
+ if not dataset_name:
1089
+ return gr.update(choices=[], value="")
1090
+ try:
1091
+ configs = get_available_configs(dataset_name)
1092
+ if configs:
1093
+ return gr.update(choices=configs, value=configs[0])
1094
+ else:
1095
+ return gr.update(choices=[], value="")
1096
+ except Exception as e:
1097
+ logger.warning(f"Could not fetch configs: {e}")
1098
+ return gr.update(choices=[], value="")
1099
+
1100
+ def load_dataset_a(dataset_name, config_name):
1101
+ """Load dataset A and update slider."""
1102
+ dataset, status = load_rbm_dataset(dataset_name, config_name)
1103
+ if dataset is not None:
1104
+ max_index = len(dataset) - 1
1105
+ return (
1106
+ dataset,
1107
+ gr.update(value=status, visible=True),
1108
+ gr.update(
1109
+ maximum=max_index, value=0, interactive=True, label=f"Trajectory Index (0 to {max_index})"
1110
+ ),
1111
+ )
1112
+ else:
1113
+ return None, gr.update(value=status, visible=True), gr.update(maximum=0, value=0, interactive=False)
1114
+
1115
+ def use_dataset_video_a(dataset, index, dataset_name):
1116
+ """Load video A from dataset and update input."""
1117
+ if dataset is None:
1118
+ return (
1119
+ None,
1120
+ gr.update(value="No dataset loaded", visible=True),
1121
+ gr.update(visible=False),
1122
+ )
1123
+
1124
+ video_path, task, quality_label, partial_success = get_trajectory_video_path(
1125
+ dataset, index, dataset_name
1126
+ )
1127
+ if video_path:
1128
+ # Build metadata text
1129
+ metadata_lines = []
1130
+ if quality_label:
1131
+ metadata_lines.append(f"**Quality Label:** {quality_label}")
1132
+ if partial_success is not None:
1133
+ metadata_lines.append(f"**Partial Success:** {partial_success:.3f}")
1134
+
1135
+ metadata_text = "\n".join(metadata_lines) if metadata_lines else ""
1136
+ status_text = f"✅ Loaded trajectory {index} from dataset for Video A"
1137
+ if metadata_text:
1138
+ status_text += f"\n\n{metadata_text}"
1139
+
1140
+ return (
1141
+ video_path,
1142
+ gr.update(value=status_text, visible=True),
1143
+ gr.update(value=metadata_text, visible=bool(metadata_text)),
1144
+ )
1145
+ else:
1146
+ return (
1147
+ None,
1148
+ gr.update(value="❌ Error loading trajectory", visible=True),
1149
+ gr.update(visible=False),
1150
+ )
1151
+
1152
+ def next_trajectory_a(dataset, current_idx, dataset_name):
1153
+ """Go to next trajectory for Video A."""
1154
+ if dataset is None:
1155
+ return 0, None, gr.update(visible=False), gr.update(visible=False)
1156
+ next_idx = min(current_idx + 1, len(dataset) - 1)
1157
+ video_path, task, quality_label, partial_success = get_trajectory_video_path(
1158
+ dataset, next_idx, dataset_name
1159
+ )
1160
+
1161
+ if video_path:
1162
+ # Build metadata text
1163
+ metadata_lines = []
1164
+ if quality_label:
1165
+ metadata_lines.append(f"**Quality Label:** {quality_label}")
1166
+ if partial_success is not None:
1167
+ metadata_lines.append(f"**Partial Success:** {partial_success:.3f}")
1168
+
1169
+ metadata_text = "\n".join(metadata_lines) if metadata_lines else ""
1170
+ return (
1171
+ next_idx,
1172
+ video_path,
1173
+ gr.update(value=metadata_text, visible=bool(metadata_text)),
1174
+ gr.update(value=f"✅ Trajectory {next_idx}/{len(dataset) - 1}", visible=True),
1175
+ )
1176
+ else:
1177
+ return current_idx, None, gr.update(visible=False), gr.update(visible=False)
1178
+
1179
+ def prev_trajectory_a(dataset, current_idx, dataset_name):
1180
+ """Go to previous trajectory for Video A."""
1181
+ if dataset is None:
1182
+ return 0, None, gr.update(visible=False), gr.update(visible=False)
1183
+ prev_idx = max(current_idx - 1, 0)
1184
+ video_path, task, quality_label, partial_success = get_trajectory_video_path(
1185
+ dataset, prev_idx, dataset_name
1186
+ )
1187
+
1188
+ if video_path:
1189
+ # Build metadata text
1190
+ metadata_lines = []
1191
+ if quality_label:
1192
+ metadata_lines.append(f"**Quality Label:** {quality_label}")
1193
+ if partial_success is not None:
1194
+ metadata_lines.append(f"**Partial Success:** {partial_success:.3f}")
1195
+
1196
+ metadata_text = "\n".join(metadata_lines) if metadata_lines else ""
1197
+ return (
1198
+ prev_idx,
1199
+ video_path,
1200
+ gr.update(value=metadata_text, visible=bool(metadata_text)),
1201
+ gr.update(value=f"✅ Trajectory {prev_idx}/{len(dataset) - 1}", visible=True),
1202
+ )
1203
+ else:
1204
+ return current_idx, None, gr.update(visible=False), gr.update(visible=False)
1205
+
1206
+ def update_trajectory_on_slider_change_a(dataset, index, dataset_name):
1207
+ """Update trajectory metadata when slider changes for Video A."""
1208
+ if dataset is None:
1209
+ return gr.update(visible=False), gr.update(visible=False)
1210
+
1211
+ video_path, task, quality_label, partial_success = get_trajectory_video_path(
1212
+ dataset, index, dataset_name
1213
+ )
1214
+ if video_path:
1215
+ # Build metadata text
1216
+ metadata_lines = []
1217
+ if quality_label:
1218
+ metadata_lines.append(f"**Quality Label:** {quality_label}")
1219
+ if partial_success is not None:
1220
+ metadata_lines.append(f"**Partial Success:** {partial_success:.3f}")
1221
+
1222
+ metadata_text = "\n".join(metadata_lines) if metadata_lines else ""
1223
+ return (
1224
+ gr.update(value=metadata_text, visible=bool(metadata_text)),
1225
+ gr.update(value=f"Trajectory {index}/{len(dataset) - 1}", visible=True),
1226
+ )
1227
+ else:
1228
+ return gr.update(visible=False), gr.update(visible=False)
1229
+
1230
+ # Helper functions for Video B (same as Video A)
1231
+ def update_config_choices_b(dataset_name):
1232
+ """Update config choices for Video B when dataset changes."""
1233
+ if not dataset_name:
1234
+ return gr.update(choices=[], value="")
1235
+ try:
1236
+ configs = get_available_configs(dataset_name)
1237
+ if configs:
1238
+ return gr.update(choices=configs, value=configs[0])
1239
+ else:
1240
+ return gr.update(choices=[], value="")
1241
+ except Exception as e:
1242
+ logger.warning(f"Could not fetch configs: {e}")
1243
+ return gr.update(choices=[], value="")
1244
+
1245
+ def load_dataset_b(dataset_name, config_name):
1246
+ """Load dataset B and update slider."""
1247
+ dataset, status = load_rbm_dataset(dataset_name, config_name)
1248
+ if dataset is not None:
1249
+ max_index = len(dataset) - 1
1250
+ return (
1251
+ dataset,
1252
+ gr.update(value=status, visible=True),
1253
+ gr.update(
1254
+ maximum=max_index, value=0, interactive=True, label=f"Trajectory Index (0 to {max_index})"
1255
+ ),
1256
+ )
1257
+ else:
1258
+ return None, gr.update(value=status, visible=True), gr.update(maximum=0, value=0, interactive=False)
1259
+
1260
+ def use_dataset_video_b(dataset, index, dataset_name):
1261
+ """Load video B from dataset and update input."""
1262
+ if dataset is None:
1263
+ return (
1264
+ None,
1265
+ gr.update(value="No dataset loaded", visible=True),
1266
+ gr.update(visible=False),
1267
+ )
1268
+
1269
+ video_path, task, quality_label, partial_success = get_trajectory_video_path(
1270
+ dataset, index, dataset_name
1271
+ )
1272
+ if video_path:
1273
+ # Build metadata text
1274
+ metadata_lines = []
1275
+ if quality_label:
1276
+ metadata_lines.append(f"**Quality Label:** {quality_label}")
1277
+ if partial_success is not None:
1278
+ metadata_lines.append(f"**Partial Success:** {partial_success:.3f}")
1279
+
1280
+ metadata_text = "\n".join(metadata_lines) if metadata_lines else ""
1281
+ status_text = f"✅ Loaded trajectory {index} from dataset for Video B"
1282
+ if metadata_text:
1283
+ status_text += f"\n\n{metadata_text}"
1284
+
1285
+ return (
1286
+ video_path,
1287
+ gr.update(value=status_text, visible=True),
1288
+ gr.update(value=metadata_text, visible=bool(metadata_text)),
1289
+ )
1290
+ else:
1291
+ return (
1292
+ None,
1293
+ gr.update(value="❌ Error loading trajectory", visible=True),
1294
+ gr.update(visible=False),
1295
+ )
1296
+
1297
+ def next_trajectory_b(dataset, current_idx, dataset_name):
1298
+ """Go to next trajectory for Video B."""
1299
+ if dataset is None:
1300
+ return 0, None, gr.update(visible=False), gr.update(visible=False)
1301
+ next_idx = min(current_idx + 1, len(dataset) - 1)
1302
+ video_path, task, quality_label, partial_success = get_trajectory_video_path(
1303
+ dataset, next_idx, dataset_name
1304
+ )
1305
+
1306
+ if video_path:
1307
+ # Build metadata text
1308
+ metadata_lines = []
1309
+ if quality_label:
1310
+ metadata_lines.append(f"**Quality Label:** {quality_label}")
1311
+ if partial_success is not None:
1312
+ metadata_lines.append(f"**Partial Success:** {partial_success:.3f}")
1313
+
1314
+ metadata_text = "\n".join(metadata_lines) if metadata_lines else ""
1315
+ return (
1316
+ next_idx,
1317
+ video_path,
1318
+ gr.update(value=metadata_text, visible=bool(metadata_text)),
1319
+ gr.update(value=f"✅ Trajectory {next_idx}/{len(dataset) - 1}", visible=True),
1320
+ )
1321
+ else:
1322
+ return current_idx, None, gr.update(visible=False), gr.update(visible=False)
1323
+
1324
+ def prev_trajectory_b(dataset, current_idx, dataset_name):
1325
+ """Go to previous trajectory for Video B."""
1326
+ if dataset is None:
1327
+ return 0, None, gr.update(visible=False), gr.update(visible=False)
1328
+ prev_idx = max(current_idx - 1, 0)
1329
+ video_path, task, quality_label, partial_success = get_trajectory_video_path(
1330
+ dataset, prev_idx, dataset_name
1331
+ )
1332
+
1333
+ if video_path:
1334
+ # Build metadata text
1335
+ metadata_lines = []
1336
+ if quality_label:
1337
+ metadata_lines.append(f"**Quality Label:** {quality_label}")
1338
+ if partial_success is not None:
1339
+ metadata_lines.append(f"**Partial Success:** {partial_success:.3f}")
1340
+
1341
+ metadata_text = "\n".join(metadata_lines) if metadata_lines else ""
1342
+ return (
1343
+ prev_idx,
1344
+ video_path,
1345
+ gr.update(value=metadata_text, visible=bool(metadata_text)),
1346
+ gr.update(value=f"✅ Trajectory {prev_idx}/{len(dataset) - 1}", visible=True),
1347
+ )
1348
+ else:
1349
+ return current_idx, None, gr.update(visible=False), gr.update(visible=False)
1350
+
1351
+ def update_trajectory_on_slider_change_b(dataset, index, dataset_name):
1352
+ """Update trajectory metadata when slider changes for Video B."""
1353
+ if dataset is None:
1354
+ return gr.update(visible=False), gr.update(visible=False)
1355
+
1356
+ video_path, task, quality_label, partial_success = get_trajectory_video_path(
1357
+ dataset, index, dataset_name
1358
+ )
1359
+ if video_path:
1360
+ # Build metadata text
1361
+ metadata_lines = []
1362
+ if quality_label:
1363
+ metadata_lines.append(f"**Quality Label:** {quality_label}")
1364
+ if partial_success is not None:
1365
+ metadata_lines.append(f"**Partial Success:** {partial_success:.3f}")
1366
+
1367
+ metadata_text = "\n".join(metadata_lines) if metadata_lines else ""
1368
+ return (
1369
+ gr.update(value=metadata_text, visible=bool(metadata_text)),
1370
+ gr.update(value=f"Trajectory {index}/{len(dataset) - 1}", visible=True),
1371
+ )
1372
+ else:
1373
+ return gr.update(visible=False), gr.update(visible=False)
1374
+
1375
+ # Video A dataset selection handlers
1376
+ dataset_name_a.change(fn=update_config_choices_a, inputs=[dataset_name_a], outputs=[config_name_a])
1377
+
1378
+ refresh_configs_btn_a.click(fn=update_config_choices_a, inputs=[dataset_name_a], outputs=[config_name_a])
1379
+
1380
+ load_dataset_btn_a.click(
1381
+ fn=load_dataset_a,
1382
+ inputs=[dataset_name_a, config_name_a],
1383
+ outputs=[current_dataset_a, dataset_status_a, trajectory_slider_a],
1384
+ )
1385
+
1386
+ use_dataset_video_btn_a.click(
1387
+ fn=use_dataset_video_a,
1388
+ inputs=[current_dataset_a, trajectory_slider_a, dataset_name_a],
1389
+ outputs=[video_a_input, dataset_status_a, trajectory_metadata_a],
1390
+ )
1391
+
1392
+ next_traj_btn_a.click(
1393
+ fn=next_trajectory_a,
1394
+ inputs=[current_dataset_a, trajectory_slider_a, dataset_name_a],
1395
+ outputs=[
1396
+ trajectory_slider_a,
1397
+ video_a_input,
1398
+ trajectory_metadata_a,
1399
+ dataset_status_a,
1400
+ ],
1401
+ )
1402
+
1403
+ prev_traj_btn_a.click(
1404
+ fn=prev_trajectory_a,
1405
+ inputs=[current_dataset_a, trajectory_slider_a, dataset_name_a],
1406
+ outputs=[
1407
+ trajectory_slider_a,
1408
+ video_a_input,
1409
+ trajectory_metadata_a,
1410
+ dataset_status_a,
1411
+ ],
1412
+ )
1413
+
1414
+ trajectory_slider_a.change(
1415
+ fn=update_trajectory_on_slider_change_a,
1416
+ inputs=[current_dataset_a, trajectory_slider_a, dataset_name_a],
1417
+ outputs=[trajectory_metadata_a, dataset_status_a],
1418
+ )
1419
+
1420
+ # Video B dataset selection handlers
1421
+ dataset_name_b.change(fn=update_config_choices_b, inputs=[dataset_name_b], outputs=[config_name_b])
1422
+
1423
+ refresh_configs_btn_b.click(fn=update_config_choices_b, inputs=[dataset_name_b], outputs=[config_name_b])
1424
+
1425
+ load_dataset_btn_b.click(
1426
+ fn=load_dataset_b,
1427
+ inputs=[dataset_name_b, config_name_b],
1428
+ outputs=[current_dataset_b, dataset_status_b, trajectory_slider_b],
1429
+ )
1430
+
1431
+ use_dataset_video_btn_b.click(
1432
+ fn=use_dataset_video_b,
1433
+ inputs=[current_dataset_b, trajectory_slider_b, dataset_name_b],
1434
+ outputs=[video_b_input, dataset_status_b, trajectory_metadata_b],
1435
+ )
1436
+
1437
+ next_traj_btn_b.click(
1438
+ fn=next_trajectory_b,
1439
+ inputs=[current_dataset_b, trajectory_slider_b, dataset_name_b],
1440
+ outputs=[
1441
+ trajectory_slider_b,
1442
+ video_b_input,
1443
+ trajectory_metadata_b,
1444
+ dataset_status_b,
1445
+ ],
1446
+ )
1447
+
1448
+ prev_traj_btn_b.click(
1449
+ fn=prev_trajectory_b,
1450
+ inputs=[current_dataset_b, trajectory_slider_b, dataset_name_b],
1451
+ outputs=[
1452
+ trajectory_slider_b,
1453
+ video_b_input,
1454
+ trajectory_metadata_b,
1455
+ dataset_status_b,
1456
+ ],
1457
+ )
1458
+
1459
+ trajectory_slider_b.change(
1460
+ fn=update_trajectory_on_slider_change_b,
1461
+ inputs=[current_dataset_b, trajectory_slider_b, dataset_name_b],
1462
+ outputs=[trajectory_metadata_b, dataset_status_b],
1463
+ )
1464
+
1465
+ def run_preference_comparison(video_a, video_b, task_text, server_url):
1466
+ result, _, _ = process_two_videos(
1467
+ video_a, video_b, task_text, "preference", server_url, fps=1.0
1468
+ )
1469
+ return result
1470
+
1471
+ analyze_dual_btn.click(
1472
+ fn=run_preference_comparison,
1473
+ inputs=[
1474
+ video_a_input,
1475
+ video_b_input,
1476
+ task_text_dual,
1477
+ server_url_state,
1478
+ ],
1479
+ outputs=[result_text],
1480
+ api_name="process_two_videos",
1481
+ )
1482
+
1483
+
1484
+ def main():
1485
+ """Launch the Gradio app."""
1486
+ demo.launch(
1487
+ server_name="0.0.0.0",
1488
+ server_port=7860,
1489
+ share=False,
1490
+ show_error=True, # Show full error messages
1491
+ )
1492
+
1493
+
1494
+ if __name__ == "__main__":
1495
+ main()