alibaba / alibaba/graphlearn-for-pytorch

[Feat] Distributed Sparse Backend

Open
#36 0 comments 0 reactions 0 assignees View on GitHub
cuda distributed feature llm model nn pytorch
Dominant language
Python
Stars
146
Forks
48
PR merge metrics
No merged PRs in 30d

Description

### 🚀 The feature, motivation and pitch

# Background
The GNN convolutions generates a lot of memory expansion during message passing and hard to make an optimal use of parallelization resources due to the sparsity, which result in insufficient computational performance and too high peak memory.
`geSpMM`, `geSDDMM` integrate graph operator, matrix calculation and reduce operator into one sparse kernel, to reduce kernel launch times and usage of memory, and then improve performance.

# Objective
Distributed Sparse Backend using Sparse Matrix Multiplication to express convolutions in GNN, replacing the commonly used Message Passing paradigm, and supporting high distributed sparse convolution.

Moreover, we can optimize the parallel implementation of the kernel based on the sparsity and feature dimensions of the input data.
When the graph data or model is too large, we can use data parallelism, model parallelism, and pipeline parallelism for distributed optimization.

# Tasks
This work includes the following major tasks, we will enrich each specific task into detailed subtasks.

Phase 1: Implementations
- [ ] Sparse Matrix representation: Convert graph data in GNN into sparse matrix format for efficient matrix computation like multiplication, softmax...
- [ ] Sparse Matrix computation kernels: like geSpMM, geSDDMM, EdgeSoftmax..
- [ ] GNN models: Implement basic GNN models and LLM-GNN models with Sparse kernels to improve computation efficiency and reduce peak memory .
- [ ] Distributed sparse modules: For commonly used GNN models, using DP, MP, PP to implement the most efficient distributed sparse convs, just like Megatron.

Phase 2: Performance optimizations
- [ ] Kernel optimization: Optimize parallelization of kernels for different workloads, half-precision and mixed-precision.
- [ ] Computation graph capture and compilation optimization: using TorchDynamo or other techniques to capture GNN operators and dynamic sparse shapes, enrich HLO to support lowering the sparse kernels mentioned above, and optimize based on input graph.
- [ ] Memory optimization: using techniques like CPU offload-ZERO.
- [ ] Distributed optimization: more efficient parallelism, cache..

### Alternatives

_No response_

### Additional context

_No response_

Contributor guide

No contributing guide indexed for this repository

Research direction

No files, tests, or entry points are named. Start by surveying the existing graph convolution and distributed training structure, then break the proposal into a narrowly scoped sparse-matrix task before attempting implementation. Done would require an agreed scope, working sparse kernels or modules, and evidence of distributed performance and memory improvements.

Written by the indexing model from the issue text.

Assessment

Tech stack
python, pytorch
Domain
distributed-systems, machine-learning
Issue type
Feature
Difficulty
5/5
Estimated time
Over a week
Activity status
Stale
Clarity
Needs clarification
Newbie friendliness
15/100

Get new issues in your inbox

A short digest of beginner-friendly GitHub issues.