dizolivemint/motion-encoder-decoder
0
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 