# 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`