Ascend-SACT/sam3_910B
模型介绍文件和版本Pull Requests讨论分析
下载使用量0

SAM3模型迁移指导文档

输出结构:SAM3模型迁移指导文档/ → SAM3模型迁移指导文档/SAM3模型迁移指导文档.md → SAM3模型迁移指导文档/SAM3模型迁移指导文档_assets/{images,files}/。 每个小节中若存在多个操作步骤,统一使用有序编号 1.、2.、3. 表达。

1. 模型概述及场景

1.1 模型介绍

SAM3(Segment Anything Model 3)是Meta AI Research开发的通用图像分割模型,支持交互式分割、目标检测和视频追踪等任务。SAM3在SAM/SAM2的基础上引入了文本-视觉双编码器架构(ViT + Text Encoder),结合RoPE(Rotary Position Embedding)和TwoWayTransformer实现prompt-to-object对齐。模型原始运行环境为GPU(CUDA)。

1.2 应用场景

用于图像与视频的交互式分割、开放词汇目标检测、Roboflow 100-VL / ODinW13等下游微调任务。本文档聚焦NPU训练迁移与PyTorch PT模式推理迁移场景。

2. 迁移环境

2.1 硬件环境

  • GPU基线:A100 80G x 1
  • NPU目标:Ascend 910B3 x 8(单卡训练指定NPU:6)
  • 存储:本地NVMe
  • HBM:64 GB / 卡

2.2 软件环境

项目版本 / 规格
操作系统 / 架构Linux / aarch64
驱动 / 固件25.2.2 / TODO(固件版本)
CANN8.1.RC1(ascend-toolkit + kernels)
Python3.10.14(编译安装至 /usr/local/python3.10/)
torch / torch_npu2.1.0 / 2.1.0.post11

3. 资源与依赖

3.1 模型权重

SAM3预训练权重文件(sam3.pt,约3.3GB),放置于 /home/sam3.pt。在训练配置中通过checkpoint_path参数指定。

3.2 数据集

  1. Roboflow 100-VL(官方推荐):https://github.com/roboflow/rf100-vl(需Roboflow API Key)
  2. COCO 2017验证集(替代方案):可用COCO 2017验证集进行smoke test验证训练流程,目录结构如下:
coco2017/
  val2017/          ← 图片文件夹(5000张)
  annotations/
    instances_val2017.json  ← 标注文件

3.3 代码仓

https://github.com/facebookresearch/sam3(基线 commit: 8e451d5)

3.4 容器镜像资源

TODO(待补充官方或自建镜像)

3.5 工具资源

npu-smi、atc、set_env.sh(CANN 环境配置)

4. 环境准备

4.1 资源准备

  1. 从 gitcode.com 镜像克隆 SAM3 源码。
  2. 安装 CANN 8.1 RC1 Toolkit + Kernels。
  3. 编译安装 Python 3.10(CANN 8.1.RC1 要求 Python 3.7-3.10)。

克隆 SAM3 源码:

git clone https://gitcode.com/gh_mirrors/sam3.git /root/sam3

4.2 环境创建

  1. 设置 CANN 环境变量。
  2. 设置 NPU 驱动 LD_LIBRARY_PATH。
  3. 设置 Python 3.10 为默认 Python。

每次运行前需执行以下基础环境配置(训练和推理通用):

source /usr/local/Ascend/ascend-toolkit/ascend-toolkit/set_env.sh
export LD_LIBRARY_PATH=/usr/local/Ascend/driver/lib64:/usr/local/Ascend/driver/lib64/common:/usr/local/Ascend/driver/lib64/driver:$LD_LIBRARY_PATH
export PATH=/usr/local/python3.10/bin:$PATH

训练场景额外配置:

export PYTHONPATH=/root/sam3:$PYTHONPATH
export ASCEND_RT_VISIBLE_DEVICES=4,5,6,7  # 多卡训练指定4卡;单卡训练指定1卡,如=6

纯 OM 推理场景额外配置(不需要 PYTHONPATH,不需要 torch_npu):

source /usr/local/Ascend/ascend-toolkit/ascend-toolkit/8.1.RC1/bin/setenv.bash
export LD_LIBRARY_PATH=/usr/local/Ascend/driver/lib64:/usr/local/Ascend/driver/lib64/common:/usr/local/Ascend/driver/lib64/driver:$LD_LIBRARY_PATH

4.3 依赖安装

  1. 安装 PyTorch CPU 版本(aarch64)及 torch_npu。
  2. 安装 SAM3 核心依赖。
  3. 安装 SAM3 训练依赖。
  4. 降级 numpy 至 1.x(torch 2.1.0 要求)。

安装 PyTorch 和 torch_npu:

pip3.10 install torch==2.1.0 --no-deps
pip3.10 install -i https://mirrors.aliyun.com/pypi/simple/ --trusted-host mirrors.aliyun.com filelock typing-extensions sympy networkx fsspec jinja2 "numpy<2"
pip3.10 install -i https://mirrors.aliyun.com/pypi/simple/ --trusted-host mirrors.aliyun.com torch-npu==2.1.0.post11 pyyaml

安装 SAM3 核心及训练依赖:

pip3.10 install -i https://mirrors.aliyun.com/pypi/simple/ --trusted-host mirrors.aliyun.com timm tqdm ftfy regex iopath huggingface_hub hydra-core submitit tensorboard zstandard torchmetrics fvcore fairscale scikit-image scikit-learn opencv-python-headless pillow torchvision==0.16.0 einops pycocotools

降级 numpy:

pip3.10 install -i https://mirrors.aliyun.com/pypi/simple/ --trusted-host mirrors.aliyun.com "numpy<2"

安装 CANN TBE 编译器所需的 Python 模块:

pip3.10 install -i https://mirrors.aliyun.com/pypi/simple/ --trusted-host mirrors.aliyun.com decorator scipy attr attrs psutil astunparse cffi pathlib2 protobuf

5. 训练迁移步骤

5.1 数据准备

方案A:使用 COCO 2017 验证集(smoke test 推荐)

  1. 下载 COCO 2017 验证集图片和标注。
  2. 解压到指定目录。
mkdir -p /root/coco2017
python3.10 -c "import zipfile; zipfile.ZipFile('val2017.zip').extractall('/root/coco2017/')"
python3.10 -c "import zipfile; zipfile.ZipFile('annotations_trainval2017.zip').extractall('/root/coco2017/')"

数据目录结构:

/root/coco2017/
  val2017/          ← 图片文件夹
  annotations/
    instances_val2017.json  ← 只需此文件

