Team Ai
Apppublic

dizolivemint/motion-encoder-decoder

sourceHugging Facemitupdated 1y agoView on Hugging Face
0likes
generate_training_data.py94 linesDownload Raw Back to model_training
1import random2import pandas as pd3import numpy as np4import time5from config import get_input_fields, get_simulation_fn, normalize_input, get_param_ranges6from utils.path_utils import resolve_path7 8def generate_training_data(physics_type, num_samples=1000, time_steps=50):9    samples = []10    11    start_total = time.time()12    slowest = 0.013    worst_case = (0.0, 0.0, 0.0)  # Initialize with default values14    skipped = 015    16    param_specs = get_input_fields(physics_type)17    param_ranges = get_param_ranges(physics_type)18    sim_fn = get_simulation_fn(physics_type)19    20    for i in range(num_samples):21        t0 = time.time()22        23        # --- Randomize parameters based on config ---24        param_values = {25            field: random.uniform(*param_ranges[field])26            for field in param_specs27        }28 29        try:30            trajectory = sim_fn(**param_values, time_steps=time_steps)31        except Exception as e:32            print(f"⚠️ Skipping sample due to error: {e}")33            skipped += 134            continue35        36        dot_coords = []37        for frame in trajectory:38            coords = np.argwhere(frame > 0.5)39            if coords.size:40                y, x = coords[0]41                x_norm, y_norm = x / 63.0, y / 63.042                dot_coords.append((x_norm, y_norm))43            else:44                dot_coords.append((0.0, 0.0))45                46        # --- Debug ---47        # if i < 3:48        #     print(f"Sample {i} raw active coords:")49        #     for t, frame in enumerate(trajectory[:10]):50        #         coords = np.argwhere(frame > 0.5)51        #         if coords.size:52        #             print(f"  t={t} dot at {coords[0].tolist()}")53        #         else:54        #             print(f"  t={t} no dot")55 56 57        # Skip if ball stays mostly in one place58        unique_coords = set(dot_coords)59 60        if len(unique_coords) < 3:61            if i < 5:62                print(f"Skipping sample {i} due to mostly stationary dot")63            skipped += 164            continue65          66        # Collect values for metadata67        row_data = {field: param_values[field] for field in param_specs}68        row_data["trajectory"] = dot_coords69        samples.append(row_data)70        71        # --- Track performance ---72        t1 = time.time()73        dt = t1 - t074        if dt > slowest:75            slowest = dt76            worst_case = param_values77        78        # Debug79        # if i % 100 == 0:80        #     print(f"Sample {i}, time={dt:.4f}s")81 82    df = pd.DataFrame(samples)83    filename = resolve_path(f"{physics_type}_data.pkl", write_mode=True)84    df.to_pickle(filename)85    print(f"⏱️ Total time: {time.time() - start_total:.2f}s")86    if len(samples) > 0:87        print(f"⏱️ Worst-case time per sample: {slowest:.4f}s → " + 88              ", ".join(f"{k}={v:.2f}" for k, v in worst_case.items() if k != "trajectory"))89    print(f"✅ Saved {len(samples)} samples to {filename} (skipped {skipped} stationary ones)")90    return f"✅ Saved {len(samples)} samples to {filename}"91 92if __name__ == "__main__":93    generate_training_data()94