Training Framework Module
RecIS’s training framework module provides comprehensive model training, evaluation, and management capabilities, simplifying the development workflow for deep learning models.
Core Components
TrainingArguments
Trainer
Saver
- class recis.framework.checkpoint_manager.Saver(options: SaverOptions)[source]
Checkpoint saver for managing model and training state persistence.
The Saver class handles the saving and loading of model checkpoints including: - Dense and sparse model parameters - Optimizer states - IO states for datasets - Checkpoint versioning and cleanup - Support for distributed filesystems
Example
>>> saver = Saver( ... model=model, ... sparse_optim=sparse_optimizer, ... output_dir="./checkpoints", ... max_keep=5, ... ) >>> saver.save("checkpoint_001")
- __init__(options: SaverOptions)[source]
Initialize the checkpoint saver.
- Parameters:
model (torch.nn.Module) – The model to save checkpoints for.
sparse_optim (Optional) – Sparse optimizer instance for sparse parameters.
output_dir (str) – Directory to save checkpoints. Defaults to “./”.
max_keep (int) – Maximum number of checkpoints to keep. Defaults to 1.
concurrency (int) – Number of concurrent save operations. Defaults to 4.
- load(ckpt_path: str | None = None, ckpt_id: str | None = None, direct_path=False, model_bank_conf: dict | None = None)[source]
根据传入的入参组合决定从哪里加载 ckpt, 真值表:
usage
branch
resolved ckpt_path
load()
by mode
see below
load(ckpt_id=”ckpt-100”)
by mode
see below
load(ckpt_path=”/x/y/”)
literal
“/x/y/”
load(ckpt_path=”/x/y/”, direct_path=1)
literal
“/x/y/” (same as above)
“by mode” 进一步分两种: - 标准 openlm_hub 用法:走 MOS 查
ckpt_id=None → MosCkptFileManager(version_uri, “r”) 拿最新
ckpt_id=”xxx” → MosCkptFileManager(version_uri/ckpt_id=xxx, “r”)
- 老协议:读 {output_dir}/checkpoint 索引文件
ckpt_id=None → 取索引最后一行
ckpt_id=”xxx” → 直接拼 {output_dir}/{ckpt_id}/
关键点:只要 caller 传了非空 ckpt_path, 就一律按字面路径加载,不会 被 MOS 查询或索引文件覆盖。这样 caller 的”我要加载这个具体路径”意图 永远不会被框架静默改写。
- register_for_checkpointing(name, obj: object)[source]
Register an object for checkpointing.
- Parameters:
- Raises:
ValueError – If the name is already registered.
- register_io_state(name, obj: object)[source]
Register an object for IO state persistence.
- Parameters:
- Raises:
ValueError – If the name is already registered.
- save(ckpt_id: str, label_key: str | None = None, label_value: str | None = None, sync_func: Callable | None = None)[source]
Save a complete checkpoint with the given ID.
流程:
save(ckpt_id) | +-- 1. 解析路径 + 文件系统 | +-- [openlm_hub] rank 0: helper.get_save_context() | | 其他 rank: broadcast 同步 ckpt_path | +-- [老协议] os.path.join(output_dir, ckpt_id) | +-- 2. makedirs(ckpt_path) | +-- 3. save_sparse_params() <- all rank | +-- 4. flush & save io_states <- all rank | +-- 5. if shard_id == 0: <- rank-0 only | +-- a. _save_rank0_states() 补空索引 / dense / extra 落盘 | +-- b. _update_ckpt_index() 写索引文件 + version_list | +-- c. _evict_old_ckpt() 淘汰旧 ckpt (len > max_keep 时) | +-- d. _register_ckpt() 注册 ckpt + 上报 metrics | +-- 6. cuda.synchronize + sync_func
ModelBankParser
TODO(lanling.ljw)
Exporter
Advanced Usage
Custom Training Pipeline
from framework.trainer import Trainer
class MyTrainer(Trainer):
def _train_step(self, data, epoch, metrics):
self.dense_optimizer.zero_grad()
if self.sparse_optimizer is not None:
self.sparse_optimizer.zero_grad()
loss = self.model(data)
metrics.update(epoch=epoch)
metrics.update(loss=loss)
metrics.update(get_global_metrics())
loss.backward()
self.dense_optimizer.step()
if self.sparse_optimizer is not None:
self.sparse_optimizer.step()
if self.dense_lr_scheduler is not None:
self.dense_lr_scheduler.step()
Gradient Accumulation Training
# Configure gradient accumulation
training_args = TrainingArguments(
output_dir="./output",
train_steps=10000,
gradient_accumulation_steps=8, # Accumulate 8 steps before update
log_steps=100
)
# Trainer will automatically handle gradient accumulation
trainer = Trainer(
model=model,
args=training_args,
train_dataset=train_dataset
)