Shortcuts

LightRFT 策略设计理念

概述

LightRFT 的策略(Strategy)模块为分布式训练策略提供了一个统一的接口,使得在保持 API 一致性的同时,能够无缝切换不同的分布式训练后端。本文档概述了策略模块的设计原则、架构和使用模式。

核心设计原则

1. 抽象与统一

原则:提供统一的接口,抽象掉不同分布式训练框架的复杂性。

实现

2. 配置驱动设计

原则:使用类型化的配置对象(Typed Configuration Objects)代替动态属性访问,以实现更好的类型安全性和代码清晰度。

实现

  • StrategyConfig 数据类提供了对所有配置参数的类型化访问。

  • 消除了对 getattr(args, "parameter", default) 模式的需求。

  • 支持 IDE 自动补全和静态类型检查。

3. 向后兼容性

原则:在引入改进的同时,保持与现有代码的兼容性。

实现

4. 可测试性

原则:无需分布式环境即可进行全面测试。

实现

  • FakeStrategy 提供了一个可用于测试的直接替代方案。

  • 所有策略方法都有用于单进程测试的 Mock 实现。

  • 单元测试验证功能和 API 的一致性。

架构

策略层次结构

StrategyBase (抽象基类)
├── DeepspeedStrategy (DeepSpeed 策略)
├── FSDPV2Strategy    (FSDP 策略)
└── FakeStrategy      (用于测试的伪策略)

关键组件

1. 策略工厂

get_strategy() 函数作为入口点,根据配置自动选择合适的策略:

from lightrft.strategy import get_strategy

# 根据 args.fsdp 自动选择 DeepSpeed 或 FSDP
strategy = get_strategy(args)

2. 配置管理

StrategyConfig 类集中管理所有配置参数:

from lightrft.strategy.config import StrategyConfig

config = StrategyConfig.from_args(args)
# 类型安全地访问参数
learning_rate = config.actor_learning_rate
use_bf16 = config.bf16

3. 通用接口

所有策略都实现相同的核心接口:

class StrategyBase(ABC):
    def setup_distributed(self, timeout=None) -> None: ...
    def create_optimizer(self, model, **kwargs) -> Optimizer: ...
    def prepare(self, *models, is_rlhf=False) -> Any: ...
    def backward(self, loss, model, optimizer, **kwargs) -> None: ...
    def optimizer_step(self, optimizer, model, scheduler, **kwargs) -> None: ...
    def save_ckpt(self, model, save_dir, **kwargs) -> None: ...
    def load_ckpt(self, model, load_dir, **kwargs) -> Any: ...

使用模式

1. 基本用法

from lightrft.strategy import get_strategy

# 初始化策略
strategy = get_strategy(args)

# 准备模型和优化器
actor, critic, reward_models, initial_model = strategy.prepare_models_and_optimizers(
    actor, critic, reward_models, initial_model, args, max_steps
)

# 训练循环
for batch in dataloader:
    loss = compute_loss(batch)
    strategy.backward(loss, actor, actor_optimizer)
    strategy.optimizer_step(actor_optimizer, actor, actor_scheduler)

2. 配置驱动用法

from lightrft.strategy.config import StrategyConfig

# 创建配置
config = StrategyConfig(
    seed=42,
    max_norm=1.0,
    micro_train_batch_size=4,
    train_batch_size=32,
    bf16=True,
    zero_stage=2
)

# 使用配置创建策略
strategy = get_strategy(config)

3. 使用 FakeStrategy 进行测试

from lightrft.strategy import get_fake_strategy

# 使用伪策略进行测试
strategy = get_fake_strategy()

# 所有操作无需分布式环境即可运行
strategy.setup_distributed()
strategy.backward(loss, model, optimizer)
strategy.save_ckpt(model, "checkpoints")

设计优势

1. 提高类型安全性

改进前 (使用 getattr):

seed = getattr(args, "seed", 42)  # 类型: Any
max_norm = getattr(args, "max_norm", 1.0)  # 类型: Any

改进后 (使用 StrategyConfig):

config = StrategyConfig.from_args(args)
seed = config.seed  # 类型: int
max_norm = config.max_norm  # 类型: float

2. 更好的代码组织

  • 配置参数在 StrategyConfig 中明确定义。

  • 策略特定的逻辑封装在具体的策略类中。

  • 通用功能在 StrategyBase 中实现。

3. 增强可测试性

  • FakeStrategy 使得无需分布式设置即可进行测试。

  • 单元测试可以验证所有策略功能。

  • Mock 实现确保了行为的一致性。

4. 未来的扩展性

  • 通过实现 StrategyBase 接口,可以轻松添加新策略。

  • 配置可以扩展而不破坏现有代码。

  • 工厂模式使得添加新策略类型变得非常简单。

最佳实践

1. 配置管理

  • 使用 StrategyConfig 进行所有参数访问。

  • 避免直接对参数对象调用 getattr

  • 对于非标准参数,使用 get_extra_arg()

2. 策略选择

  • 使用 get_strategy() 工厂函数创建策略。

  • 让工厂根据配置确定适当的策略。

  • 在开发和测试中使用 FakeStrategy

3. 错误处理

  • 策略应对不支持的操作提供清晰的错误消息。

  • 使用策略的 print() 方法进行日志记录。

  • 在上下文管理器中实现适当的清理。

4. 测试

  • 在单元测试中使用 FakeStrategy

  • 测试策略特定的功能和通用功能。

  • 验证所有策略是否都实现了所需的接口。

结论

LightRFT 策略模块优化了分布式训练的抽象设计,旨在提升 RLHF 系统的灵活性、类型安全与开发效率。通过统一抽象层与配置驱动开发,该模块实现了不同训练框架间的互操作性。

核心设计

  • 统一接口架构:封装 DeepSpeed、FSDP 等分布式后端,提供一致的 API。开发者无需修改业务代码即可切换底层策略。

  • 类型安全配置:通过 StrategyConfig 将动态配置转为强类型对象,减少运行时错误,并支持 IDE 自动补全。

  • 工厂模式选择get_strategy() 根据配置参数自动实例化策略,在简化调用的同时保留了对后端的控制权。

功能特性

  • 推理引擎集成:通过统一接口支持纯文本及多模态生成,兼容 vLLM 和 SGLang 后端。

  • 便捷测试支持:提供 FakeStrategy,允许在无分布式环境下测试训练工作流,降低调试成本。

  • 资源效率优化:支持推理引擎休眠/唤醒、梯度累积及内存感知检查点,优化大规模训练时的资源使用。

总结

该模块在易用性与灵活性之间取得了平衡。其架构设计兼顾了开发体验与性能需求,既降低了分布式训练的上手门槛,也为从 RLHF 到 RLVR 的多样化场景提供了可扩展的技术基础。