Skip to content

feat: add PrivateUse1 backend extension support - #206

Open
chen2021673 wants to merge 6 commits into
masterfrom
refactor/privateuse1-backend
Open

chen2021673 wants to merge 6 commits into
masterfrom
refactor/privateuse1-backend

Conversation

@chen2021673

@chen2021673 chen2021673 commented Aug 12, 2026

Copy link
Copy Markdown
Contributor

背景

InfiniTrain 原有设备体系只包含 CPU 和 CUDA。接入新后端时,需要在核心框架中增加厂商专属的设备枚举、runtime、CCL、kernel、测试和模型入口判断,导致核心代码与具体厂商耦合。

本 PR 引入通用的 DeviceType::kPrivateUse1 扩展槽位。外部 Provider 可以注册自己的运行时和算子实现,MACA 后端则通过独立仓库接入,公共 CMake 中不包含 MACA SDK 或厂商构建选项。

主要改动

PrivateUse1 注册接口

新增 PrivateUse1BackendRegistrationRegisterPrivateUse1Backend(),统一完成:

  • 注册厂商名称,例如 maca
  • 声明默认 autocast dtype,目前支持 FP16 或 BF16
  • 注册 DeviceGuardImpl
  • 注册 backend kernels
  • 可选注册 CclImpl
  • 校验 runtime、CCL 和基础 kernel 是否注册完整
  • 限制一个进程最多注册一个 PrivateUse1 Provider

Provider 名称仅允许小写 ASCII 字母、数字和下划线,且不能占用 cpucuda

继续复用 REGISTER_KERNELINFINI_TRAIN_REGISTER_DEVICE_GUARD_IMPLINFINI_TRAIN_REGISTER_CCL_IMPL

设备解析与模型入口

新增统一的 Device::ParseType()

  • cpu 映射到 kCPU
  • cuda 映射到 kCUDA
  • privateuse1 映射到 kPrivateUse1
  • 已注册的厂商名称映射到 kPrivateUse1

Device::ToString() 使用注册后的厂商名称。AutocastGuard 根据 Provider 注册信息取得 PrivateUse1 的默认计算类型,CPU 和 CUDA 的现有行为保持不变。

GPT-2、Llama 3 和 Mixtral 使用 Device::ParseType() 解析设备。外部工程可通过编译定义注入 Provider 头文件和注册入口,并在 gflags 校验 --device 前完成注册。并行模型使用用户选择的 accelerator backend,不再固定为 CUDA。

运行脚本新增 DEVICE_BACKEND,默认值仍为 cuda

GPT-2 和 Llama 3 暂时保留少量仅在 Provider 名称为 maca 时启用的同步和进程退出 workaround,并保留 FIXME;这些逻辑不会影响其他 PrivateUse1 Provider。

构建与静态注册

新增并明确以下 CMake 接口:

  • InfiniTrain::infini_train:供库和 Provider 使用的核心接口
  • InfiniTrain::cpu_kernels:CPU kernel target
  • InfiniTrain::infini_train_executable:最终可执行文件的完整链接接口

最终可执行文件通过 archive group 和 --whole-archive 保留 DeviceGuard、CCL 和 kernel 的静态注册对象。Provider 可以通过自己的 executable interface 或 EXTRA_ARCHIVES 加入 Provider 注册 archive。

InfiniTrain 作为 submodule 使用时不再构建自身 examples 和 tools,并隔离 glog 的测试选项,避免污染上层工程。

测试复用

公共测试的 suite 声明与设备实例化分离,每个测试二进制只实例化一个设备。CMake 为目标注入设备类型、GTest 前缀和 DEVICE_INDEX,其中设备序号默认是 0

新增 infini_train_add_privateuse1_test_suites(),允许外部 Provider 复用全量公共测试:

  • Provider 构建要求 USE_CUDA=OFF,冲突时直接报错
  • 保留 InfiniTrain 原有 CPU、fake Provider 和 CPU-only 测试
  • 追加 test_*_<BACKEND_NAME> Provider 测试
  • Provider 注册入口由共享 test_main 在测试环境初始化前显式调用
  • 删除 ONLY_CUDA,copy 等测试改为通用 accelerator 测试

Provider 测试使用厂商名作为 CTest label,例如 MACA 使用 ctest -L maca;不提供 ctest -L privateuse1 标签。

新增无需真实硬件的 fake PrivateUse1 测试,覆盖注册校验、名称解析、默认 autocast dtype、延迟 runtime 初始化和基础 kernel 调度。

Test

image image

