refactor: split train context build steps

- separate checkpoint, model, data, and strategy setup\n- keep build orchestration concise and readable
This commit is contained in:
2026-08-08 18:15:14 +08:00
parent d7cd69fef5
commit d9240ab149
+149 -129
View File
@@ -66,6 +66,15 @@ class TrainContext:
) )
@dataclass
class _PreloadedState:
model_config: dict = field(default_factory=dict)
state_dict: Optional[dict] = None
epoch: int = 0
consumed_samples: int = 0
checkpoint: Optional[Checkpoint] = None
class TrainContextBuilder: class TrainContextBuilder:
def __init__( def __init__(
self, self,
@@ -81,112 +90,149 @@ class TrainContextBuilder:
return self return self
def build(self) -> TrainContext: def build(self) -> TrainContext:
cfg = self.config # Resolve persisted state.
device = get_current_device() preloaded_state = self._load_preloaded_state()
executor = ExecutorFactory.create( # Build the core training components and restore their persisted state.
executor = self._create_executor()
context = self._create_context(preloaded_state, executor)
self._prepare_model(context, executor, preloaded_state)
self._restore_optimizer_state(context)
# Resolve datasets.
train_dataset, val_dataset = self._get_datasets()
self._create_dataloaders(context, train_dataset, val_dataset)
# Strategies depend on the prepared model; online rollout depends on both.
strategy_kwargs = self._create_strategy(context, executor)
self._configure_rollout(context, strategy_kwargs)
return context
def _create_executor(self) -> BaseExecutor:
cfg = self.config
return ExecutorFactory.create(
cfg.parallel_mode, cfg.parallel_mode,
grad_accum_steps=cfg.grad_accum_steps, grad_accum_steps=cfg.grad_accum_steps,
**cfg.executor_kwargs, **cfg.executor_kwargs,
) )
model_config = {} def _load_preloaded_state(self) -> _PreloadedState:
cfg = self.config
state = _PreloadedState(
epoch=cfg.start_epoch,
consumed_samples=cfg.start_samples * get_world_size(),
)
if self._param_path: if self._param_path:
config_path = Path(self._param_path) / "config.json" config_path = Path(self._param_path) / "config.json"
if config_path.exists(): if config_path.exists():
model_config = load_json(config_path) state.model_config = load_json(config_path)
preloaded_state_dict = None
preloaded_epoch = cfg.start_epoch
preloaded_consumed = cfg.start_samples * get_world_size()
preloaded_checkpoint = None
if self._param_path:
checkpoint = Checkpoint.load_any(self._param_path) checkpoint = Checkpoint.load_any(self._param_path)
if checkpoint is not None: if checkpoint is not None:
preloaded_state_dict = checkpoint.state_dict state.state_dict = checkpoint.state_dict
if checkpoint.config: state.model_config = checkpoint.config or state.model_config
model_config = checkpoint.config
if self._resume: if self._resume:
preloaded_epoch = checkpoint.epoch state.epoch = checkpoint.epoch
per_step = ( per_step = (
cfg.batch_per_device * get_world_size() * cfg.grad_accum_steps cfg.batch_per_device * get_world_size() * cfg.grad_accum_steps
) )
preloaded_consumed = ( state.consumed_samples = (
checkpoint.consumed_samples // per_step checkpoint.consumed_samples // per_step * per_step
) * per_step )
preloaded_checkpoint = checkpoint state.checkpoint = checkpoint
if not state.model_config and hasattr(cfg.model_fn(), "config"):
state.model_config = cfg.model_fn().config.to_dict()
return state
if not model_config and hasattr(cfg.model_fn(), "config"): def _create_context(
model_config = cfg.model_fn().config.to_dict() self, state: _PreloadedState, executor: BaseExecutor
) -> TrainContext:
return TrainContext(
world_size=get_world_size(),
rank=get_rank(),
config=self.config,
model_config=state.model_config,
executor=executor,
epoch=state.epoch,
consumed_samples=state.consumed_samples,
checkpoint=state.checkpoint,
)
def _before_wrap(m): def _prepare_model(
m = m.to(device=device) self, context: TrainContext, executor: BaseExecutor, state: _PreloadedState
) -> None:
cfg = self.config
device = get_current_device()
def before_wrap(model):
model = model.to(device=device)
if cfg.lora is not None: if cfg.lora is not None:
inject_lora( inject_lora(
m, model,
r=cfg.lora.r, r=cfg.lora.r,
alpha=cfg.lora.alpha, alpha=cfg.lora.alpha,
target_modules=set(cfg.lora.target_modules), target_modules=set(cfg.lora.target_modules),
) )
if preloaded_state_dict is not None: if state.state_dict is not None:
m.load_state_dict(preloaded_state_dict, strict=False) model.load_state_dict(state.state_dict, strict=False)
return m return model
def _after_wrap(m): def after_wrap(model):
if cfg.compile_mode is not None: if cfg.compile_mode is not None:
logger.info("torch.compile enabled (mode=%s)", cfg.compile_mode) logger.info("torch.compile enabled (mode=%s)", cfg.compile_mode)
m = torch.compile(m, mode=cfg.compile_mode) model = torch.compile(model, mode=cfg.compile_mode)
return m return model
context = TrainContext(
world_size=get_world_size(),
rank=get_rank(),
config=cfg,
model_config=model_config,
executor=executor,
epoch=preloaded_epoch,
consumed_samples=preloaded_consumed,
checkpoint=preloaded_checkpoint,
)
context.model, context.optimizer, context.scheduler = executor.prepare( context.model, context.optimizer, context.scheduler = executor.prepare(
cfg.model_fn, cfg.model_fn,
cfg.optimizer_fn, cfg.optimizer_fn,
cfg.scheduler_fn, cfg.scheduler_fn,
before_wrap=_before_wrap, before_wrap=before_wrap,
after_wrap=_after_wrap, after_wrap=after_wrap,
) )
train_dataset = cfg.dataset def _get_datasets(self):
val_dataset = cfg.val_dataset cfg = self.config
if cfg.val_dataset is not None or cfg.val_split is None:
if val_dataset is None and cfg.val_split is not None: return cfg.dataset, cfg.val_dataset
n_total = len(cfg.dataset) n_val = max(1, int(len(cfg.dataset) * cfg.val_split))
n_val = max(1, int(n_total * cfg.val_split))
n_train = n_total - n_val
generator = torch.Generator().manual_seed(cfg.random_seed) generator = torch.Generator().manual_seed(cfg.random_seed)
train_dataset, val_dataset = random_split( return random_split(
cfg.dataset, [n_train, n_val], generator=generator cfg.dataset, [len(cfg.dataset) - n_val, n_val], generator=generator
) )
def _create_dataloaders(
self, context: TrainContext, train_dataset, val_dataset
) -> None:
cfg = self.config
sampler_offset = context.consumed_samples // context.world_size sampler_offset = context.consumed_samples // context.world_size
if self._resume and sampler_offset > 0: if self._resume and sampler_offset > 0:
offset = context.world_size - 1 samples_per_replica = (
num_samples_per_replica = ( len(train_dataset) + context.world_size - 1
len(train_dataset) + offset
) // context.world_size ) // context.world_size
if num_samples_per_replica > 0: if samples_per_replica > 0:
context.epoch = sampler_offset // num_samples_per_replica context.epoch = sampler_offset // samples_per_replica
context.dataloader = self._create_dataloader(
sampler = RDSampler( train_dataset, context.epoch, sampler_offset
data_source=train_dataset,
start_epoch=context.epoch,
start_iter=sampler_offset,
seed=cfg.random_seed,
) )
context.dataloader = DataLoader( if val_dataset is not None:
train_dataset, context.val_dataloader = self._create_dataloader(
val_dataset, 0, 0, shuffle=False
)
def _create_dataloader(
self, dataset, epoch: int, start_iter: int, shuffle: bool = True
):
cfg = self.config
sampler = RDSampler(
dataset,
start_epoch=epoch,
start_iter=start_iter,
seed=cfg.random_seed,
shuffle=shuffle,
)
return DataLoader(
dataset,
batch_size=cfg.batch_per_device, batch_size=cfg.batch_per_device,
sampler=sampler, sampler=sampler,
num_workers=cfg.num_workers, num_workers=cfg.num_workers,
@@ -195,85 +241,60 @@ class TrainContextBuilder:
collate_fn=cfg.collate_fn, collate_fn=cfg.collate_fn,
) )
if val_dataset is not None: def _restore_optimizer_state(self, context: TrainContext) -> None:
val_sampler = RDSampler(
data_source=val_dataset,
start_epoch=0,
start_iter=0,
seed=cfg.random_seed,
shuffle=False,
)
context.val_dataloader = DataLoader(
val_dataset,
batch_size=cfg.batch_per_device,
sampler=val_sampler,
num_workers=cfg.num_workers,
pin_memory=cfg.pin_memory,
prefetch_factor=cfg.prefetch_factor,
collate_fn=cfg.collate_fn,
)
if context.checkpoint and context.checkpoint.extra: if context.checkpoint and context.checkpoint.extra:
extra = context.checkpoint.extra
for name in ("optimizer", "scheduler"): for name in ("optimizer", "scheduler"):
if name in extra: if (
obj = getattr(context, name, None) name in context.checkpoint.extra
if obj is not None: and getattr(context, name, None) is not None
obj.load_state_dict(extra[name]) ):
getattr(context, name).load_state_dict(
strategy_kwargs = dict(cfg.extra_kwargs) context.checkpoint.extra[name]
strategy_kwargs.setdefault("moe_aux_loss_coef", cfg.moe_aux_loss_coef)
needs_ref = cfg.strategy in (
"dpo",
"grpo",
"online_grpo",
"online_dpo",
)
needs_old = cfg.strategy in ("grpo", "online_grpo")
if needs_ref:
strategy_kwargs["ref_model"] = create_ref_model(
cfg.model_fn, executor=executor, model=context.model, device=device
) )
if needs_old: def _create_strategy(self, context: TrainContext, executor: BaseExecutor) -> dict:
strategy_kwargs["old_model"] = create_ref_model( cfg = self.config
cfg.model_fn, executor=executor, model=context.model, device=device kwargs = dict(cfg.extra_kwargs)
kwargs.setdefault("moe_aux_loss_coef", cfg.moe_aux_loss_coef)
if cfg.strategy in ("dpo", "grpo", "online_grpo", "online_dpo"):
kwargs["ref_model"] = create_ref_model(
cfg.model_fn,
executor=executor,
model=context.model,
device=get_current_device(),
)
if cfg.strategy in ("grpo", "online_grpo"):
kwargs["old_model"] = create_ref_model(
cfg.model_fn,
executor=executor,
model=context.model,
device=get_current_device(),
) )
context.strategy = StrategyFactory.create( context.strategy = StrategyFactory.create(
cfg.strategy, cfg.strategy,
model=context.model, model=context.model,
device=device, device=get_current_device(),
executor=executor, executor=executor,
**strategy_kwargs, **kwargs,
) )
return kwargs
# Enable online rollout when the train_type is an ``online_*`` variant. def _configure_rollout(self, context: TrainContext, strategy_kwargs: dict) -> None:
is_online = cfg.strategy.startswith("online_") cfg = self.config
if is_online: if not cfg.strategy.startswith("online_"):
return
if not context.strategy.supports_online(): if not context.strategy.supports_online():
raise ValueError( raise ValueError(
f"Strategy '{cfg.strategy}' does not support online rollout" f"Strategy '{cfg.strategy}' does not support online rollout"
) )
if cfg.reward_model_fn is None:
raise ValueError("reward_model_fn is required for online RL strategies")
tokenizer = AutoTokenizer.from_pretrained(self._param_path) tokenizer = AutoTokenizer.from_pretrained(self._param_path)
reward_model = cfg.reward_model_fn()
group_size = strategy_kwargs.get("group_size", 1) 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_position_embeddings", None)
scheduler = InferenceScheduler( scheduler = InferenceScheduler(
model=context.model, model=context.model,
tokenizer=tokenizer, tokenizer=tokenizer,
max_batch_size=rollout_batch_size, max_batch_size=group_size * max(1, cfg.batch_per_device),
max_seq_len=max_seq_len, max_seq_len=getattr(context.model.config, "max_position_embeddings", None),
) )
generator = RolloutGenerator( generator = RolloutGenerator(
scheduler=scheduler, scheduler=scheduler,
tokenizer=tokenizer, tokenizer=tokenizer,
@@ -283,11 +304,10 @@ class TrainContextBuilder:
top_k=cfg.rollout_top_k, top_k=cfg.rollout_top_k,
top_p=cfg.rollout_top_p, top_p=cfg.rollout_top_p,
) )
runner = RolloutRunner( context.strategy.set_rollout_runner(
RolloutRunner(
generator=generator, generator=generator,
reward_model=reward_model, reward_model=cfg.reward_model_fn(),
rollout_interval=cfg.rollout_interval, rollout_interval=cfg.rollout_interval,
) )
context.strategy.set_rollout_runner(runner) )
return context