方案B:使用 Roboflow 100-VL(正式训练)

  1. 安装 rf100vl pip 包。
  2. 设置 Roboflow API Key。
  3. 下载 Roboflow 100-VL 数据集。
pip3.10 install rf100vl
export ROBOFLOW_API_KEY=TODO(用户提供)
python3.10 -c "from rf100vl import download_rf100vl; download_rf100vl(path='./rf100-vl/')"

5.2 配置修改

SAM3 源码存在多处 GPU/CUDA 硬编码,需进行 NPU 兼容性适配。以下为已完成的代码修改清单:

  1. torch.nn.attention 兼容适配层:torch 2.1.0 缺少该模块,创建兼容替代。

    适配文件:torch_nn_attention_compat.py

    from contextlib import contextmanager
    from enum import Enum
    
    class SDPBackend(Enum):
        MATH = 0
        FLASH_ATTENTION = 1
        EFFICIENT_ATTENTION = 2
        CUDNN_ATTENTION = 3
    
    @contextmanager
    def sdpa_kernel(*backends):
        yield

    在 sam3/model/decoder.py 和 sam3/model/vl_combiner.py 中修改导入:

    try:
        from torch.nn.attention import sdpa_kernel, SDPBackend
    except ImportError:
        from sam3.compat.torch_nn_attention import sdpa_kernel, SDPBackend
  2. pytree.register_pytree_node 兼容性处理:torch 2.1.0 使用 _register_pytree_node。

    在 sam3/model/data_misc.py 中修改:

    _register_fn = getattr(pytree, 'register_pytree_node', None) or getattr(pytree, '_register_pytree_node')
    _register_fn(NestedTensor, ...)
  3. EDT 算子 NPU 替代实现:triton 为 CUDA 专用,用 scipy 替代。

    适配文件:edt_npu.py

    在 sam3/model/sam3_tracker_utils.py 中修改导入:

    try:
        from sam3.model.edt import edt_triton
    except ImportError:
        from sam3.model.edt_npu import edt_triton
  4. device="cuda" 硬编码替换:将多处 "cuda" 改为动态设备选择。

    设备工具:device_utils.py

    在 sam3/model/position_encoding.py 中修改:

    import sam3.compat
    tensors = torch.zeros((1, 1) + size, device=sam3.compat.get_device())
  5. 训练 YAML 配置修改:需修改以下字段。

    • trainer.accelerator → cuda 改为 npu
    • trainer.distributed.backend → nccl 改为 hccl
    • trainer.model.device → cuda 改为 npu
    • trainer.model.load_from_HF → 设置为 false(本地加载权重)
    • trainer.model.checkpoint_path → 设置为 /home/sam3.pt
    • submitit.use_cluster → 设置为 False(本地执行)
    • launcher.gpus_per_node → 设置为 1(单卡)
    • 数据集路径 → 指向实际数据目录
  6. NPU 设备初始化:在 sam3/train/trainer.py 的 _setup_device 方法中增加 NPU 分支。

    elif accelerator == "npu":
        self.device = torch.device("npu", self.local_rank)
        torch.npu.set_device(self.local_rank)

    注意:CANN 8.1.RC1 已安装预编译 kernels(8869 个 .o 文件),不再需要 torch.npu.set_compile_mode(jit_compile=True)。如果使用 CANN 8.0.RC2,则需要开启 JIT 编译模式。

  7. RoPE 实数运算替代:NPU 不支持在 NPU 上直接创建 Complex dtype 张量,需将 torch.polar(复数)改为 torch.cos/torch.sin(实数)。

    注意:CANN 8.1.RC1 下 torch.randn(..., dtype=torch.complex64, device='npu:0') 仍会报错,Complex 张量只能在 CPU 上创建后通过 view_as_real 转为 float32 再移到 NPU。因此 use_rope_real=True 解决方案在所有 CANN 版本中均需保留。

    在 sam3/sam/rope.py 的 compute_axial_cis 中增加 use_real=True 参数:

    if use_real:
        freqs_cis_x_real = torch.cos(freqs_x)
        freqs_cis_x_imag = torch.sin(freqs_x)
        freqs_cis_y_real = torch.cos(freqs_y)
        freqs_cis_y_imag = torch.sin(freqs_y)
        freqs_cis_real = torch.cat([freqs_cis_x_real, freqs_cis_y_real], dim=-1)
        freqs_cis_imag = torch.cat([freqs_cis_x_imag, freqs_cis_y_imag], dim=-1)
        return None, freqs_cis_real, freqs_cis_imag

    在 sam3/model_builder.py 中设置全局变量自动启用:

    _USE_ROPE_REAL = sam3.compat.is_npu_available()

    涉及文件:sam3/sam/rope.py、sam3/model/decoder.py(RoPEAttention)、sam3/sam/transformer.py(RoPEAttention)、sam3/model/vitdet.py(Attention _setup_rope_freqs 和 compute_axial_cis)

  8. GradScaler NPU 适配:torch 2.1.0 CPU 版本无 torch.amp.GradScaler。

    在 sam3/train/trainer.py 中根据设备类型选择:

    if sam3.compat.is_npu_available():
        self.scaler = torch.npu.amp.GradScaler(enabled=...)
    elif torch.cuda.is_available():
        self.scaler = torch.amp.GradScaler(self.device, enabled=...)
  9. triton sigmoid_focal_loss 兼容处理:NPU 无 triton,需条件导入并走纯 PyTorch 路径。

    在 sam3/train/loss/loss_fns.py 中修改:

    try:
        from sam3.train.loss.sigmoid_focal_loss import triton_sigmoid_focal_loss, triton_sigmoid_focal_loss_reduce
    except ImportError:
        triton_sigmoid_focal_loss = None
        triton_sigmoid_focal_loss_reduce = None

    并在 sigmoid_focal_loss 函数中:if triton and triton_sigmoid_focal_loss is not None:

  10. Trainer NPU 加速器支持:在 sam3/train/trainer.py 中增加 NPU 分支。

    • _setup_device:增加 accelerator == "npu" 分支
    • _infer_distributed_backend_if_none:NPU 使用 hccl 后端
    • _setup_ddp_distributed_training:NPU 传入 device_ids
    • torch.amp.autocast:device_type 改为 sam3.compat.get_device_type_for_autocast()
    • torch.cuda.is_available() → sam3.compat.is_accelerator_available()
    • torch.cuda.empty_cache() → sam3.compat.empty_cache()
  11. multiprocessing set_start_method 兼容处理:重复调用会抛 RuntimeError。

    在 sam3/train/train.py 中修改:

    try:
        torch.multiprocessing.set_start_method("spawn")
    except RuntimeError:
        pass
  12. functional_attention RoPE 条件判断修复:freqs_cis is not None 条件在 use_rope_real=True 时为 False(freqs_cis=None),导致 RoPE 被跳过。

    在 sam3/model/decoder.py 中修改:

    if freqs_cis is not None or freqs_cis_real is not None:
  13. 其他文件适配:

    • sam3/train/utils/train_utils.py:torch.cuda.manual_seed_all → sam3.compat.manual_seed_all;MemMeter 中 torch.cuda.max_memory_allocated → sam3.compat.max_memory_allocated
    • sam3/train/utils/distributed.py:set_cuda_device_index 增加 NPU 分支;DDP device_ids 增加条件判断
    • sam3/model/model_misc.py:get_sdpa_settings 增加 NPU 分支(old_gpu=True, use_flash_attn=False)
    • sam3/model/vl_combiner.py:device="cuda" 默认参数改为 sam3.compat.get_device()
    • sam3/train/loss/sam3_loss.py:device="cuda" 默认改为 sam3.compat.get_device()
    • sam3/train/masks_ops.py、sam3/perflib/nms.py、sam3/perflib/connected_components.py:.is_cuda → .is_cuda or (hasattr(x, 'is_npu') and x.is_npu)
  14. addmm_act NPU 降级处理:CUDA 融合算子 _addmm_activation 在 NPU 无高效实现,原始代码还会断梯度。

    适配文件:fused_npu.py

    在 sam3/perflib/fused.py 中,NPU 环境下用 linear(mat1) + activation(out) 替代融合算子:

    _USE_NPU_ADDMM_FALLBACK = sam3.compat.is_npu_available()
    
    def addmm_act(activation, linear, mat1):
        if _USE_NPU_ADDMM_FALLBACK:
            out = linear(mat1)
            if activation in [torch.nn.functional.relu, torch.nn.ReLU]:
                return torch.nn.functional.relu(out)
            if activation in [torch.nn.functional.gelu, torch.nn.GELU]:
                return torch.nn.functional.gelu(out)
            raise ValueError(f"Unexpected activation {activation}")
        # CUDA original path...
  15. torch.compiler.is_dynamo_compiling 兼容处理:torch 2.1.0 不含此函数。

    在 sam3/model/decoder.py 顶部添加猴子补丁:

    if not hasattr(torch.compiler, 'is_dynamo_compiling'):
        torch.compiler.is_dynamo_compiling = lambda: False

