Team Ai
Apppublic

dizolivemint/motion-encoder-decoder

sourceHugging Facemitupdated 1y agoView on Hugging Face
0likes
app.py405 linesDownload Raw Back to root
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