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:
  • name (str) – Name identifier for the checkpointed object.

  • obj (object) – Object to include in checkpoints.

Raises:

ValueError – If the name is already registered.

register_io_state(name, obj: object)[source]

Register an object for IO state persistence.

Parameters:
  • name (str) – Name identifier for the IO state.

  • obj (object) – Object that supports IO state dump/load operations.

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
Parameters:
  • ckpt_id (str) – Unique identifier for this checkpoint.

  • label_key (str) – Key for the label when saving to MOS. Defaults to None.

  • label_value (str) – Value for the label when saving to MOS. Defaults to None.

  • sync_func (Callable) – Function to sync files. Defaults to None.

ModelBankParser

TODO(lanling.ljw)

class recis.framework.model_bank.ModelBankParser(output_dir: str, model_bank_content: list[Dict[str, Any]], model_names: set[str], sparse_model_names: set[str], sparse_tables: set[str], dense_model_names: set[str], extra_fields)[source]
__init__(output_dir: str, model_bank_content: list[Dict[str, Any]], model_names: set[str], sparse_model_names: set[str], sparse_tables: set[str], dense_model_names: set[str], extra_fields)[source]

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
)