5.3 单卡训练

  1. 设置环境变量(包含 ASCEND_RT_VISIBLE_DEVICES 指定物理卡号)。
  2. 启动单卡训练。

启动单卡训练(指定物理 NPU 卡6):

source /usr/local/Ascend/ascend-toolkit/ascend-toolkit/set_env.sh
export LD_LIBRARY_PATH=/usr/local/Ascend/driver/lib64:/usr/local/Ascend/driver/lib64/driver:$LD_LIBRARY_PATH
export PATH=/usr/local/python3.10/bin:$PATH
export ASCEND_RT_VISIBLE_DEVICES=6
export MASTER_ADDR=localhost
export MASTER_PORT=29500
export RANK=0
export LOCAL_RANK=0
export WORLD_SIZE=1

cd /root/sam3
python3.10 -m sam3.train.train \
    -c configs/coco2017/coco2017_npu_real_train.yaml \
    --use-cluster 0 \
    --num-gpus 1

说明:ASCEND_RT_VISIBLE_DEVICES=6 让进程只看到物理卡6,此时 local_rank=0 映射到 npu:0(即物理卡6)。

COCO 冒烟测试训练配置已创建于 sam3/train/configs/coco2017/coco2017_npu_smoke_test.yaml,关键参数:

  • accelerator: npu、backend: hccl、device: npu、gpus_per_node: 1
  • load_from_HF: false、checkpoint_path: /home/sam3.pt
  • max_epochs: 2、limit_ids: 50(仅50张图片做快速验证) 单卡训练成功

5.4 多卡训练

  1. 使用 ASCEND_RT_VISIBLE_DEVICES 指定可见卡。
  2. 通过 --num-gpus 参数控制每节点卡数。
  3. SAM3 使用 torch.multiprocessing.start_processes(spawn 方式)启动多进程,不需要 torchrun。

使用4卡训练(指定物理卡4,5,6,7):

source /usr/local/Ascend/ascend-toolkit/ascend-toolkit/set_env.sh
export LD_LIBRARY_PATH=/usr/local/Ascend/driver/lib64:/usr/local/Ascend/driver/lib64/common:/usr/local/Ascend/driver/lib64/driver:$LD_LIBRARY_PATH
export PATH=/usr/local/python3.10/bin:$PATH
export ASCEND_RT_VISIBLE_DEVICES=4,5,6,7
export PYTHONPATH=/root/sam3:$PYTHONPATH

cd /root/sam3
python3.10 -m sam3.train.train \
    -c configs/coco2017/coco2017_npu_real_train.yaml \
    --use-cluster 0 \
    --num-gpus 4

注意:

  • ASCEND_RT_VISIBLE_DEVICES=4,5,6,7 让进程只看到4张卡,逻辑编号 0-3 对应物理卡4-7
  • --num-gpus 4 传递给 cfg.launcher.gpus_per_node,控制 spawn 进程数
  • SAM3 的 Sam3VideoPredictorMultiGPU 使用 torch.multiprocessing.start_processes 启动多进程(不是 torchrun)
  • HCCL 后端由训练配置 backend: hccl 指定,无需手动设置
  • 如需使用不同卡数组合,只需修改 ASCEND_RT_VISIBLE_DEVICES 和 --num-gpus 多卡训练成功

5.5 训练结果分析

COCO 2017 验证集 smoke test 结果(50张图片,2个epoch,NPU:6,bfloat16 AMP):

  • 训练启动成功,模型构建+权重加载+数据加载+forward+backward+optimizer step 全流程通过
  • Loss 正常下降:loss_bbox、loss_giou、loss_ce 等分量均有数值
  • 无 Complex dtype 错误(use_rope_real=True 生效)
  • 无 triton 相关错误(条件导入生效)
  • 单张图片训练耗时约 10-15 秒(含 JIT 编译首次开销)

