feat: add record-mode to Store for DPO/GRPO

- Store gains fetch_record/num_records alongside stream fetch/__len__
- save_bin/load_bin support per-record offsets via record_keys param
- H5Store/MmapStore/JsonlStore all support dual stream+record access
- DPODataset/GRPODataset use fetch_record, no cross-record concat
- dpo_collate_fn + collate_fn wired through TrainConfig
- fixes attention context leakage in DPO from windowed concatenation
This commit is contained in:
2026-07-18 21:02:29 +08:00
parent 28886e4241
commit a74e5b91a3
12 changed files with 334 additions and 92 deletions
+8 -1
View File
@@ -8,7 +8,7 @@ import torch.optim as optim
from torch import Tensor, nn
from astrai.config import AutoRegressiveLMConfig, TrainConfig
from astrai.dataset import DatasetFactory
from astrai.dataset import DatasetFactory, dpo_collate_fn, grpo_collate_fn
from astrai.model import AutoRegressiveLM
from astrai.model.components.decoder_block import DecoderBlock
from astrai.trainer import SchedulerFactory, Trainer
@@ -504,6 +504,12 @@ def train(
grad_ckpt_modules = [DecoderBlock] if gradient_checkpointing else []
collate_fn = None
if train_type == "dpo":
collate_fn = dpo_collate_fn
elif train_type == "grpo":
collate_fn = grpo_collate_fn
train_config = TrainConfig(
model_fn=model_fn,
strategy=train_type,
@@ -536,6 +542,7 @@ def train(
executor_kwargs=executor_kwargs,
extra_kwargs=strategy_kwargs,
neftune_alpha=neftune_alpha,
collate_fn=collate_fn,
)
trainer = Trainer(train_config)