Context Parallelism
Tìm hiểu cách xử lý song song ngữ cảnh phân phối các chuỗi dài trên nhiều GPU để giảm mức sử dụng bộ nhớ, mở rộng quy mô huấn luyện Transformer và hỗ trợ workload AI với tài liệu dài và video.
Context parallelism là kỹ thuật tính toán phân tán, chia một chuỗi đầu vào dài giữa nhiều bộ tăng tốc. Mỗi GPU xử lý một phần chuỗi trong khi phối hợp với các GPU khác trong quá trình attention. Cách này giảm bộ nhớ activation trên từng thiết bị, cho phép huấn luyện Transformer với đầu vào có thể vượt quá bộ nhớ của một GPU, chẳng hạn như tài liệu rất dài, video mở rộng hoặc tập hợp lớn các patch hình ảnh.
Khác với việc chỉ mở rộng cửa sổ ngữ cảnh của model, context parallelism không thay đổi lượng thông tin mà về mặt lý thuyết kiến trúc có thể tiếp nhận. Thay vào đó, kỹ thuật này giúp việc xử lý ngữ cảnh đó trở nên khả thi về mặt tính toán bằng cách phân phối chiều chuỗi.
Cách Context Parallelism hoạt động#
Giả sử một chuỗi có 32.000 token và context parallelism sử dụng bốn GPU. Ban đầu, mỗi thiết bị nhận khoảng 8.000 token và lưu các activation trung gian tương ứng.
Hầu hết phép toán, chẳng hạn như chuẩn hóa và các lớp feed-forward, có thể xử lý độc lập các phần chuỗi cục bộ này. Thách thức nằm ở cơ chế attention: một query cục bộ có thể cần attention đến key và value được lưu trên mọi thiết bị khác.
Do đó, các triển khai trao đổi các khối key-value, hay KV, giữa các GPU. Trong ring attention, mỗi thiết bị tính attention một phần bằng dữ liệu cục bộ, chuyển một khối KV đến thiết bị tiếp theo và lặp lại cho đến khi xử lý toàn bộ chuỗi. Hướng dẫn context parallel của PyTorch minh họa hành vi này thông qua attention phân tán với tích vô hướng được scale, trong khi gói context parallel của NVIDIA mô tả các tùy chọn giao tiếp dựa trên all-gather, reduce-scatter và ring.
Kết quả cuối cùng tương đương về mặt toán học với attention trên toàn chuỗi, có tính đến các sai khác số học thông thường, nhưng không có GPU đơn lẻ nào phải giữ mọi activation của chuỗi.
Tầm quan trọng của Context Parallelism#
Chuỗi dài gây ra hai vấn đề mở rộng chính. Thứ nhất, activation được lưu tiêu tốn nhiều bộ nhớ hơn khi độ dài chuỗi tăng. Thứ hai, self-attention tiêu chuẩn so sánh các token trong toàn chuỗi, tạo ra lượng tính toán và dữ liệu tạm thời đáng kể.
Context parallelism xử lý vấn đề bộ nhớ bằng cách phân chia activation giữa các thiết bị. Kỹ thuật này cũng có thể phân chia khối lượng tính toán attention, dù giao tiếp làm phát sinh chi phí mới. Các hệ thống hiệu quả chồng lấp việc truyền KV với tính toán bằng các thao tác như được ghi lại trong hướng dẫn các phép toán collective của NCCL.
Kỹ thuật này hữu ích nhất khi độ dài chuỗi, chứ không phải trọng số model hoặc kích thước batch, gây ra lỗi hết bộ nhớ. Đây là một phần của lĩnh vực rộng hơn về huấn luyện phân tán và thường được kết hợp với các chiến lược khác để đồng thời mở rộng nhiều chiều.
Context Parallelism và các kỹ thuật liên quan#
- Tensor parallelism chia các phép toán hoặc ma trận trọng số bên trong từng lớp riêng lẻ. Ngược lại, context parallelism chia token theo chiều chuỗi.
- Pipeline parallelism phân công các nhóm lớp model khác nhau cho các thiết bị khác nhau. Kỹ thuật này phân vùng độ sâu của model thay vì độ dài chuỗi.
- Data parallelism nhân bản model và cung cấp cho mỗi bản sao các ví dụ huấn luyện khác nhau. Tổng quan về distributed của PyTorch khuyến nghị phương pháp này khi toàn bộ model và từng mẫu đều vừa với một GPU.
- Sequence parallelism thường phân mảnh activation cho các phép toán được chọn liên quan đến tensor parallelism. Context parallelism áp dụng việc phân vùng chuỗi rộng hơn trên các đầu vào và activation của mạng.
Các phương pháp này bổ trợ lẫn nhau. Hướng dẫn chiến lược parallelism của NVIDIA cho thấy cách context, tensor, pipeline và data parallelism có thể tạo thành bố cục thiết bị đa chiều.
Ứng dụng thực tế#
-
AI xử lý tài liệu dài: Model ngôn ngữ pháp lý hoặc y tế có thể cần xử lý toàn bộ hồ sơ vụ việc, bệnh sử hoặc tài liệu kỹ thuật. Context parallelism phân phối hàng nghìn token của tài liệu trên các bộ tăng tốc, giảm áp lực lên bộ nhớ activation trong khi vẫn duy trì attention giữa các phần cách xa nhau.
-
Hiểu video dài và dữ liệu đa phương thức: Transformer xử lý video và model thị giác lớn có thể biểu diễn frame, patch hình ảnh, đoạn âm thanh và văn bản thành một chuỗi token dài. Phân phối chuỗi đó giúp model phân tích các bản ghi dài mà không cần giảm mạnh số lượng frame hoặc chi tiết không gian. Tổng quan context parallelism của AWS Neuron minh họa cách các nhóm bộ tăng tốc trao đổi các phân mảnh KV cho những khối lượng công việc có ngữ cảnh dài này.
Đối với các kiến trúc thị giác máy tính gọn nhẹ như Ultralytics YOLO26, context parallelism thường không cần thiết. Data parallelism tiêu chuẩn trên nhiều GPU thông qua quy trình huấn luyện model của Ultralytics nhìn chung là cách phù hợp hơn để tăng tốc huấn luyện.
Ứng dụng thực tế và sự đánh đổi#
Ví dụ tối giản sau sử dụng API context-parallel thử nghiệm của PyTorch và các tùy chọn kiểm soát attention với tích vô hướng được scale. Lưu ví dụ này thành cp_example.py và chạy trên hai GPU bằng torchrun --standalone --nproc-per-node=2 cp_example.py.
import os
import torch
import torch.distributed as dist
import torch.nn.functional as F
from torch.distributed.device_mesh import init_device_mesh
from torch.distributed.tensor.experimental import context_parallel
from torch.nn.attention import SDPBackend, sdpa_kernel
rank = int(os.environ["RANK"])
world_size = int(os.environ["WORLD_SIZE"])
torch.cuda.set_device(rank)
torch.cuda.manual_seed(0)
dist.init_process_group("nccl")
mesh = init_device_mesh("cuda", (world_size,))
qkv = [torch.randn(1, 4, 4096, 64, device="cuda", dtype=torch.bfloat16, requires_grad=True) for _ in range(3)]
with sdpa_kernel(SDPBackend.FLASH_ATTENTION), context_parallel(mesh, buffers=tuple(qkv), buffer_seq_dims=(2, 2, 2)):
output = F.scaled_dot_product_attention(*qkv, is_causal=True)
output.float().square().mean().backward()
dist.destroy_process_group()Ở đây, chiều 2 là chiều chuỗi, vì vậy mỗi tiến trình nhận một phân mảnh chuỗi trong khi attention được phối hợp trên lưới thiết bị.
Trong thực tế, kỹ sư nên xác nhận rằng mức giảm bộ nhớ lớn hơn chi phí giao tiếp, sử dụng kết nối tốc độ cao và benchmark các độ dài chuỗi đại diện. Đối với các dự án thị giác tiêu chuẩn, Ultralytics Platform cung cấp các quy trình cloud và cục bộ đơn giản hơn cho gán nhãn tập dữ liệu, huấn luyện, triển khai và giám sát mà không cần cấu hình context parallel thủ công.









