dizolivemint/motion-encoder-decoder
0
1# gradio_full_system.py2 3import gradio as gr4import torch5import pandas as pd6import random7import matplotlib.pyplot as plt8from model_training.model_torch import EncoderDecoder9from model_training.train import train_model10from model_training.generate_training_data import generate_training_data11from model_training.train import PhysicsTrajectoryDataset12from video_sequencer.generate_frames_and_video import generate_frames_and_video13import os14import torch.nn as nn15import numpy as np16from config import normalize_input, denormalize_input, get_input_fields, get_physics_types, get_param_ranges17from utils.path_utils import resolve_path18 19# --- Available Physics Types ---20physics_types = get_physics_types()21 22# --- Inspect Training Trajectories ---23def inspect_training_trajectories(physics_type, frame_size=64):24 dataset_path = resolve_path(f"{physics_type}_data.pkl")25 dataset = PhysicsTrajectoryDataset(dataset_path, physics_type)26 27 num_samples = 328 fig, axes = plt.subplots(num_samples, 3, figsize=(10, 3 * num_samples), constrained_layout=True)29 30 for i in range(num_samples):31 row = dataset.df.iloc[i]32 input_features, trajectory = dataset[i] # trajectory: [T, 2]33 T = trajectory.shape[0]34 angle = row.get("angle", None)35 36 for j, t in enumerate([0, T // 2, T - 1]):37 ax = axes[i, j] if num_samples > 1 else axes[j]38 frame = torch.zeros((frame_size, frame_size))39 40 # Get (x, y) and denormalize back to pixel space41 x, y = trajectory[t].numpy()42 px = int(x * (frame_size - 1))43 py = int(y * (frame_size - 1))44 45 # Plot the point46 ax.imshow(frame, cmap="gray")47 ax.plot(px, py, "ro", markersize=5)48 49 # Initialize title50 title = f"t={t}"51 52 # Dynamically append available metadata fields53 for key in ['mass', 'angle', 'friction', 'initial_velocity', 'acceleration', 'gravity']:54 if key in row:55 val = row[key]56 if key == 'angle':57 title += f"\n{key.capitalize()}: {val:.1f}ยฐ"58 else:59 title += f"\n{key.capitalize()}: {val:.2f}"60 if angle is not None:61 title += f"\nAngle: {angle:.1f}ยฐ"62 63 ax.set_title(title, pad=10)64 ax.axis("off")65 66 return fig67 68# --- Predict Normalized Coordinates ---69def predict_trajectory(physics_type, *inputs, debug=False):70 model_file = f"{physics_type}_model.pth"71 model_path = resolve_path(model_file)72 if not os.path.exists(model_path):73 return None74 75 sample_data = pd.read_pickle(resolve_path(f"{physics_type}_data.pkl"))76 input_dim = len(sample_data.columns) - 177 output_seq_len = len(sample_data.iloc[0]['trajectory']) # [T, 2]78 79 model = EncoderDecoder(80 input_dim=input_dim,81 output_seq_len=output_seq_len,82 output_shape=None # Use coordinate mode83 )84 checkpoint = torch.load(model_path)85 model.load_state_dict(checkpoint['model_state'])86 model.eval()87 88 inputs_tensor = torch.tensor([inputs], dtype=torch.float32)89 with torch.no_grad():90 prediction = model(inputs_tensor)91 prediction = prediction.cpu().numpy()[0] # [T, 2]92 93 if debug:94 print("๐ Debug output โ predicted (x, y) per frame:")95 for t, (x, y) in enumerate(prediction):96 print(f"t={t}: ({x:.3f}, {y:.3f})")97 98 return prediction # [T, 2]99 100# --- Plotting for Normalized (x, y) ---101def plot_coordinates_over_time(coords, title="Predicted Dot Trajectory", frame_size=64):102 T = len(coords)103 fig, axes = plt.subplots(1, 3, figsize=(12, 4))104 indices = [0, T // 2, T - 1]105 106 for ax, idx in zip(axes, indices):107 x, y = coords[idx]108 px = int(x * (frame_size - 1))109 py = int(y * (frame_size - 1))110 111 frame = np.zeros((frame_size, frame_size))112 ax.imshow(frame, cmap='gray')113 ax.plot(px, py, 'ro')114 ax.set_title(f"t={idx} | ({px}, {py})")115 ax.axis('off')116 117 fig.suptitle(title)118 return fig119 120def plot_multiple_predictions(predictions, inputs_list, frame_size=64, denorm=None):121 N = len(predictions)122 fig, axes = plt.subplots(N, 3, figsize=(12, 3 * N), constrained_layout=True)123 124 for row in range(N):125 indices = [0, len(predictions[row]) // 2, len(predictions[row]) - 1]126 inp = inputs_list[row]127 128 # Denormalize if provided129 labels = denorm(inp) if denorm else inp130 131 for col, t in enumerate(indices):132 x, y = predictions[row][t]133 px = int(x * (frame_size - 1))134 py = int(y * (frame_size - 1))135 136 px = round(px, 2)137 py = round(py, 2)138 139 frame = np.zeros((frame_size, frame_size))140 ax = axes[row, col] if N > 1 else axes[col]141 ax.imshow(frame, cmap='gray')142 ax.plot(px, py, 'ro')143 ax.set_title(f"t={t} | inputs={labels}\nDot: ({px}, {py})")144 ax.axis("off")145 146 return fig147 148def plot_trajectory(physics_type, *inputs, debug=False):149 pred = predict_trajectory(physics_type, *inputs, debug=debug)150 if pred is None:151 return None152 fig = plot_coordinates_over_time(pred, title=f"{physics_type.title()} Prediction")153 return fig, pred154 155# --- Video Generation ---156def predict_plot_video(physics_type, *inputs, debug=False):157 norm_input = normalize_input(physics_type, *inputs)158 159 fig, pred = plot_trajectory(physics_type, *norm_input, debug=debug)160 if pred is None:161 return None, None162 video_mp4_path = generate_frames_and_video(pred)163 return fig, video_mp4_path164 165def test_input_sensitivity(physics_type):166 model_path = resolve_path(f"{physics_type}_model.pth")167 if not os.path.exists(model_path):168 return None169 170 # Load sample data to get dimensions171 sample_data = pd.read_pickle(f"data/{physics_type}_data.pkl")172 input_dim = len(sample_data.columns) - 1173 output_seq_len = len(sample_data.iloc[0]['trajectory']) # [T, 2]174 175 # Load trained model176 model = EncoderDecoder(177 input_dim=input_dim,178 output_seq_len=output_seq_len,179 output_shape=None180 )181 checkpoint = torch.load(model_path)182 model.load_state_dict(checkpoint['model_state'])183 model.eval()184 185 # Get parameter ranges and field names186 param_ranges = get_param_ranges(physics_type)187 param_specs = get_input_fields(physics_type)188 189 # Dynamically generate test inputs (midpoints or meaningful values)190 test_inputs = []191 for variation in range(4):192 base = []193 for field in param_specs:194 min_val, max_val = param_ranges[field]195 mid = (min_val + max_val) / 2196 val = mid + ((variation - 1.5) * (max_val - min_val) / 4)197 val = max(min_val, min(max_val, val)) # Clamp198 base.append(val)199 test_inputs.append(normalize_input(physics_type, *base))200 201 preds = []202 for inp in test_inputs:203 inp_tensor = torch.tensor([inp], dtype=torch.float32)204 with torch.no_grad():205 output = model(inp_tensor).cpu().numpy()[0] # [T, 2]206 preds.append(output)207 208 # Create a closure that captures physics_type209 def denorm_fn(inputs):210 return denormalize_input(physics_type, inputs)211 212 fig = plot_multiple_predictions(preds, test_inputs, denorm=denorm_fn)213 return fig214 215def load_uploaded_dataset(uploaded_file, physics_type):216 tmp_path = os.path.join("/tmp", f"{physics_type}_data.pkl")217 with open(uploaded_file.name, "rb") as src, open(tmp_path, "wb") as dst:218 dst.write(src.read())219 return f"โ
Uploaded and saved dataset to /tmp for '{physics_type}'"220 221def load_uploaded_model(uploaded_file, physics_type):222 tmp_path = os.path.join("/tmp", "", f"{physics_type}_model.pth")223 os.makedirs(os.path.dirname(tmp_path), exist_ok=True)224 with open(uploaded_file.name, "rb") as src, open(tmp_path, "wb") as dst:225 dst.write(src.read())226 return f"โ
Uploaded and saved model to /tmp for '{physics_type}'"227 228# --- Gradio Interface ---229with gr.Blocks() as demo:230 gr.Markdown("# ๐ง Full Physics ML System")231 gr.Markdown("## ๐งช Data โ Train โ Predict")232 gr.Markdown("### ๐จโ๐ป Developed by [Miles Exner](https://www.linkedin.com/in/milesexner/)")233 234 # ๐ Global dropdown visible to user235 physics_dropdown = gr.Dropdown(236 choices=physics_types,237 label="Physics Type",238 value=physics_types[0]239 )240 241 with gr.Tab("Data Generation"):242 with gr.Row():243 num_samples = gr.Slider(100, 5000, value=1000, label="Number of Samples", step=100)244 time_steps = gr.Slider(5, 100, value=50, label="Time Steps", step=5)245 gen_output = gr.Textbox(label="Output Log")246 generate_btn = gr.Button("Generate Data")247 generate_btn.click(248 fn=generate_training_data,249 inputs=[physics_dropdown, num_samples, time_steps],250 outputs=gen_output251 )252 253 gr.Markdown("#### ๐ค Upload Existing Dataset (.pkl)")254 upload_data = gr.File(file_types=[".pkl"], label="Upload .pkl")255 upload_data.upload(256 fn=load_uploaded_dataset,257 inputs=[upload_data, physics_dropdown],258 outputs=gen_output259 )260 261 gr.Markdown("#### ๐ฅ Download Generated Dataset")262 download_data_btn = gr.Button("Download Dataset")263 download_data_file = gr.File(label="Download Link")264 265 def return_dataset_path(physics_type):266 return resolve_path(f"{physics_type}_data.pkl", write_mode=True)267 268 download_data_btn.click(269 fn=return_dataset_path,270 inputs=[physics_dropdown],271 outputs=download_data_file272 )273 274 275 with gr.Tab("Data Inspection"):276 gr.Markdown("Visualize 3 samples from the training dataset to debug dot position.")277 278 inspect_btn = gr.Button("Show Sample Trajectories")279 output_fig = gr.Plot()280 281 inspect_btn.click(282 fn=inspect_training_trajectories,283 inputs=[physics_dropdown],284 outputs=output_fig285 )286 287 with gr.Tab("Training"):288 with gr.Row():289 epochs = gr.Slider(5, 100, value=20, label="Epochs", step=1)290 291 early_stopping_checkbox = gr.Checkbox(label="Enable Early Stopping", value=False)292 patience_slider = gr.Slider(1, 20, value=5, label="Patience Steps", step=1)293 294 train_output = gr.Textbox(label="Training Log")295 loss_plot = gr.Plot(label="Training Loss Curve")296 train_btn = gr.Button("Train Model")297 298 def run_training(physics_type, epochs, early_stopping, patience):299 msg, losses = train_model(300 physics_type=physics_type,301 epochs=epochs,302 early_stopping=early_stopping,303 patience=patience304 )305 fig, ax = plt.subplots()306 ax.plot(losses)307 ax.set_title("Training Loss Curve")308 ax.set_xlabel("Epoch")309 ax.set_ylabel("Loss")310 return msg, fig311 312 train_btn.click(313 fn=run_training,314 inputs=[physics_dropdown, epochs, early_stopping_checkbox, patience_slider],315 outputs=[train_output, loss_plot]316 )317 318 gr.Markdown("#### ๐ค Upload Trained Model (.pth)")319 upload_model = gr.File(file_types=[".pth"], label="Upload .pth")320 upload_model.upload(321 fn=load_uploaded_model,322 inputs=[upload_model, physics_dropdown],323 outputs=train_output324 )325 326 gr.Markdown("#### ๐ฅ Download Trained Model")327 download_model_btn = gr.Button("Download Model")328 download_model_file = gr.File(label="Download Link")329 330 def return_model_path(physics_type):331 return resolve_path(f"{physics_type}_model.pth", write_mode=True)332 333 download_model_btn.click(334 fn=return_model_path,335 inputs=[physics_dropdown],336 outputs=download_model_file337 )338 339 340 with gr.Tab("Input Sensitivity Test"):341 gr.Markdown("Compare how different inputs affect predicted trajectories.")342 sensitivity_btn = gr.Button("Run Sensitivity Test")343 output_plot = gr.Plot()344 345 sensitivity_btn.click(346 fn=test_input_sensitivity,347 inputs=[physics_dropdown],348 outputs=output_plot349 )350 351 with gr.Tab("Prediction"):352 with gr.Row():353 slider_outputs = [354 gr.Slider(visible=False),355 gr.Slider(visible=False),356 gr.Slider(visible=False)357 ]358 359 debug_checkbox = gr.Checkbox(label="Debug", value=False)360 predict_btn = gr.Button("Predict Trajectory")361 pred_plot = gr.Plot(label="Trajectory Prediction")362 pred_video = gr.Video(label="Generated Video")363 364 def refresh_inputs(physics_type):365 sliders = get_input_fields(physics_type) # Returns configured sliders366 updates = []367 param_ranges = get_param_ranges(physics_type) 368 for template, target in zip(sliders, slider_outputs):369 min_val, max_val = param_ranges[template]370 default_val = (min_val + max_val) / 2371 updates.append(gr.update(372 visible=True,373 label=template.replace('_', ' ').title(),374 minimum=min_val,375 maximum=max_val,376 value=default_val377 ))378 379 # Hide any unused sliders380 for _ in range(len(sliders), len(slider_outputs)):381 updates.append(gr.update(visible=False))382 383 return updates384 385 physics_dropdown.change(386 fn=refresh_inputs,387 inputs=physics_dropdown,388 outputs=slider_outputs389 )390 391 demo.load(fn=refresh_inputs, inputs=physics_dropdown, outputs=slider_outputs)392 393 def predict_switch(physics_type, *args):394 *slider_vals, debug = args395 return predict_plot_video(physics_type, *slider_vals, debug=debug)396 397 predict_btn.click(398 fn=predict_switch,399 inputs=[physics_dropdown] + slider_outputs + [debug_checkbox],400 outputs=[pred_plot, pred_video]401 )402 403if __name__ == "__main__":404 demo.launch()405 