真实训练结果(CANN 8.1.RC1 + torch_npu 2.1.0.post11,全量5000张COCO图片,10 epochs,NPU:6,bfloat16 AMP):

  • 训练启动成功,无崩溃
  • Epoch 0 步骤 0→750 的 loss 趋势:9.26 → 97.3(loss 在下降,有波动是正常的,因为 SAM3 loss 包含多个分量如 focal loss、box loss 等)
  • 平均 batch time:6.9-7.0 秒(稳定,初始化后首 batch 61.85 秒为一次性开销)
  • 内存占用:46 GB / 卡
  • CPU 回退算子(不影响训练正确性,但可能影响性能):torchvision::roi_align、aten::_assert_async、torchvision::_roi_align_backward
  • 预估完整训练时间:~10 小时/epoch × 10 epochs ≈ 100 小时

4卡完整训练结果(CANN 8.1.RC1 + torch_npu 2.1.0.post11,NPU:4,5,6,7,bfloat16 AMP,10 epochs 全量完成):

  • 训练启动成功,4卡全部参与(spawn 多进程方式)
  • 10 个 epoch 全部正常完成,无崩溃、无 OOM、无算子异常
  • 训练总时长:1天13小时10分钟
  • 平均 batch time:7.31s(稳定,首 batch 35.28s 为一次性初始化开销)
  • 内存占用:46 GB / 卡
  • 总 steps:1250/epoch(4卡分摊,每卡处理 5000/4=1250 samples)

4卡训练 Loss 趋势(o2o,未加权平均值):

EpochTotal Lossloss_celoss_bboxloss_gioupresence_lossce_f1presence_dec_acc
0137.720.0025790.0276420.20460.01190.23140.9913
1127.460.0024290.0238540.18990.01010.30010.9927
2120.700.0022690.0227340.18010.00960.33320.9927
3126.000.0024790.0237300.18950.00900.35320.9934
4123.660.0024120.0236900.18000.00930.35660.9929
5123.370.0024220.0231340.18410.00870.35140.9932
6118.560.0023340.0214330.17560.00910.33510.9931
7115.230.0022680.0215390.17230.00810.32530.9938
8119.600.0023350.0225890.17840.00820.33800.9937
9112.260.0022620.0212290.16400.00810.33110.9937

loss曲线

Loss 趋势分析:

  • Total Loss:137.72 → 112.26(下降 18.5%),整体下降趋势正确,Epoch 3/8 有约 5% 的波动属于正常训练方差
  • loss_giou:0.2046 → 0.1640(下降 19.8%),最大的 o2o 分项,主导 total loss 趋势
  • loss_bbox:0.0276 → 0.0212(下降 23.2%),框回归稳步改善
  • loss_ce:0.0026 → 0.0023(下降 12.3%),值很小但趋势正确
  • presence_loss:0.0119 → 0.0081(下降 31.4%),目标存在性检测改善明显
  • ce_f1:0.23 → 0.33(上升 43.1%),分类 F1 指标持续提升
  • presence_dec_acc:0.9913 → 0.9937,始终接近 1.0

4卡训练 COCO 检测 AP 结果(val_epoch_freq=5,仅在 Epoch 5 和 Epoch 9 执行验证):

EpochAPAP_50AP_75AP_smallAP_mediumAP_large
50.6200.7960.6760.4610.6680.794
90.6290.8040.6860.4770.6790.802

AP 从 Epoch 5 到 Epoch 9 提升 0.009(0.620 → 0.629),训练趋势正确。

训练产出文件:

文件路径说明
训练日志/root/sam3_logs/coco2017_npu_real_train/logs/coco/log.txt完整训练过程日志
训练统计/root/sam3_logs/coco2017_npu_real_train/logs/coco/train_stats.json每 epoch 一行 JSON(loss、F1、accuracy)
验证统计/root/sam3_logs/coco2017_npu_real_train/logs/coco/val_stats.json每 val epoch 一行 JSON(AP 指标)
验证预测/root/sam3_logs/coco2017_npu_real_train/dumps/coco/coco_predictions_bbox.jsonCOCO bbox 预测 JSON(55MB)
Loss 曲线图/root/loss_curves.png6 子图:Total Loss + 4 分项 + ce_f1 + AP
Checkpoint(最新)/root/sam3_logs/coco2017_npu_real_train/checkpoints/checkpoint.pt约 10GB
Checkpoint(Epoch 5)/root/sam3_logs/coco2017_npu_real_train/checkpoints/checkpoint_5.pt约 10GB
Checkpoint(Epoch 10)/root/sam3_logs/coco2017_npu_real_train/checkpoints/checkpoint_10.pt约 10GB

说明:Total Loss 为加权 core_loss(包含 weight_dict 权重、aux loss、o2m loss 和 o2m_weight 乘数),与 4 个 o2o 分项(loss_ce、loss_bbox、loss_giou、presence_loss)的未加权加和之间存在约 5.4 倍的固定比率(因权重和层数乘数),不影响趋势分析正确性。 单卡 vs 4卡训练对比:

指标单卡(NPU:6)4卡(NPU:4,5,6,7)
batch time7.0s7.31s
每 epoch steps50001250
每 epoch 时间~10h~2.5h
内存 / 卡46 GB46 GB
加速比1x~4x

6. 推理迁移步骤

SAM3 推理采用 PyTorch PT 模式直接在 NPU 上运行,无需 ONNX/OM 转换。核心适配手段是 torch_npu.contrib.transfer_to_npu——一行代码即可全局将 .cuda() 重定向为 .npu()、torch.cuda.* 重定向为 torch.npu.*、torch.autocast(device_type="cuda") 自动适配 NPU、torch.distributed.init_process_group(backend="nccl") 自动改为 hccl。

6.1 PyTorch 推理适配

  1. 在推理脚本顶部导入 transfer_to_npu(必须在 import torch 之前)。
  2. 修复 _setup_tf32() 和 get_sdpa_settings() 中 NPU 设备属性没有 major 字段的问题。
  3. 处理 freqs_cis buffer key 不匹配问题(strict_state_dict_loading=False)。
  4. 修复 FA3(Flash Attention 3)在 NPU 上不可用的问题。

