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:
+168
-148
@@ -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:
|
||||||
|
return cfg.dataset, cfg.val_dataset
|
||||||
|
n_val = max(1, int(len(cfg.dataset) * cfg.val_split))
|
||||||
|
generator = torch.Generator().manual_seed(cfg.random_seed)
|
||||||
|
return random_split(
|
||||||
|
cfg.dataset, [len(cfg.dataset) - n_val, n_val], generator=generator
|
||||||
|
)
|
||||||
|
|
||||||
if val_dataset is None and cfg.val_split is not None:
|
def _create_dataloaders(
|
||||||
n_total = len(cfg.dataset)
|
self, context: TrainContext, train_dataset, val_dataset
|
||||||
n_val = max(1, int(n_total * cfg.val_split))
|
) -> None:
|
||||||
n_train = n_total - n_val
|
cfg = self.config
|
||||||
generator = torch.Generator().manual_seed(cfg.random_seed)
|
sampler_offset = context.consumed_samples // context.world_size
|
||||||
train_dataset, val_dataset = random_split(
|
if self._resume and sampler_offset > 0:
|
||||||
cfg.dataset, [n_train, n_val], generator=generator
|
samples_per_replica = (
|
||||||
|
len(train_dataset) + context.world_size - 1
|
||||||
|
) // context.world_size
|
||||||
|
if samples_per_replica > 0:
|
||||||
|
context.epoch = sampler_offset // samples_per_replica
|
||||||
|
context.dataloader = self._create_dataloader(
|
||||||
|
train_dataset, context.epoch, sampler_offset
|
||||||
|
)
|
||||||
|
if val_dataset is not None:
|
||||||
|
context.val_dataloader = self._create_dataloader(
|
||||||
|
val_dataset, 0, 0, shuffle=False
|
||||||
)
|
)
|
||||||
|
|
||||||
sampler_offset = context.consumed_samples // context.world_size
|
def _create_dataloader(
|
||||||
|
self, dataset, epoch: int, start_iter: int, shuffle: bool = True
|
||||||
if self._resume and sampler_offset > 0:
|
):
|
||||||
offset = context.world_size - 1
|
cfg = self.config
|
||||||
num_samples_per_replica = (
|
|
||||||
len(train_dataset) + offset
|
|
||||||
) // context.world_size
|
|
||||||
if num_samples_per_replica > 0:
|
|
||||||
context.epoch = sampler_offset // num_samples_per_replica
|
|
||||||
|
|
||||||
sampler = RDSampler(
|
sampler = RDSampler(
|
||||||
data_source=train_dataset,
|
dataset,
|
||||||
start_epoch=context.epoch,
|
start_epoch=epoch,
|
||||||
start_iter=sampler_offset,
|
start_iter=start_iter,
|
||||||
seed=cfg.random_seed,
|
seed=cfg.random_seed,
|
||||||
|
shuffle=shuffle,
|
||||||
)
|
)
|
||||||
context.dataloader = DataLoader(
|
return DataLoader(
|
||||||
train_dataset,
|
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,99 +241,73 @@ 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(
|
||||||
|
context.checkpoint.extra[name]
|
||||||
|
)
|
||||||
|
|
||||||
strategy_kwargs = dict(cfg.extra_kwargs)
|
def _create_strategy(self, context: TrainContext, executor: BaseExecutor) -> dict:
|
||||||
strategy_kwargs.setdefault("moe_aux_loss_coef", cfg.moe_aux_loss_coef)
|
cfg = self.config
|
||||||
|
kwargs = dict(cfg.extra_kwargs)
|
||||||
needs_ref = cfg.strategy in (
|
kwargs.setdefault("moe_aux_loss_coef", cfg.moe_aux_loss_coef)
|
||||||
"dpo",
|
if cfg.strategy in ("dpo", "grpo", "online_grpo", "online_dpo"):
|
||||||
"grpo",
|
kwargs["ref_model"] = create_ref_model(
|
||||||
"online_grpo",
|
cfg.model_fn,
|
||||||
"online_dpo",
|
executor=executor,
|
||||||
)
|
model=context.model,
|
||||||
needs_old = cfg.strategy in ("grpo", "online_grpo")
|
device=get_current_device(),
|
||||||
|
|
||||||
if needs_ref:
|
|
||||||
strategy_kwargs["ref_model"] = create_ref_model(
|
|
||||||
cfg.model_fn, executor=executor, model=context.model, device=device
|
|
||||||
)
|
)
|
||||||
|
if cfg.strategy in ("grpo", "online_grpo"):
|
||||||
if needs_old:
|
kwargs["old_model"] = create_ref_model(
|
||||||
strategy_kwargs["old_model"] = create_ref_model(
|
cfg.model_fn,
|
||||||
cfg.model_fn, executor=executor, model=context.model, device=device
|
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_"):
|
||||||
if not context.strategy.supports_online():
|
return
|
||||||
raise ValueError(
|
if not context.strategy.supports_online():
|
||||||
f"Strategy '{cfg.strategy}' does not support online rollout"
|
raise ValueError(
|
||||||
)
|
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)
|
|
||||||
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_position_embeddings", None)
|
|
||||||
|
|
||||||
scheduler = InferenceScheduler(
|
|
||||||
model=context.model,
|
|
||||||
tokenizer=tokenizer,
|
|
||||||
max_batch_size=rollout_batch_size,
|
|
||||||
max_seq_len=max_seq_len,
|
|
||||||
)
|
)
|
||||||
|
tokenizer = AutoTokenizer.from_pretrained(self._param_path)
|
||||||
generator = RolloutGenerator(
|
group_size = strategy_kwargs.get("group_size", 1)
|
||||||
scheduler=scheduler,
|
scheduler = InferenceScheduler(
|
||||||
tokenizer=tokenizer,
|
model=context.model,
|
||||||
max_tokens=cfg.rollout_max_tokens,
|
tokenizer=tokenizer,
|
||||||
group_size=group_size,
|
max_batch_size=group_size * max(1, cfg.batch_per_device),
|
||||||
temperature=cfg.rollout_temperature,
|
max_seq_len=getattr(context.model.config, "max_position_embeddings", None),
|
||||||
top_k=cfg.rollout_top_k,
|
)
|
||||||
top_p=cfg.rollout_top_p,
|
generator = RolloutGenerator(
|
||||||
)
|
scheduler=scheduler,
|
||||||
runner = RolloutRunner(
|
tokenizer=tokenizer,
|
||||||
|
max_tokens=cfg.rollout_max_tokens,
|
||||||
|
group_size=group_size,
|
||||||
|
temperature=cfg.rollout_temperature,
|
||||||
|
top_k=cfg.rollout_top_k,
|
||||||
|
top_p=cfg.rollout_top_p,
|
||||||
|
)
|
||||||
|
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
|
|
||||||
|
|||||||
Reference in New Issue
Block a user