diff --git a/README.md b/README.md
index 20c2c02..4141f79 100644
--- a/README.md
+++ b/README.md
@@ -183,15 +183,18 @@ Before trying your own inputs, we highly recommend going through the sanity chec
| **T2V** | | | |
| **V2V** | | | |
-### ✨ Parallel Inference on Multiple GPUs
-For example, let's take Helios-Base with 2 GPUs.
+### ✨ Context Parallelism on Multiple GPUs
+Helios supports various Context Parallelism mechanisms, including Ulysses Attention, Ring Attention, Unified Attention, and Ulysses Anything Attention. For more details, please refer to the [documentation](https://huggingface.co/docs/diffusers/v0.37.0/en/training/distributed_inference#context-parallelism).
+
+For example, let's take Helios-Base with 4 GPUs.
Click to expand the code
```bash
- CUDA_VISIBLE_DEVICES=0,1 torchrun --nproc_per_node 2 infer_helios.py \
- --enable_parallelism \
+ CUDA_VISIBLE_DEVICES=0,1,2,3 torchrun --nproc_per_node 4 infer_helios.py \
+ --enable_parallelism \ # remember to enable this config
+ --cp_backend "ulysses" \ # ["ring", "ulysses", "unified", "ulysses_anything"]
--base_model_path "BestWishYsh/Helios-Base" \
--transformer_path "BestWishYsh/Helios-Base" \
--sample_type "t2v" \
diff --git a/infer_helios.py b/infer_helios.py
index ec2cb10..c19328d 100644
--- a/infer_helios.py
+++ b/infer_helios.py
@@ -139,6 +139,16 @@ def parse_args():
default=None,
)
+ # === Context parallelism ===
+ # Please refer to https://huggingface.co/docs/diffusers/v0.37.0/en/training/distributed_inference#context-parallelism
+ parser.add_argument(
+ "--cp_backend",
+ type=str,
+ choices=["ring", "ulysses", "unified", "ulysses_anything"],
+ default="ulysses",
+ help="Context parallel backend to use.",
+ )
+
return parser.parse_args()
@@ -159,7 +169,10 @@ def main():
os.makedirs(args.output_folder, exist_ok=True)
if dist.is_available() and "RANK" in os.environ:
- dist.init_process_group(backend="nccl")
+ if args.cp_backend == "ulysses_anything":
+ dist.init_process_group(backend="cpu:gloo,cuda:nccl")
+ else:
+ dist.init_process_group(backend="nccl")
rank = dist.get_rank()
device = torch.device("cuda", rank % torch.cuda.device_count())
world_size = dist.get_world_size()
@@ -271,8 +284,18 @@ def main():
pipe = pipe.to(device)
if world_size > 1 and args.enable_parallelism:
- # transformer.set_attention_backend("flash")
- pipe.transformer.enable_parallelism(config=ContextParallelConfig(ulysses_degree=world_size))
+ if args.cp_backend == "ring":
+ cp_config = ContextParallelConfig(ring_degree=world_size)
+ elif args.cp_backend == "unified":
+ cp_config = ContextParallelConfig(ring_degree=world_size // 2, ulysses_degree=world_size // 2)
+ elif args.cp_backend == "ulysses":
+ cp_config = ContextParallelConfig(ulysses_degree=world_size)
+ elif args.cp_backend == "ulysses_anything":
+ cp_config = ContextParallelConfig(ulysses_degree=world_size, ulysses_anything=True)
+ else:
+ raise ValueError(f"Unsupported cp_backend: {args.cp_backend}")
+
+ pipe.transformer.enable_parallelism(config=cp_config)
if args.debug_mode:
diff --git a/scripts/inference/helios-base_i2v.sh b/scripts/inference/helios-base_i2v.sh
index 5ccf67d..f1afe9d 100644
--- a/scripts/inference/helios-base_i2v.sh
+++ b/scripts/inference/helios-base_i2v.sh
@@ -1,6 +1,7 @@
# Example: Running inference with 2-GPU parallelism
# CUDA_VISIBLE_DEVICES=0,1 torchrun --nproc_per_node 2 infer_helios.py \
# --enable_parallelism \
+# --cp_backend "ulysses" \ # ["ring", "ulysses", "unified", "ulysses_anything"]
CUDA_VISIBLE_DEVICES=0 python infer_helios.py \
--base_model_path "BestWishYsh/Helios-Base" \
diff --git a/scripts/inference/helios-base_t2v.sh b/scripts/inference/helios-base_t2v.sh
index 1d0914a..57a173a 100644
--- a/scripts/inference/helios-base_t2v.sh
+++ b/scripts/inference/helios-base_t2v.sh
@@ -1,6 +1,7 @@
# Example: Running inference with 2-GPU parallelism
# CUDA_VISIBLE_DEVICES=0,1 torchrun --nproc_per_node 2 infer_helios.py \
# --enable_parallelism \
+# --cp_backend "ulysses" \ # ["ring", "ulysses", "unified", "ulysses_anything"]
CUDA_VISIBLE_DEVICES=0 python infer_helios.py \
--base_model_path "BestWishYsh/Helios-Base" \
diff --git a/scripts/inference/helios-base_v2v.sh b/scripts/inference/helios-base_v2v.sh
index 4188f00..87e2c0d 100644
--- a/scripts/inference/helios-base_v2v.sh
+++ b/scripts/inference/helios-base_v2v.sh
@@ -1,6 +1,7 @@
# Example: Running inference with 2-GPU parallelism
# CUDA_VISIBLE_DEVICES=0,1 torchrun --nproc_per_node 2 infer_helios.py \
# --enable_parallelism \
+# --cp_backend "ulysses" \ # ["ring", "ulysses", "unified", "ulysses_anything"]
CUDA_VISIBLE_DEVICES=0 python infer_helios.py \
--base_model_path "BestWishYsh/Helios-Base" \
diff --git a/scripts/inference/helios-distilled_i2v.sh b/scripts/inference/helios-distilled_i2v.sh
index ae34f7a..aecd398 100644
--- a/scripts/inference/helios-distilled_i2v.sh
+++ b/scripts/inference/helios-distilled_i2v.sh
@@ -1,6 +1,7 @@
# Example: Running inference with 2-GPU parallelism
# CUDA_VISIBLE_DEVICES=0,1 torchrun --nproc_per_node 2 infer_helios.py \
# --enable_parallelism \
+# --cp_backend "ulysses" \ # ["ring", "ulysses", "unified", "ulysses_anything"]
CUDA_VISIBLE_DEVICES=0 python infer_helios.py \
--base_model_path "BestWishYsh/Helios-Distilled" \
diff --git a/scripts/inference/helios-distilled_t2v.sh b/scripts/inference/helios-distilled_t2v.sh
index 46d6079..562d2ca 100644
--- a/scripts/inference/helios-distilled_t2v.sh
+++ b/scripts/inference/helios-distilled_t2v.sh
@@ -1,6 +1,7 @@
# Example: Running inference with 2-GPU parallelism
# CUDA_VISIBLE_DEVICES=0,1 torchrun --nproc_per_node 2 infer_helios.py \
# --enable_parallelism \
+# --cp_backend "ulysses" \ # ["ring", "ulysses", "unified", "ulysses_anything"]
CUDA_VISIBLE_DEVICES=0 python infer_helios.py \
--base_model_path "BestWishYsh/Helios-Distilled" \
diff --git a/scripts/inference/helios-distilled_v2v.sh b/scripts/inference/helios-distilled_v2v.sh
index 299b596..4aa3134 100644
--- a/scripts/inference/helios-distilled_v2v.sh
+++ b/scripts/inference/helios-distilled_v2v.sh
@@ -1,6 +1,7 @@
# Example: Running inference with 2-GPU parallelism
# CUDA_VISIBLE_DEVICES=0,1 torchrun --nproc_per_node 2 infer_helios.py \
# --enable_parallelism \
+# --cp_backend "ulysses" \ # ["ring", "ulysses", "unified", "ulysses_anything"]
CUDA_VISIBLE_DEVICES=0 python infer_helios.py \
--base_model_path "BestWishYsh/Helios-Distilled" \
diff --git a/scripts/inference/helios-mid_i2v.sh b/scripts/inference/helios-mid_i2v.sh
index e512aff..2d4e89e 100644
--- a/scripts/inference/helios-mid_i2v.sh
+++ b/scripts/inference/helios-mid_i2v.sh
@@ -1,6 +1,7 @@
# Example: Running inference with 2-GPU parallelism
# CUDA_VISIBLE_DEVICES=0,1 torchrun --nproc_per_node 2 infer_helios.py \
# --enable_parallelism \
+# --cp_backend "ulysses" \ # ["ring", "ulysses", "unified", "ulysses_anything"]
CUDA_VISIBLE_DEVICES=0 python infer_helios.py \
--base_model_path "BestWishYsh/Helios-Mid" \
diff --git a/scripts/inference/helios-mid_t2v.sh b/scripts/inference/helios-mid_t2v.sh
index b5f3f20..7ad84e3 100644
--- a/scripts/inference/helios-mid_t2v.sh
+++ b/scripts/inference/helios-mid_t2v.sh
@@ -1,6 +1,7 @@
# Example: Running inference with 2-GPU parallelism
# CUDA_VISIBLE_DEVICES=0,1 torchrun --nproc_per_node 2 infer_helios.py \
# --enable_parallelism \
+# --cp_backend "ulysses" \ # ["ring", "ulysses", "unified", "ulysses_anything"]
CUDA_VISIBLE_DEVICES=0 python infer_helios.py \
--base_model_path "BestWishYsh/Helios-Mid" \
diff --git a/scripts/inference/helios-mid_v2v.sh b/scripts/inference/helios-mid_v2v.sh
index f60911b..2a41a99 100644
--- a/scripts/inference/helios-mid_v2v.sh
+++ b/scripts/inference/helios-mid_v2v.sh
@@ -1,6 +1,7 @@
# Example: Running inference with 2-GPU parallelism
# CUDA_VISIBLE_DEVICES=0,1 torchrun --nproc_per_node 2 infer_helios.py \
# --enable_parallelism \
+# --cp_backend "ulysses" \ # ["ring", "ulysses", "unified", "ulysses_anything"]
CUDA_VISIBLE_DEVICES=0 python infer_helios.py \
--base_model_path "BestWishYsh/Helios-Mid" \