推理适配代码修改清单:

  1. transfer_to_npu 全局重定向:在推理脚本最顶部(import torch 之前)添加一行即可自动处理所有 .cuda() → .npu()、torch.cuda.* → torch.npu.*、autocast(device_type="cuda") → NPU 兼容、backend="nccl" → backend="hccl"。

    from torch_npu.contrib import transfer_to_npu  # 必须在 import torch 之前

    注意:transfer_to_npu 会禁用 torch.jit.script 和 torch.jit.script_method,如需启用则不可使用此方式。

  2. _setup_tf32() NPU 保护:transfer_to_npu 使 torch.cuda.is_available() 返回 True,但 NPU 设备属性没有 major 字段,需先检测 NPU 再跳过 CUDA 逻辑。

    在 sam3/model_builder.py 中修改 _setup_tf32():

    def _setup_tf32() -> None:
        try:
            import torch_npu
            if torch.npu.is_available():
                return
        except ImportError:
            pass
        if torch.cuda.is_available():
            device_props = torch.cuda.get_device_properties(0)
            if device_props.major >= 8:
                torch.backends.cuda.matmul.allow_tf32 = True
                torch.backends.cudnn.allow_tf32 = True
  3. get_sdpa_settings() NPU 优先:transfer_to_npu 使 torch.cuda.is_available() 返回 True,需先检测 NPU 跳过 CUDA major 字段检查。

    在 sam3/model/model_misc.py 中修改 get_sdpa_settings():

    def get_sdpa_settings():
        try:
            import torch_npu
            is_npu = torch.npu.is_available()
        except ImportError:
            is_npu = False
        if is_npu:
            old_gpu = True
            use_flash_attn = False
            math_kernel_on = True
        elif torch.cuda.is_available():
            old_gpu = torch.cuda.get_device_properties(0).major < 7
            use_flash_attn = torch.cuda.get_device_properties(0).major >= 8
            ...
        else:
            ...
  4. freqs_cis buffer key 不匹配:NPU 上 _USE_ROPE_REAL=True,模型注册 freqs_cis_real/freqs_cis_imag buffer,但权重文件存储 freqs_cis(complex 格式)。这些 buffer 在运行时由 compute_axial_cis 重新生成,不影响推理正确性。

    在 sam3/model_builder.py 的 build_sam3_video_model 中,NPU 推理时自动关闭 strict loading:

    if _USE_ROPE_REAL and strict_state_dict_loading:
        strict_state_dict_loading = False
  5. FA3 在 NPU 上禁用:Flash Attention 3 使用 NVIDIA Hopper float8 格式,NPU 不支持。在 sam3/sam/transformer.py 中修改:

    if self.use_fa3 and not sam3.compat.is_npu_available():
        from sam3.perflib.fa3 import flash_attn_func
        ...
    else:
        out = F.scaled_dot_product_attention(q, k, v, dropout_p=dropout_p)

    同时删除 torch.backends.cuda.enable_flash_sdp(True) 等 CUDA-specific SDPA 控制(NPU 的 SDPA 由 CANN kernels 提供)。

6.2 ONNX导出与ATC转换

SAM3 推理链路拆分为 3 个独立 ONNX 模型:Image Encoder、Text Encoder、Grounding Head。每个模型单独导出 ONNX 后通过 ATC 转换为 OM,最终以纯 OM 方式在 NPU 上链式推理。

6.2.1 Image Encoder ONNX 导出

  1. 关键适配:不导入 torch_npu,通过动态修补(monkey-patch).cuda() 和 torch.nn.Module.cuda() 使模型保持在 CPU 上,修补 torch.compiler.is_dynamo_compiling。

  2. 导出脚本核心逻辑:

torch.cuda.is_available = lambda: False

def cuda_to_cpu(self, device=None, non_blocking=False):
    return self
torch.Tensor.cuda = cuda_to_cpu
torch.nn.Module.cuda = cuda_to_cpu

torch.compiler.is_dynamo_compiling = lambda: False
  1. Wrapper 设计:将 Image Encoder 的 backbone.forward_image() 输出封装为 7 个输出 tensor。
class ImageEncoderONNXWrapper(torch.nn.Module):
    def forward(self, image):
        backbone_out = self.model.backbone.forward_image(image)
        vision_features = backbone_out["vision_features"]
        vision_pos_enc = backbone_out["vision_pos_enc"]
        backbone_fpn = backbone_out["backbone_fpn"]
        return (vision_features,
                vision_pos_enc[0], vision_pos_enc[1], vision_pos_enc[2],
                backbone_fpn[0], backbone_fpn[1], backbone_fpn[2])
  1. 导出参数:
torch.onnx.export(
    wrapper, (image_tensor,), ONNX_PATH,
    input_names=["image"],
    output_names=["vision_features", "vision_pos_enc_0", "vision_pos_enc_1",
                  "vision_pos_enc_2", "backbone_fpn_0", "backbone_fpn_1", "backbone_fpn_2"],
    opset_version=17, do_constant_folding=True,
)
  1. 导出结果:/root/sam3_image_encoder.onnx(1.7GB),输入 image:1,3,1008,1008,输出 7 个 tensor。

6.2.2 Text Encoder ONNX 导出

  1. Wrapper 设计:封装 VETextEncoder 的 encoder(token_ids) + resizer(text_memory) 流程,输出 attention_mask(padding mask)和 text_memory_resized。
class TextEncoderONNXWrapper(torch.nn.Module):
    def forward(self, token_ids):
        _, text_memory = self.encoder(token_ids)
        attention_mask = (token_ids == 0)  # padding mask: 1=padding, 0=real token
        text_memory_seq = text_memory.transpose(0, 1)
        text_memory_resized = self.resizer(text_memory_seq)
        return attention_mask.to(torch.float32), text_memory_resized

注意:attention_mask = (token_ids == 0) 是 padding mask(True 表示应忽略的 padding 位置),不是 attention mask(True 表示应关注的 token)。Grounding Head 中将其作为 prompt_key_padding_mask 使用。

  1. 导出结果:/root/sam3_text_encoder.onnx(1.4GB),输入 token_ids:1,32(int64),输出 attention_mask:1,32(float32)+ text_memory_resized:32,1,256(float32)。

6.2.3 Grounding Head ONNX 导出

  1. 关键适配:除与 Image Encoder 相同的 CPU 强制 patch 外,还需绕过 oneDNN F.linear output_dim=1 崩溃 bug(详见问题排查第 18 条)。
import torch.nn.functional as F
original_linear = F.linear
def patched_linear(input, weight, bias=None):
    if weight.shape[0] == 1:
        weight_padded = torch.cat([weight, torch.zeros(1, weight.shape[1], device=weight.device, dtype=weight.dtype)], dim=0)
        bias_padded = torch.cat([bias, torch.zeros(1, device=bias.device, dtype=bias.dtype)], dim=0) if bias is not None else None
        result = original_linear(input, weight_padded, bias_padded)
        return result[..., :1]
    return original_linear(input, weight, bias)
F.linear = patched_linear
  1. import 修复:inverse_sigmoid 应从 sam3.model.model_misc 导入(不是 sam3.model.data_misc);box_cxcywh_to_xyxy 应从 sam3.model.box_ops 导入。

  2. Wrapper 设计:将 Transformer Encoder Fusion + Decoder + scoring + bbox_head + segmentation_head 封装为单模型。输入 6 个 tensor(因 num_feature_levels=1,ONNX 自动消除了未使用的 vision_pos_enc_0/1)。

