perf(ml-alpha): cache loaded file across next_sequence calls

The MultiHorizonLoader was calling load_or_predecode_mbp10 on every
next_sequence() call, deserializing millions of MBP-10 snapshots per
sequence. With 8000 sequences and 9 files this gave ~50+ hour
training time on what should be IO-trivial work.

Now keep one file cached (LoadedFile { snapshots, labels_full }),
yield ceil(n_max_sequences / n_files) sequences from it before
advancing. Per-horizon labels are computed once per file load and
sliced cheaply per anchor. With 8000/9: ~890 loads → 9 loads.

Co-Authored-By: Claude Opus 4.7 <noreply@anthropic.com>
This commit is contained in:
jgrusewski
2026-05-17 00:19:06 +02:00
parent 4514313793
commit 586d1e7820

View File

@@ -48,12 +48,32 @@ pub struct LabeledSequence {
pub labels: [Vec<f32>; 5],
}
/// Loaded-file cache. Holds the deserialized snapshots + pre-computed
/// per-horizon labels so subsequent `next_sequence()` calls on the
/// same file avoid the multi-second bincode deserialization.
struct LoadedFile {
snapshots: Vec<Mbp10Snapshot>,
/// `labels_full[h]` holds the full-stream label vector for horizon
/// `cfg.horizons[h]`, indexed the same as `snapshots`. Sliced per
/// anchor on each `next_sequence()` call.
labels_full: [Vec<f32>; 5],
}
pub struct MultiHorizonLoader {
cfg: MultiHorizonLoaderConfig,
files: Vec<PathBuf>,
rng: ChaCha8Rng,
file_idx: usize,
yielded: usize,
/// Currently-loaded file. Cleared and replaced when we advance.
current: Option<LoadedFile>,
/// Sequences yielded from `current` since it was loaded.
yielded_from_current: usize,
/// How many sequences to yield from each file before advancing. Set
/// from `n_max_sequences / n_files` (rounded up) at construction, so
/// each file is loaded exactly once per "pass" through the loader.
sequences_per_file: usize,
}
impl MultiHorizonLoader {
@@ -75,12 +95,20 @@ impl MultiHorizonLoader {
);
let mut rng = ChaCha8Rng::seed_from_u64(cfg.seed);
files.shuffle(&mut rng);
let n_files = files.len();
// Yield (n_max_sequences / n_files) sequences per file, rounded
// up so we always honor n_max_sequences. With 8000 max and 9
// files: 8000/9 = 889/file, total 9 loads instead of 8000.
let sequences_per_file = cfg.n_max_sequences.div_ceil(n_files).max(1);
Ok(Self {
cfg: cfg.clone(),
files,
rng,
file_idx: 0,
yielded: 0,
current: None,
yielded_from_current: 0,
sequences_per_file,
})
}
@@ -96,64 +124,72 @@ impl MultiHorizonLoader {
if self.yielded >= self.cfg.n_max_sequences {
return Ok(None);
}
// Loop until we either yield or exhaust the file list (re-shuffle and
// try again if files are exhausted but we haven't hit n_max_sequences).
for _ in 0..self.files.len() * 2 {
if self.file_idx >= self.files.len() {
self.files.shuffle(&mut self.rng);
self.file_idx = 0;
}
let path = self.files[self.file_idx].clone();
self.file_idx += 1;
let max_horizon = *self.cfg.horizons.iter().max().expect("non-empty horizons");
let snapshots = load_or_predecode_mbp10(&path, &self.cfg.predecoded_dir)
.with_context(|| format!("load mbp10 {}", path.display()))?;
let max_horizon = *self.cfg.horizons.iter().max().expect("non-empty horizons");
if snapshots.len() < self.cfg.seq_len + max_horizon + 1 {
continue;
}
let max_anchor = snapshots.len() - self.cfg.seq_len - max_horizon;
let anchor: usize = self.rng.gen_range(0..max_anchor);
// Compute mid prices over the entire file (label generation needs
// the trailing window past the anchor + seq_len).
let prices: Vec<f32> = snapshots
.iter()
.map(|s| mid_price_f32(s))
.collect();
let mut labels: [Vec<f32>; 5] = Default::default();
for (h_idx, &h) in self.cfg.horizons.iter().enumerate() {
let mut out = vec![f32::NAN; self.cfg.seq_len];
let raw = generate_labels(&prices, h);
for (i, &t) in raw.valid_indices.iter().enumerate() {
if t >= anchor && t < anchor + self.cfg.seq_len {
out[t - anchor] = raw.labels[i];
}
// Advance to a new file if we have none cached or have yielded our
// budget from the current one.
let need_new = self.current.is_none()
|| self.yielded_from_current >= self.sequences_per_file;
if need_new {
// Try files until one is big enough; advance file_idx each time.
let mut found: Option<LoadedFile> = None;
for _ in 0..self.files.len() * 2 {
if self.file_idx >= self.files.len() {
self.files.shuffle(&mut self.rng);
self.file_idx = 0;
}
let path = &self.files[self.file_idx];
let snapshots = load_or_predecode_mbp10(path, &self.cfg.predecoded_dir)
.with_context(|| format!("load mbp10 {}", path.display()))?;
self.file_idx += 1;
if snapshots.len() < self.cfg.seq_len + max_horizon + 1 {
continue;
}
labels[h_idx] = out;
}
// Build sequence of Mbp10RawInput. Each input carries its own
// prev_mid + prev_ts_ns + trade_signed_vol derived from the
// PRIOR snapshot in the source stream (not from `anchor`).
let mut sequence = Vec::with_capacity(self.cfg.seq_len);
for k in 0..self.cfg.seq_len {
let idx = anchor + k;
let cur = &snapshots[idx];
let prev_idx = if idx == 0 { 0 } else { idx - 1 };
let prev = &snapshots[prev_idx];
sequence.push(convert(cur, prev));
// Precompute mid prices + per-horizon labels for the WHOLE file.
// We pay this once per file load, then slice cheaply per
// anchor on subsequent next_sequence calls.
let prices: Vec<f32> = snapshots.iter().map(mid_price_f32).collect();
let mut labels_full: [Vec<f32>; 5] = Default::default();
for (h_idx, &h) in self.cfg.horizons.iter().enumerate() {
let mut full = vec![f32::NAN; snapshots.len()];
let raw = generate_labels(&prices, h);
for (i, &t) in raw.valid_indices.iter().enumerate() {
full[t] = raw.labels[i];
}
labels_full[h_idx] = full;
}
found = Some(LoadedFile { snapshots, labels_full });
break;
}
self.current = found;
self.yielded_from_current = 0;
if self.current.is_none() {
return Ok(None); // no file in the inventory has enough snapshots
}
self.yielded += 1;
return Ok(Some(LabeledSequence {
snapshots: sequence,
labels,
}));
}
Ok(None) // exhausted files and none had enough snapshots
// Sample an anchor inside the current file and build the sequence.
let lf = self.current.as_ref().expect("set above");
let max_anchor = lf.snapshots.len() - self.cfg.seq_len - max_horizon;
let anchor: usize = self.rng.gen_range(0..max_anchor);
let mut labels: [Vec<f32>; 5] = Default::default();
for h in 0..5 {
labels[h] = lf.labels_full[h][anchor..anchor + self.cfg.seq_len].to_vec();
}
let mut sequence = Vec::with_capacity(self.cfg.seq_len);
for k in 0..self.cfg.seq_len {
let idx = anchor + k;
let cur = &lf.snapshots[idx];
let prev_idx = if idx == 0 { 0 } else { idx - 1 };
let prev = &lf.snapshots[prev_idx];
sequence.push(convert(cur, prev));
}
self.yielded += 1;
self.yielded_from_current += 1;
Ok(Some(LabeledSequence { snapshots: sequence, labels }))
}
}