@JYMiracle305
JYMiracle305 self-requested a review August 17, 2026 07:17
Comment thread infini_train/src/nn/parallel/ddp/distributed_data_parallel.cc
Comment thread infini_train/src/nn/parallel/ddp/reducer.cc
Comment thread infini_train/src/core/runtime/device_guard.cc
Comment thread CMakeLists.txt
Comment thread tests/backend/test_privateuse1_backend.cc Outdated
Comment thread infini_train/src/core/runtime/device_guard.cc Outdated
@chen2021673
chen2021673 force-pushed the refactor/privateuse1-backend branch 3 times, most recently from 4b1042c to 66cc91e Compare August 27, 2026 05:39
@chen2021673
chen2021673 changed the base branch from master to fix/backend-independent-correctness August 31, 2026 08:18
@chen2021673
chen2021673 force-pushed the refactor/privateuse1-backend branch from 66cc91e to 54dc847 Compare August 31, 2026 09:30
@chen2021673
chen2021673 force-pushed the refactor/privateuse1-backend branch from 54dc847 to e8a25ee Compare September 2, 2026 09:31
Base automatically changed from fix/backend-independent-correctness to master September 11, 2026 06:46
@chen2021673
chen2021673 force-pushed the refactor/privateuse1-backend branch from d9fe648 to 9b4a33c Compare September 14, 2026 02:36
Comment thread example/gpt2/main.cc Outdated
DEFINE_bool(overfit_single_batch, true, "overfit just one batch of data");
// memory management
DEFINE_string(device, "cuda", "device type (cpu/cuda), useless if using parallel training mode");
DEFINE_string(device, "cuda", "device type, useless if using parallel training mode");

Copy link
Copy Markdown
Collaborator

Choose a reason for hiding this comment

The reason will be displayed to describe this comment to others. Learn more.

cpu/cuda/privateuse1/<registered privateuse1 backend name>
改成这样吧,避免用户对大小写或合法取值产生误解。

Copy link
Copy Markdown
Contributor Author

Choose a reason for hiding this comment

The reason will be displayed to describe this comment to others. Learn more.

done

// rely solely on file-scope static initialization because an unreferenced static archive member may be discarded.
using PrivateUse1RegistrationCallback = void (*)();

struct PrivateUse1BackendRegistration {

Copy link
Copy Markdown
Collaborator

Choose a reason for hiding this comment

The reason will be displayed to describe this comment to others. Learn more.

如果走静态注册的话,这部分是不是不需要了?

Copy link
Copy Markdown
Contributor Author

Choose a reason for hiding this comment

The reason will be displayed to describe this comment to others. Learn more.

name和default_autocast_dtype应该还是需要,三个回调函数不需要了,我删掉。

Comment thread example/gpt2/main.cc
#ifdef INFINITRAIN_EXAMPLE_EXTERNAL_BACKEND_WORKAROUNDS
// FIXME(cx): MACA needs synchronization before training teardown.
if (INFINITRAIN_EXAMPLE_EXTERNAL_BACKEND_WORKAROUNDS && device.type() == Device::DeviceType::kPrivateUse1) {
impl->SynchronizeDevice(device);

Copy link
Copy Markdown
Collaborator

Choose a reason for hiding this comment

The reason will be displayed to describe this comment to others. Learn more.

这里麻烦再确认下 main.cc 里这两处 workrounds 是否仍然有必要:

  1. 沐曦平台改成多进程启动后,是否还会复现原问题;
  2. 沐曦在多线程分布式启动模式下,确认这两处是否仍然必要(考虑到框架经过一段时间迭代,可能有些 bug 被修复了)。

@chen2021673 chen2021673 Sep 17, 2026

Copy link
Copy Markdown
Contributor Author

Choose a reason for hiding this comment

The reason will be displayed to describe this comment to others. Learn more.

  1. 已验证多线程问题依然存在,多进程不会复现原问题,但需要一处新的修复
  2. 因为多线程分布式启动模式下问题依然存在,所以两处修复仍有必要,目前在多进程启动模式下单独关闭这两处 workrounds 修复及异步 malloc/free 修复,来观察性能数据。


TEST(PrivateUse1BackendValidationTest, RejectsNonAsciiBackendNames) {
auto registration = MinimalRegistration();
registration.name = "invalid-name";

Copy link
Copy Markdown
Collaborator

Choose a reason for hiding this comment

The reason will be displayed to describe this comment to others. Learn more.

这里测试的并不是 non-ASCII,而是不支持的 backend name 字符(如 -),建议改下测试名。

@chen2021673 chen2021673 Sep 17, 2026

Copy link
Copy Markdown
Contributor Author

Choose a reason for hiding this comment

The reason will be displayed to describe this comment to others. Learn more.

改为 RejectsUnsupportedBackendNameCharacters

- add a provider-neutral PrivateUse1 device type and registration API
- validate runtime, kernel, and optional CCL backend registrations
- initialize external device runtimes lazily on first use
- support provider names in device parsing and display
- require explicit autocast dtype for PrivateUse1 devices
- allow examples to register an external backend before flag parsing
- honor average_in_collective consistently across DDP paths
- expose embeddable CMake targets and add fake backend tests
- separate test suite declaration from CPU, CUDA, and provider instantiation
- support provider-injected registration, linkage, target names, and CTest labels
- add provider-defined default autocast dtype and backend-neutral test helpers
- generalize accelerator copy tests and remove the ineffective CUDA optimizer test
- fix Exp/Add backward, CUDA bias reduction, and DDP device validation
- scope runtime workarounds to the registered MACA backend
- document external backend test integration and usage
- group CCL initialization only for multiple local communicators
Register provider metadata explicitly while leaving runtime, kernel, and CCL
implementations in their statically initialized registries. Validate the
runtime and required kernels after metadata registration.
@chen2021673
chen2021673 force-pushed the refactor/privateuse1-backend branch 2 times, most recently from 99b1438 to e552299 Compare September 17, 2026 03:18
Sign up for free to join this conversation on GitHub. Already have an account? Sign in to comment

Labels

None yet

Projects

None yet

Development

Successfully merging this pull request may close these issues.

3 participants