class GroundingHeadONNXWrapper(torch.nn.Module):
    def forward(self, backbone_fpn_0, backbone_fpn_1, backbone_fpn_2,
                vision_pos_enc_2, text_memory_resized, text_attention_mask_float):
        # encoder → decoder → scoring → bbox → seg_head
        return pred_logits, pred_boxes, pred_masks, presence_logit_dec
  1. 导出结果:/root/sam3_grounding_head.onnx(97.4MB),6 个输入,4 个输出:
    • pred_logits:6,1,200,1(6 层 decoder 输出,取最后一层 [5])
    • pred_boxes:6,1,200,4(同上,取 [5])
    • pred_masks:1,200,288,288(ONNX 动态维度标记为 [0,0,0,0])
    • presence_logit_dec:6,1,1

6.2.4 ATC 转换

  1. 设置 CANN 环境(含驱动 LD_LIBRARY_PATH)。
source /usr/local/Ascend/ascend-toolkit/ascend-toolkit/set_env.sh
export LD_LIBRARY_PATH=/usr/local/Ascend/driver/lib64/driver:/usr/local/Ascend/ascend-toolkit/ascend-toolkit/latest/devlib:/usr/local/Ascend/ascend-toolkit/ascend-toolkit/latest/lib64:/usr/local/Ascend/ascend-toolkit/ascend-toolkit/latest/hccl/lib64:$LD_LIBRARY_PATH
  1. 图像编码器ATC(注意:OM输出尺寸为0,需要硬编码shape分配缓冲区)。
atc --framework=5 --soc_version=Ascend910B3 \
    --model=/root/sam3_image_encoder.onnx \
    --output=/root/sam3_om_static \
    --input_shape="image:1,3,1008,1008" \
    --input_format=NCHW
# 结果:sam3_om_static_linux_aarch64.om(1003MB)
  1. 文本编码器ATC。
atc --framework=5 --soc_version=Ascend910B3 \
    --model=/root/sam3_text_encoder.onnx \
    --output=/root/sam3_text_encoder \
    --input_shape="token_ids:1,32" \
    --input_format=NCHW
# 结果:sam3_text_encoder_linux_aarch64.om(746MB)
  1. 接地头 ATC。
atc --framework=5 --soc_version=Ascend910B3 \
    --model=/root/sam3_grounding_head.onnx \
    --output=/root/sam3_grounding_head \
    --input_shape="backbone_fpn_0:1,256,288,288;backbone_fpn_1:1,256,144,144;backbone_fpn_2:1,256,72,72;vision_pos_enc_2:1,256,72,72;text_memory_resized:32,1,256;text_attention_mask_float:1,32" \
    --input_format=NCHW
# 结果:sam3_grounding_head_linux_aarch64.om(112MB)

ATC 警告信息(不影响推理正确性,仅影响部分算子性能):Op [/decoder/layers.*/ca_text/Mod] does not hit the high-priority operator information library。

6.3 OM模型推理

纯 OM 推理使用 acl API 链式调用 3 个 OM 模型,不依赖 PyTorch 或 torch_npu。

6.3.1 推理流程

  1. acl.init() + acl.rt.set_device(0) 初始化 NPU。
  2. 加载 3 个 OM 模型(acl.mdl.load_from_file)。
  3. 准备输入:加载图片 → 调整大小至 1008x1008 → ImageNet 归一化 → [1,3,1008,1008] float32 格式;文本 BPE 分词 → [1,32] int64 格式。
  4. 图像编码器 OM 推理:将图像数据复制到设备 → 执行推理 → 将 7 个输出复制到主机。
  5. 文本编码器 OM 推理:将 token_ids 复制到设备 → 执行推理 → 将 2 个输出复制到主机(attention_mask 为 float32 填充掩码)。
  6. 定位头 OM 推理:将 6 个输入复制到设备 → 执行推理 → 将 4 个输出复制到主机。
  7. 后处理:对 logits[5] 进行 sigmoid 运算得到分数,将 cxcywh 格式转换为 xyxy 格式的框坐标,将掩码调整至原图尺寸。
  8. 清理:卸载模型 → 重置设备 → 结束。

6.3.2 acl API 关键要点

  • acl.rt.malloc(size, policy):2 个参数,policy=0(ACL_MEM_MALLOC_HUGE_FIRST)。
  • acl.rt.memcpy(dst, dst_size, src, src_size, kind):5 个参数,kind=1(主机到设备),kind=2(设备到主机)。
  • acl.create_data_buffer(dev, size):返回单个整数。
  • OM 动态维度输出 size=0 时,必须硬编码已知形状以分配输出缓冲区。
  • acl.util.numpy_to_ptr 将 numpy 数组转换为指针(bytes_to_ptr 是新版 API)。
  • acl.init() 需先执行,推理后不调用 acl.finalize(),直至最后结束。

6.3.3 推理脚本

推理脚本路径:pure_om_inference_trained.py

运行命令:

source /usr/local/Ascend/ascend-toolkit/ascend-toolkit/set_env.sh
export LD_LIBRARY_PATH=/usr/local/Ascend/driver/lib64/driver:/usr/local/Ascend/ascend-toolkit/ascend-toolkit/latest/devlib:/usr/local/Ascend/ascend-toolkit/ascend-toolkit/latest/lib64:/usr/local/Ascend/ascend-toolkit/ascend-toolkit/latest/hccl/lib64:$LD_LIBRARY_PATH
/usr/local/python3.10/bin/python3.10 /root/pure_om_inference_trained.py

6.3.4 修改检测类别

纯 OM 推理脚本 /root/pure_om_inference.py 第 32 行定义了检测类别:

TEXT_PROMPT = "person"

修改 TEXT_PROMPT 即可改变检测目标,例如:

  • TEXT_PROMPT = "car" → 检测汽车
  • TEXT_PROMPT = "dog" → 检测狗
  • TEXT_PROMPT = "person car" → 同时检测人和汽车(一条文本中包含多个类别词)

注意:SAM3 是 grounding 检测模型,每次推理对应一条文本描述。每次只能检测当前 TEXT_PROMPT 中描述的类别。如需做完整的 COCO 80 类评估,需对每个类别分别运行一次 Text Encoder OM + Grounding Head OM(Image Encoder OM 只需运行一次,然后复用其 7 个输出),汇总所有类别的检测结果。

6.3.5 OM 模型输入输出规格

