# torch_scatter **Repository Path**: onescience-ai/torch_scatter ## Basic Information - **Project Name**: torch_scatter - **Description**: 本项目为基于flagos统一中间层实现的torch_scatter库 - **Primary Language**: Unknown - **License**: MIT - **Default Branch**: master - **Homepage**: None - **GVP Project**: No ## Statistics - **Stars**: 0 - **Forks**: 0 - **Created**: 2026-08-25 - **Last Updated**: 2026-09-30 ## Categories & Tags **Categories**: Uncategorized **Tags**: None ## README # PyTorch Scatter

-------------------------------------------------------------------------------- 这是面向 FlagOS 的 torch_scatter 实现,保留上游 `torch_scatter` 的公开导入名和主要 API。加速设备上的计算默认使用 FlagTree Triton kernel,CPU 使用 PyTorch reference 路径。本版本不包含 C++、CUDA/HIP 扩展源码。 文档参考:[torch-scatter Documentation](https://pytorch-scatter.readthedocs.io) ## 算子库介绍 本库用于将源张量中的值按照 `index` 指定的分组写入输出张量,并对同一分组执行归约。 支持的归约类型为 `sum`、`mean`、`mul`、`min` 和 `max`。 - **Scatter**:使用任意顺序的索引进行分组归约,不要求 `index` 排序。 - **Segment COO**:使用 COO 形式的分组索引进行归约,通常要求索引按分组有序。 - **Segment CSR**:使用 `indptr` 指针描述每个分段的范围,适合已有 CSR/压缩行结构的数据。 - **Gather COO/CSR**:根据 COO 索引或 CSR 指针,从源张量提取对应分组的数据。 例如,`scatter_sum(src, index, dim=0)` 会把所有满足 `index[i] == j` 的 `src[i]` 累加到输出的第 `j` 个位置;`dim_size` 用于显式指定输出分组数。上述基础算子还被 `scatter_std`、`scatter_logsumexp`、`scatter_softmax` 和 `scatter_log_softmax` 等复合 算子调用。 ## 环境要求 - Python >= 3.8 - 已安装与当前软件栈匹配的 PyTorch 和 Triton - 运行 benchmark 还需要 `scipy`、`wget` ## 安装 进入本目录: ```bash cd /path_to_dir/torch_scatter ``` ### 按硬件后端安装 源码安装或构建 wheel 时,可以通过通用环境变量 `FLAGOS_BACKEND` 选择目标硬件: ```bash # NVIDIA(默认值;不设置变量时使用) python -m pip install . --no-deps --no-build-isolation # NVIDIA/GPU FLAGOS_BACKEND=nvidia python -m pip install . --no-deps --no-build-isolation # Hygon/DCU FLAGOS_BACKEND=hygon python -m pip install . --no-deps --no-build-isolation # MThreads/MUSA FLAGOS_BACKEND=mthreads python -m pip install . --no-deps --no-build-isolation ``` 开发和调试推荐使用 editable 安装: ```bash # 以下以 Hygon 为例;请根据目标硬件替换 hygon FLAGOS_BACKEND=hygon python -m pip install -e . --no-deps --no-build-isolation ``` 构建 wheel 时使用相同的变量: ```bash FLAGOS_BACKEND=hygon python -m pip wheel . --no-deps --no-build-isolation --wheel-dir dist ``` 当前可选值只有 `nvidia`、`hygon` 和 `mthreads`,默认是 `nvidia`。该变量只在安装/构建 阶段读取;安装完成后不会通过修改环境变量动态切换后端。生成的 wheel 会带有对应的 版本标记(例如 `2.1.2+triton.hygon`),可以直接安装,安装时无需再次设置变量: ```bash python -m pip install --force-reinstall --no-deps dist/torch_scatter-2.1.2+triton.hygon-*.whl ``` 检查已安装的后端: ```bash python -c "import torch_scatter; print(torch_scatter.__version__); print(torch_scatter.backend_name())" ``` ## 基本用法 ```python import torch from torch_scatter import scatter_sum, scatter_mean src = torch.tensor([1., 2., 3., 4.], device="cuda") index = torch.tensor([0, 0, 1, 1], device="cuda") result = scatter_sum(src, index, dim=0, dim_size=2) # tensor([3., 7.], device='cuda') mean = scatter_mean(src, index, dim=0, dim_size=2) # tensor([1.5, 3.5], device='cuda') ``` ## 算子示例 下面示例展示 `scatter_max` 返回归约值及其输入位置: ```python import torch from torch_scatter import scatter_max src = torch.tensor([[2, 0, 1, 4, 3], [0, 2, 1, 3, 4]]) index = torch.tensor([[4, 5, 4, 2, 3], [0, 0, 2, 2, 1]]) out, argmax = scatter_max(src, index, dim=-1) print(out) print(argmax) ``` 输出示例: ```text tensor([[0, 0, 4, 3, 2, 0], [2, 4, 3, 0, 0, 0]]) tensor([[5, 5, 3, 4, 0, 1], [1, 4, 3, 5, 5, 5]]) ``` ## 正确性测试 在本目录执行完整测试: ```bash cd /path_to_dir/torch_scatter pytest -q ``` 只运行 Triton 边界测试: ```bash pytest -q test/test_triton_edges.py ``` 测试覆盖前向、反向、广播、空张量、非连续 `out`、NaN/Inf、索引边界以及多卡场景。 ## 性能测试 benchmark 默认使用 `citationCiteseer` 数据集和 feature size `1, 16, 32`。数据集文件不存在时会自动下载。 运行 scatter/segment benchmark: ```bash cd /path_to_dir/torch_scatter python scatter_segment.py --reduce sum --iters 100 ``` 运行 gather benchmark: ```bash python gather.py --iters 100 ``` 运行固定规模的 Triton scatter benchmark: ```bash python triton_benchmark.py --reduce sum --warmup 10 --iters 100 ``` 可通过 `--sizes`、`--datasets`、`--all-datasets`、`--warmup` 和 `--iters` 调整测试规模。benchmark 结果受计算卡型号、索引分布、缓存状态和首次 JIT 编译影响,应在同一设备和相同参数下比较。 ### 模型实际输入测例 `mattergen_mp20_sparse_batch256.npz` 来自 OneScience MatterGen 的 MP-20 cache,包含 256 个结构的 GemNet 邻接输入(`row`/`col`、`edge_batch`、`num_atoms` 等字段),不是随机构造的形状。 `model_benchmark.py` 测量 GemNet `get_triplets` 所需的 `SparseTensor` 行查询、triplet 过滤及 SpMM 前向/反向;默认使用 32 个结构、512 个通道、预热 10 次并迭代 50 次: ```bash python model_benchmark.py --device cuda --warmup 10 --iters 50 ``` `--input` 可指定其他同字段的 `.npz` 快照,`--num-structures 0` 使用全部结构,`--compare-torch`