Files
uzh-fpv-sv-test/src/velocity_prediction/config.py
T
hexone2086 479c2b1488 refactor: strided conv encoder with 80x60 input resolution
- 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>
2026-08-01 17:47:45 +08:00

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()