feat(generalization): domain randomization — per-epoch jitter on sim params
Prevents memorization of fixed simulation parameters that exist only in training. Each epoch randomizes: - tx_cost_multiplier: U[0.5, 2.5] (was fixed 1.0) - spread: U[0.5x, 3.0x] base spread - fill_ioc_prob: U[0.65, 0.95] (was fixed 0.85) - fill_limit_min/max: randomized ranges - spread_cost: U[0.3, 0.8] fraction - Episode starts: ±25% stride jitter (was deterministic) - Episode length: ±25% base (was fixed) Controlled by enable_domain_randomization flag (default: true). Co-Authored-By: Claude Opus 4.6 (1M context) <noreply@anthropic.com>
This commit is contained in:
@@ -596,7 +596,8 @@ impl DQNAgentType {
|
||||
/// Get state dimension from agent configuration
|
||||
///
|
||||
/// WAVE 10.4: Added to fix hardcoded STATE_DIM bug
|
||||
/// Returns the actual state dimension (45 without OFI: 42 market + 3 portfolio, or 53 with OFI: +8 OFI)
|
||||
/// Returns the actual state dimension (72 without OFI, 80 with OFI)
|
||||
/// Layout: 42 market + 8 portfolio + 16 multi-timeframe + (8 OFI), 8-aligned
|
||||
pub fn get_state_dim(&self) -> usize {
|
||||
match self {
|
||||
Self::Standard(agent) => agent.get_state_dim(),
|
||||
@@ -926,6 +927,10 @@ pub struct DQNHyperparameters {
|
||||
/// Enable circuit breaker (5-failure trip mechanism)
|
||||
pub enable_circuit_breaker: bool,
|
||||
|
||||
/// Domain randomization: per-epoch jitter on spread, tx_cost, fills, episode starts.
|
||||
/// Prevents memorization of fixed simulation parameters.
|
||||
pub enable_domain_randomization: bool,
|
||||
|
||||
// Wave 16 Portfolio Features
|
||||
/// Enable action masking (filters invalid actions based on position limits)
|
||||
pub enable_action_masking: bool,
|
||||
@@ -1406,6 +1411,9 @@ impl DQNHyperparameters {
|
||||
enable_position_limits: true, // Default: 3-tier position limits enabled
|
||||
enable_circuit_breaker: true, // Default: circuit breaker enabled (5-failure trip)
|
||||
|
||||
// Generalization: domain randomization (per-epoch jitter on sim params)
|
||||
enable_domain_randomization: true, // Default: enabled (prevents memorization of fixed sim params)
|
||||
|
||||
// Wave 16 Portfolio Features (default: ALL ENABLED)
|
||||
enable_action_masking: true, // Default: action masking enabled
|
||||
enable_entropy_regularization: true, // Default: entropy regularization enabled
|
||||
|
||||
@@ -928,12 +928,33 @@ impl DQNTrainer {
|
||||
computed
|
||||
};
|
||||
|
||||
let timesteps = self.hyperparams.gpu_timesteps_per_episode.min(1000) as i32;
|
||||
// Domain randomization: variable episode length per epoch
|
||||
let dr = self.hyperparams.enable_domain_randomization;
|
||||
use rand::Rng;
|
||||
let mut epoch_rng = rand::thread_rng();
|
||||
let timesteps = {
|
||||
let base = self.hyperparams.gpu_timesteps_per_episode.min(1000) as i32;
|
||||
if dr {
|
||||
let jitter = epoch_rng.gen_range(-(base / 4)..(base / 4));
|
||||
(base + jitter).max(50)
|
||||
} else {
|
||||
base
|
||||
}
|
||||
};
|
||||
let total_bars = training_data.len() as i32;
|
||||
let usable_bars = (total_bars - timesteps).max(1);
|
||||
let stride = (usable_bars / n_episodes).max(1);
|
||||
// Domain randomization: jittered episode starts (±25% stride)
|
||||
let mut episode_starts: Vec<i32> = (0..n_episodes)
|
||||
.map(|i| (i * stride).rem_euclid(usable_bars))
|
||||
.map(|i| {
|
||||
let base = (i * stride).rem_euclid(usable_bars);
|
||||
if dr {
|
||||
let jitter = epoch_rng.gen_range(-(stride / 4)..(stride / 4));
|
||||
(base + jitter).rem_euclid(usable_bars)
|
||||
} else {
|
||||
base
|
||||
}
|
||||
})
|
||||
.collect();
|
||||
|
||||
// ── Curriculum learning: filter episode starts by difficulty ──────
|
||||
@@ -1014,7 +1035,11 @@ impl DQNTrainer {
|
||||
enable_action_masking: self.enable_action_masking,
|
||||
curiosity_scale: if self.curiosity_module.is_some() { 1.0 } else { 0.0 },
|
||||
loss_aversion: self.hyperparams.loss_aversion as f32,
|
||||
tx_cost_multiplier: self.hyperparams.transaction_cost_multiplier as f32,
|
||||
tx_cost_multiplier: if dr {
|
||||
epoch_rng.gen_range(0.5_f32..2.5)
|
||||
} else {
|
||||
self.hyperparams.transaction_cost_multiplier as f32
|
||||
},
|
||||
count_bonus_coefficient: self.hyperparams.count_bonus_coefficient
|
||||
.unwrap_or(0.0) as f32,
|
||||
q_clip_min: self.hyperparams.q_clip_min as f32,
|
||||
@@ -1030,11 +1055,27 @@ impl DQNTrainer {
|
||||
num_atoms: self.hyperparams.num_atoms as i32,
|
||||
v_min: self.hyperparams.v_min as f32,
|
||||
v_max: self.hyperparams.v_max as f32,
|
||||
fill_median_spread: self.hyperparams.avg_spread as f32,
|
||||
fill_median_spread: if dr {
|
||||
self.hyperparams.avg_spread as f32 * epoch_rng.gen_range(0.5_f32..3.0)
|
||||
} else {
|
||||
self.hyperparams.avg_spread as f32
|
||||
},
|
||||
fill_median_vol: self.median_vol as f32,
|
||||
fill_ioc_fill_prob: self.hyperparams.fill_ioc_fill_prob as f32,
|
||||
fill_limit_fill_min: self.hyperparams.fill_limit_fill_min as f32,
|
||||
fill_limit_fill_max: self.hyperparams.fill_limit_fill_max as f32,
|
||||
fill_ioc_fill_prob: if dr {
|
||||
epoch_rng.gen_range(0.65_f32..0.95)
|
||||
} else {
|
||||
self.hyperparams.fill_ioc_fill_prob as f32
|
||||
},
|
||||
fill_limit_fill_min: if dr {
|
||||
epoch_rng.gen_range(0.15_f32..0.45)
|
||||
} else {
|
||||
self.hyperparams.fill_limit_fill_min as f32
|
||||
},
|
||||
fill_limit_fill_max: if dr {
|
||||
epoch_rng.gen_range(0.60_f32..0.95)
|
||||
} else {
|
||||
self.hyperparams.fill_limit_fill_max as f32
|
||||
},
|
||||
fill_spread_cost_frac: self.hyperparams.fill_spread_cost_frac as f32,
|
||||
fill_spread_capture_frac: self.hyperparams.fill_spread_capture_frac as f32,
|
||||
fill_simulation_enabled: self.median_vol > 0.0,
|
||||
@@ -1052,7 +1093,13 @@ impl DQNTrainer {
|
||||
min_hold_bars: self.hyperparams.min_hold_bars as i32,
|
||||
// spread_cost = tick_size * multiplier * fraction
|
||||
// Matches backtest_env_kernel's spread_cost from GpuBacktestConfig
|
||||
spread_cost: (self.hyperparams.tick_size * self.hyperparams.contract_multiplier * self.hyperparams.fill_spread_cost_frac) as f32,
|
||||
spread_cost: if dr {
|
||||
(self.hyperparams.tick_size * self.hyperparams.contract_multiplier
|
||||
* epoch_rng.gen_range(0.3..0.8) as f64) as f32
|
||||
} else {
|
||||
(self.hyperparams.tick_size * self.hyperparams.contract_multiplier
|
||||
* self.hyperparams.fill_spread_cost_frac) as f32
|
||||
},
|
||||
contract_multiplier: self.hyperparams.contract_multiplier as f32,
|
||||
margin_pct: self.hyperparams.margin_pct as f32,
|
||||
dd_threshold: self.hyperparams.dd_threshold as f32,
|
||||
|
||||
Reference in New Issue
Block a user