fix: use batch_per_device for rollout scheduler batch sizing

- train_context.py referenced non-existent cfg.batch_size, replaced with cfg.batch_per_device
- default group_size lowered from 8 to 1: without a group concept (DPO), scheduler batch equals batch_per_device; rollout-based DPO can opt in via extra_kwargs['group_size']>=2
- inline expressions (rollout_batch_size, max_seq_len) extracted for readability
- add tests/trainer/test_online_e2e.py: end-to-end online_dpo via Trainer.train, exercising KV-cache-backed rollout path
This commit is contained in:
2026-07-20 13:32:04 +08:00
parent a6e1f26cd4
commit e8ff7f5321
2 changed files with 120 additions and 5 deletions
+8 -5
View File
@@ -238,20 +238,23 @@ class TrainContextBuilder:
tokenizer = AutoTokenizer.from_pretrained(self._param_path)
reward_model = cfg.reward_model_fn()
group_size = strategy_kwargs.get("group_size", 1)
rollout_batch_size = group_size * max(1, cfg.batch_per_device)
max_seq_len = getattr(context.model.config, "max_len", None)
scheduler = InferenceScheduler(
model=context.model,
tokenizer=tokenizer,
max_batch_size=strategy_kwargs.get("group_size", 8)
* max(1, cfg.batch_size or 1),
max_seq_len=getattr(context.model.config, "max_len", None),
max_prompt_len=getattr(context.model.config, "max_len", 4096),
max_batch_size=rollout_batch_size,
max_seq_len=max_seq_len,
max_prompt_len=max_seq_len or 4096,
)
generator = RolloutGenerator(
scheduler=scheduler,
tokenizer=tokenizer,
max_tokens=cfg.rollout_max_tokens,
group_size=strategy_kwargs.get("group_size", 8),
group_size=group_size,
temperature=cfg.rollout_temperature,
top_k=cfg.rollout_top_k,
top_p=cfg.rollout_top_p,