479c2b1488
- CNNEncoder: stride=2 convs replace Conv2d+MaxPool2d pattern - 3 layers (32,64,128) instead of 4 (32,64,128,256), GRU input 192 - DecodeSample resizes grayscale frames to 80x60 via INTER_AREA - Model params: 227K (was 1.5M), input 80x60 (was 320x240) Generated by Mistral Vibe. Co-Authored-By: Mistral Vibe <vibe@mistral.ai>
125 lines
4.0 KiB
Python
125 lines
4.0 KiB
Python
"""
|
|
Global configuration for velocity prediction.
|
|
"""
|
|
|
|
from dataclasses import dataclass, field
|
|
from pathlib import Path
|
|
|
|
|
|
# ──────────────────────────── Dataset paths ────────────────────────────
|
|
|
|
DATASET_ROOT = Path(__file__).resolve().parents[2] / "dataset"
|
|
|
|
# Velocity normalization stats (computed from forward scenes only, 28363 frames)
|
|
# Yaw-compensated horizontal velocity: [v_right, v_forward]
|
|
VELOCITY_MEAN = [-2.902497, 3.837231] # [v_right, v_forward]
|
|
VELOCITY_STD = [3.453774, 3.722085] # [v_right, v_forward]
|
|
|
|
# TRAIN_SCENES = [
|
|
# "indoor_forward_3", "indoor_forward_5", "indoor_forward_6",
|
|
# "indoor_forward_7", "indoor_forward_9", "indoor_forward_10",
|
|
# "indoor_45_2", "indoor_45_4", "indoor_45_9", "indoor_45_12",
|
|
# ]
|
|
# VAL_SCENES = [
|
|
# "indoor_45_13", "indoor_45_14",
|
|
# ]
|
|
# TEST_SCENES = [
|
|
# "outdoor_forward_1", "outdoor_forward_3", "outdoor_forward_5",
|
|
# "outdoor_45_1",
|
|
# ]
|
|
|
|
TRAIN_SCENES = [
|
|
"indoor_forward_3", "indoor_forward_9", "indoor_forward_10", # Easy
|
|
"indoor_forward_5", "indoor_forward_6", # Medium
|
|
"outdoor_forward_3" # Medium 室外
|
|
]
|
|
VAL_SCENES = [
|
|
"indoor_forward_7", # Hard 室内
|
|
"outdoor_forward_1" # Easy 室外
|
|
# "indoor_forward_3", "indoor_forward_9", "indoor_forward_10", # Easy
|
|
]
|
|
TEST_SCENES = [
|
|
# "indoor_forward_9",
|
|
# "indoor_forward_9","indoor_forward_3",
|
|
# "indoor_forward_7", # Hard 室内
|
|
"outdoor_forward_1", # Easy 室外
|
|
# "outdoor_forward_5" # Hard 室外
|
|
# "indoor_forward_3", "indoor_forward_9", "indoor_forward_10", # Easy
|
|
]
|
|
|
|
|
|
# ──────────────────────────── Model architecture ────────────────────────────
|
|
|
|
@dataclass
|
|
class CNNConfig:
|
|
in_channels: int = 1
|
|
channels: tuple = (32, 64, 128) # per-layer output channels
|
|
kernel_size: int = 3
|
|
stride: int = 2 # strided conv replaces conv+pool
|
|
use_bn: bool = True
|
|
|
|
|
|
@dataclass
|
|
class PoseMLPConfig:
|
|
input_dim: int = 3
|
|
hidden_dim: int = 32
|
|
output_dim: int = 64
|
|
|
|
|
|
@dataclass
|
|
class GRUConfig:
|
|
input_size: int = 192 # CNN(128) + PoseMLP(64)
|
|
hidden_size: int = 128
|
|
num_layers: int = 1
|
|
dropout: float = 0.0
|
|
|
|
|
|
@dataclass
|
|
class HeadConfig:
|
|
input_dim: int = 128 # GRU hidden_size
|
|
hidden_dim: int = 64
|
|
output_dim: int = 2 # [v_right, v_forward]
|
|
|
|
|
|
@dataclass
|
|
class ModelConfig:
|
|
cnn: CNNConfig = field(default_factory=CNNConfig)
|
|
pose_mlp: PoseMLPConfig = field(default_factory=PoseMLPConfig)
|
|
gru: GRUConfig = field(default_factory=GRUConfig)
|
|
head: HeadConfig = field(default_factory=HeadConfig)
|
|
|
|
|
|
# ──────────────────────────── Training ────────────────────────────
|
|
|
|
@dataclass
|
|
class TrainConfig:
|
|
seq_len: int = 128 # frames per training sequence
|
|
batch_size: int = 4
|
|
epochs: int = 300
|
|
lr: float = 1e-3
|
|
weight_decay: float = 1e-5
|
|
lr_scheduler_step: int = 30
|
|
lr_scheduler_gamma: float = 0.5
|
|
num_workers: int = 4
|
|
seed: int = 42
|
|
|
|
# Sliding window: stride=1 → full overlap, stride=seq_len → non-overlapping
|
|
sliding_window_stride: int = 64
|
|
|
|
# Event simulation
|
|
event_threshold: float = 0.1
|
|
event_use_log: bool = True
|
|
event_auto_threshold: bool = False
|
|
|
|
# Logging / checkpoint
|
|
log_dir: str = "logs"
|
|
checkpoint_dir: str = "checkpoints"
|
|
log_interval: int = 10
|
|
save_interval: int = 10
|
|
|
|
|
|
# ──────────────────────────── Singleton instances ────────────────────────────
|
|
|
|
model_cfg = ModelConfig()
|
|
train_cfg = TrainConfig()
|