模型输入输出
Image Encoderimage:1,3,1008,1008(float32)vision_features:1,256,72,72、vision_pos_enc_0/1/2、backbone_fpn_0/1/2(7个,均为float32)
Text Encodertoken_ids:1,32(int64)attention_mask:1,32(float32,padding mask)、text_memory_resized:32,1,256(float32)
Grounding Headbackbone_fpn_0/1/2、vision_pos_enc_2、text_memory_resized、text_attention_mask_float(6个)pred_logits:6,1,200,1、pred_boxes:6,1,200,4、pred_masks:1,200,288,288、presence_logit_dec:6,1,1(4个)

6.3.5 推理链路数据流

图片 → [Image Encoder OM] → backbone_fpn_0/1/2, vision_pos_enc_2
                                ↓
文本 → [Text Encoder OM] → text_memory_resized, attention_mask_float
                                ↓
        [Grounding Head OM] → pred_logits, pred_boxes, pred_masks, presence_logit_dec
                                ↓
                        后处理 → 检测框 + 分割mask

6.4 推理结果分析

PT 模式 NPU 推理结果(CANN 8.1.RC1 + torch_npu 2.1.0.post11,NPU:0,bfloat16 AMP):

  • 推理启动成功,模型加载至 NPU 设备
  • start_session(单张图片)耗时约 0.1 秒
  • add_prompt(文本提示"person")返回 mask 结果,mask 非零像素数约 11789
  • 缺失的键(均为 freqs_cis 缓冲区,运行时重新生成,不影响推理):detector.backbone.vision_backbone.trunk.blocks.*.attn.freqs_cis_real/imag
  • 意外的键(权重中原始 freqs_cis,已被 strict=False 忽略):detector.backbone.vision_backbone.trunk.blocks.*.attn.freqs_cis
  • CPU 回退算子(不影响推理正确性):torchvision::roi_align、aten::_assert_async

纯 OM 推理结果(acl API,CANN 8.1.RC1,Ascend 910B3):

测试配置:COCO 2017 图片 000000000139.jpg(640x426),文本提示 "person"

模型推理耗时
Image Encoder OM0.72s
Text Encoder OM0.02s
Grounding Head OM0.04s
总推理耗时0.78s(不含模型加载)

检测结果:2 个 person 目标

检测目标ScoreBox (xyxy, 原图坐标)
Query 250.9428[413.1, 157.5, 464.4, 295.1]
Query 550.9080[384.5, 172.1, 400.5, 206.9]

OM 与混合 OM+PT 推理对比:

方案检测数最高 Score推理方式
纯 OM(3 个 OM 链式 acl)20.9428acl API,无 PyTorch 依赖
混合 OM+PT(Image Encoder OM + PT Grounding)20.929Image Encoder OM + PyTorch 推理
纯 PT(PyTorch NPU)2~0.93全量 PyTorch 推理

三种方案检测结果一致(均为 2 个 person),score 差异 < 0.02,在 float16 舍入精度范围内。

输出文件保存在 /root/pure_om_outputs/:

  • img_enc_*.npy:Image Encoder 的 7 个输出
  • txt_enc_*.npy:Text Encoder 的 2 个输出
  • grd_head_*.npy:Grounding Head 的 4 个输出
  • result.jpg:可视化结果图片(检测框 + mask 叠加)
  • scores.npy、boxes_xyxy.npy:后处理结果

6.5 训练后 Checkpoint OM 导出与推理

训练保存的 checkpoint(如 checkpoint_10.pt)需额外处理才能正确导出 ONNX/OM 并推理。

6.5.1 Checkpoint Key 前缀问题

训练 checkpoint 的 state_dict key 没有 detector. 前缀(DDP unwrap 后去除),而原始 sam3.pt 的 key 有 detector. 前缀。_load_checkpoint 函数只处理含 detector. 的 key,导致训练 checkpoint 的所有权重被忽略。

来源Key 示例
原始 sam3.ptdetector.backbone.language_backbone.encoder.positional_embedding
训练 checkpoint_10.ptbackbone.language_backbone.encoder.positional_embedding

影响:positional_embedding 等 key 被视为缺失键,模型用 torch.empty 初始化(含 NaN),导致 Text Encoder 输出全为 NaN。

6.5.2 修复方法

在导出脚本中手动加载 checkpoint 并修复 key 前缀,不依赖 _load_checkpoint:

CKPT = '/root/sam3_logs/coco2017_npu_real_train/checkpoints/checkpoint_10.pt'

model = build_sam3_image_model(
    bpe_path='.../bpe_simple_vocab_16e6.txt.gz',
    checkpoint_path=None,
    load_from_HF=False,
    device='cpu',
)

ckpt_data = torch.load(CKPT, map_location='cpu', weights_only=True)
raw_state = ckpt_data['model'] if 'model' in ckpt_data else ckpt_data

fixed_state = {}
for k, v in raw_state.items():
    if 'detector.' not in k:
        fixed_state['detector.' + k] = v
    else:
        fixed_state[k] = v

sam3_ckpt = {k.replace("detector.", ""): v for k, v in fixed_state.items() if "detector" in k}
missing, unexpected = model.load_state_dict(sam3_ckpt, strict=False)
print(f"Loaded: missing={len(missing)}, unexpected={len(unexpected)}")
model.eval()

注意:missing 中应只有 freqs_cis_real/freqs_cis_imag(运行时重新生成,不影响推理)。

6.5.3 导出脚本 Patch 清单

训练后 checkpoint ONNX 导出需以下 patch(与原始权重导出相同,但额外需要 key 前缀修复):

  1. torch.cuda.is_available = lambda: False:禁用 CUDA 检测
  2. .cuda() → 保持在 CPU:torch.Tensor.cuda = lambda self, *a, **kw: self、torch.nn.Module.cuda = lambda self, *a, **kw: self
  3. .npu() → 保持在 CPU:torch.Tensor.npu = lambda self, *a, **kw: self、torch.nn.Module.npu = lambda self, *a, **kw: self
  4. torch.Tensor.to() 拦截 npu/cuda:拦截 device='npu'/device='cuda' 参数,返回 CPU tensor
  5. sam3.compat.get_device = lambda: 'cpu':设备选择强制 CPU
  6. sam3.compat.is_npu_available = lambda: False:禁用 NPU 检测(Grounding Head 导出必需)
  7. device='cpu' 参数:build_sam3_image_model(..., device='cpu')
  8. torch.compiler.is_dynamo_compiling = lambda: False:torch 2.1.0 兼容
  9. F.linear padding patch(Grounding Head 导出必需):oneDNN output_dim=1 bug 规避方法
  10. Checkpoint key 前缀修复(训练后 checkpoint 导出新增):给所有 key 加 detector. 前缀

6.5.4 ATC 转换

与原始权重 OM 转换命令相同,只需替换输入 ONNX 文件路径:

# Image Encoder
atc --framework=5 --soc_version=Ascend910B3 \
    --model=/root/sam3_image_encoder_trained.onnx \
    --output=/root/sam3_image_encoder_trained \
    --input_shape="image:1,3,1008,1008" --input_format=NCHW
# 结果:sam3_image_encoder_trained_linux_aarch64.om(1004MB)

# Text Encoder
atc --framework=5 --soc_version=Ascend910B3 \
    --model=/root/sam3_text_encoder_trained.onnx \
    --output=/root/sam3_text_encoder_trained \
    --input_shape="token_ids:1,32" --input_format=NCHW
# 结果:sam3_text_encoder_trained_linux_aarch64.om(746MB)

# Grounding Head
atc --framework=5 --soc_version=Ascend910B3 \
    --model=/root/sam3_grounding_head_trained.onnx \
    --output=/root/sam3_grounding_head_trained \
    --input_shape="backbone_fpn_0:1,256,288,288;backbone_fpn_1:1,256,144,144;backbone_fpn_2:1,256,72,72;vision_pos_enc_2:1,256,72,72;text_memory_resized:32,1,256;text_attention_mask_float:1,32" \
    --input_format=NCHW
# 结果:sam3_grounding_head_trained_linux_aarch64.om(112MB)

6.5.5 训练后 OM 推理结果

推理结果 测试配置:COCO 2017 图片 000000000139.jpg(640x426),文本提示 "person",训练 checkpoint 为 checkpoint_10.pt(10 轮次)。

模型推理耗时
Image Encoder OM0.72s
Text Encoder OM0.02s
Grounding Head OM0.04s

检测结果:2 个 person 目标

检测目标ScoreBox (xyxy, 原图坐标)
Query 250.8942[409.3, 157.0, 465.7, 296.5]
Query 1660.8604[384.4, 172.1, 401.2, 206.9]

与原始 OM 推理对比:

方案检测数最高 ScoreBox
原始 OM(sam3.pt)20.9428[413.1, 157.5, 464.4, 295.1] + [384.5, 172.1, 400.5, 206.9]
训练 OM(checkpoint_10.pt)20.8942[409.3, 157.0, 465.7, 296.5] + [384.4, 172.1, 401.2, 206.9]

两者检测目标一致(均为 2 个 person),框坐标接近,得分略有下降(0.89 vs 0.94)。这符合预期——在 COCO 上训练 10 轮次的检测 AP=0.629,而原始预训练权重在 COCO 上的 AP 通常更高(预训练数据量更大)。

推理脚本路径:/root/pure_om_inference_trained.py(指向训练后的 OM 文件,输出目录 pure_om_outputs_trained)。

7. 问题排查

问题 1:NPU 算子缺失 kernel binary

  • 现象:调用 Add、Mul(标量乘)、AsStrided 等算子时报错 Op xxx does not has any binary 或 FlashAttention AsStrided 内核执行失败。

  • 可能原因:未安装 Ascend-cann-kernels 包,缺少预编译算子二进制文件。

  • 处理方法:安装 CANN kernels 包(推荐 8.1.RC1,包含 8869 个预编译算子),安装后不再需要 jit_compile=True。如使用 CANN 8.0.RC2(无 kernels 包),则需开启 JIT 编译模式:torch.npu.set_compile_mode(jit_compile=True)。

  • 安装 CANN 8.1.RC1 kernels 步骤:

    1. 下载 kernels 包(.zip 双层打包)。
    curl -L -o /tmp/Ascend-cann-kernels-910b_8.1.RC1_linux-aarch64.zip \
        "https://ascend-repo.obs.cn-east-2.myhuaweicloud.com/CANN/CANN%208.1.RC1/Ascend-cann-kernels-910b_8.1.RC1_linux-aarch64.zip"
    1. 解压两层 .zip,得到 .run 文件。
    python3 -c "import zipfile; zipfile.ZipFile('/tmp/Ascend-cann-kernels-910b_8.1.RC1_linux-aarch64.zip').extractall('/tmp/kernels_stage1')"
    python3 -c "import zipfile; zipfile.ZipFile('/tmp/kernels_stage1/Ascend-cann-kernels-910b_8.1.RC1_linux-aarch64.zip').extractall('/tmp/kernels_stage2')"
    1. 安装 .run 文件。
    chmod +x /tmp/kernels_stage2/Ascend-cann-kernels-910b_8.1.RC1_linux-aarch64.run
    /tmp/kernels_stage2/Ascend-cann-kernels-910b_8.1.RC1_linux-aarch64.run --install --quiet --force
  • 校验命令:

    python3.10 -c "
    import torch, torch_npu
    torch.npu.set_device(0)
    # AsStrided 是之前缺失的算子
    x = torch.randn(4, 4, device='npu:0')
    y = torch.as_strided(x, (2, 2), (4, 1), 0)
    print('AsStrided OK:', y.shape)
    # FlashAttention
    q = torch.randn(1, 8, 64, 128, device='npu:0')
    k = torch.randn(1, 8, 64, 128, device='npu:0')
    v = torch.randn(1, 8, 64, 128, device='npu:0')
    out = torch.nn.functional.scaled_dot_product_attention(q, k, v)
    print('SDPA OK:', out.shape)
    "

问题 2:CANN TBE 初始化缺少 Python 模块

  • 现象:aclrtCreateContext 失败,日志提示 No module named 'decorator'、No module named 'scipy'、No module named 'attr'。
  • 可能原因:CANN TBE 编译器需要 decorator、scipy、attr 等模块但未安装。
  • 处理方法:安装缺失依赖:pip3.10 install decorator scipy attr attrs psutil astunparse cffi pathlib2 protobuf。
  • 校验命令:
python3.10 -c "import decorator, scipy, attr, attrs, psutil, astunparse, cffi, pathlib2, protobuf; print('OK')"

问题 3:torch 2.1.0 缺少 torch.nn.attention 模块

  • 现象:ModuleNotFoundError: No module named 'torch.nn.attention'。
  • 可能原因:torch.nn.attention(含 sdpa_kernel、SDPBackend)在 PyTorch 2.2 才引入。
  • 处理方法:创建兼容 shim 文件 sam3/compat/torch_nn_attention.py,在 decoder.py 和 vl_combiner.py 中用 try/except 导入。
  • 校验命令:
python3.10 -c "from sam3.compat.torch_nn_attention import sdpa_kernel, SDPBackend; print('OK')"