Gradient Checkpointing
Tìm hiểu cách gradient checkpointing giảm mức sử dụng bộ nhớ GPU bằng cách tính toán lại các activation trong quá trình lan truyền ngược, cùng ví dụ PyTorch, các đánh đổi và hướng dẫn huấn luyện thực tiễn.
Gradient checkpointing là kỹ thuật huấn luyện tiết kiệm bộ nhớ, chỉ lưu một số activation trung gian được chọn trong lượt truyền xuôi và tính toán lại các activation còn lại trong quá trình lan truyền ngược. Kỹ thuật này còn được gọi là activation checkpointing, đánh đổi việc tăng tính toán để giảm mức sử dụng bộ nhớ đỉnh. Dù có tên như vậy, kỹ thuật này checkpoint activation chứ không phải gradient tham số hay tệp model, nên đặc biệt hữu ích khi tensor activation khiến mạng neural không thể vừa trong bộ nhớ GPU hiện có.
Cách Gradient Checkpointing hoạt động#
Trong lượt truyền xuôi thông thường, mạng neural tính toán các tensor trung gian gọi là activation. Hệ thống tự động vi phân giữ lại các activation cần thiết để tính gradient sau này, như mô tả trong cơ chế autograd của PyTorch. Mạng sâu, batch lớn, hình ảnh độ phân giải cao và chuỗi đầu vào dài có thể khiến các tensor được lưu này tiêu tốn nhiều bộ nhớ.
Gradient checkpointing chia mạng thành các phân đoạn:
- Lượt truyền xuôi lưu đầu vào hoặc activation ranh giới của các phân đoạn được chọn.
- Các activation trung gian khác bên trong những phân đoạn đó sẽ bị loại bỏ.
- Trong lượt truyền ngược, mỗi phân đoạn được checkpoint chạy lại theo chiều xuôi để khôi phục các giá trị bị thiếu.
- Các activation được khôi phục sẽ được dùng ngay để tính gradient.
API activation checkpointing của PyTorch cung cấp hành vi này thông qua torch.utils.checkpoint. Các khái niệm tương đương xuất hiện dưới dạng gradient checkpointing và rematerialization của JAX và tape checkpointing của TensorFlow.
Vị trí checkpoint quyết định sự đánh đổi. Checkpoint nhiều vùng hơn thường tiết kiệm nhiều bộ nhớ hơn nhưng lặp lại nhiều phép tính hơn. Bố trí checkpoint có chọn lọc quanh các khối tiêu tốn nhiều activation có thể cân bằng tốt hơn so với việc tính toán lại toàn bộ mạng.
Sự đánh đổi và các kỹ thuật liên quan#
Gradient checkpointing thường không làm thay đổi tham số model, gradient và trạng thái optimizer. Kỹ thuật này chủ yếu nhắm đến bộ nhớ activation, đồng thời làm chậm quá trình huấn luyện vì một số phép tính truyền xuôi được chạy hai lần. Kết quả cụ thể phụ thuộc vào kiến trúc, ranh giới checkpoint, hình dạng batch và phần cứng.
Kỹ thuật này khác với một số kỹ thuật liên quan:
- Tích lũy gradient xử lý nhiều microbatch trước khi cập nhật trọng số model, tạo batch hiệu dụng lớn hơn mà không cần nạp đồng thời mọi mẫu. Ngược lại, gradient checkpointing giảm số activation được giữ lại cho mỗi microbatch.
- Độ chính xác hỗn hợp sử dụng kiểu dữ liệu có độ chính xác thấp hơn cho một số phép toán được chọn. Quy trình automatic mixed precision của PyTorch có thể giảm bộ nhớ và tăng tốc các phép toán tương thích, trong khi checkpointing chủ động làm tăng lượng tính toán.
- Giảm kích thước batch giảm bộ nhớ bằng cách xử lý đồng thời ít mẫu hơn. Checkpointing có thể giúp duy trì kích thước batch hoặc độ phân giải đầu vào lớn hơn.
- Checkpoint huấn luyện lưu trọng số và trạng thái optimizer để khôi phục hoặc suy luận sau này. Hướng dẫn checkpoint model của PyTorch mô tả cơ chế lưu trữ này, không liên quan đến việc tính toán lại activation.
Mã được checkpoint cần nhất quán về chức năng giữa lần thực thi truyền xuôi ban đầu và lần tính toán lại. Trạng thái có thể thay đổi, việc chuyển thiết bị hoặc tính ngẫu nhiên không được kiểm soát bên trong vùng checkpoint có thể gây lỗi hoặc tạo gradient không chính xác. Trong Transformer kiểu decoder, có thể cần tắt cache key-value trong quá trình huấn luyện có checkpoint vì trạng thái suy luận đã lưu có thể xung đột với việc xây dựng lại đồ thị truyền xuôi.
Ví dụ PyTorch#
Ví dụ sau checkpoint một khối tiêu tốn nhiều bộ nhớ trong một bước huấn luyện:
import torch
from torch import nn
from torch.utils.checkpoint import checkpoint
torch.manual_seed(0)
block = nn.Sequential(
nn.Linear(1024, 4096),
nn.ReLU(),
nn.Linear(4096, 1024),
)
optimizer = torch.optim.AdamW(block.parameters())
inputs = torch.randn(8, 1024, requires_grad=True)
targets = torch.zeros_like(inputs)
optimizer.zero_grad(set_to_none=True)
outputs = checkpoint(block, inputs, use_reentrant=False)
loss = nn.functional.mse_loss(outputs, targets)
loss.backward()
optimizer.step()Chỉ các đầu vào của khối và thông tin ranh giới cần thiết được giữ lại; các activation bên trong được khôi phục trong loss.backward(). Trong các dự án thực tế, nên so sánh các lần chạy có và không có checkpointing bằng phép đo như bộ nhớ GPU được cấp phát đỉnh của PyTorch.
Ứng dụng thực tế#
-
Thị giác máy tính độ phân giải cao: Phân đoạn y tế, phát hiện trên ảnh hàng không và kiểm tra công nghiệp có thể cần huấn luyện trên ảnh lớn, trong đó feature map tiêu tốn nhiều bộ nhớ hơn trọng số model. Checkpoint các giai đoạn backbone được chọn có thể duy trì độ phân giải ảnh hoặc cho phép tăng số mẫu mỗi batch thay vì giảm mạnh kích thước đầu vào.
-
Huấn luyện Transformer với chuỗi dài: Bộ nhớ lưu activation tăng nhanh khi độ dài chuỗi và số lớp tăng. Tính toán lại các khối Transformer có thể giúp ngữ cảnh dài hơn hoặc microbatch lớn hơn vừa với cùng một bộ tăng tốc. Hướng dẫn tính toán lại activation của NVIDIA minh họa việc tính toán lại toàn phần và có chọn lọc cho các lớp Transformer.
Hướng dẫn thực tiễn#
Sử dụng gradient checkpointing khi profiling cho thấy activation, chứ không phải tham số hoặc trạng thái optimizer, chiếm phần lớn bộ nhớ. Bắt đầu với các khối lớn được lặp lại và benchmark bộ nhớ đỉnh, thời gian mỗi vòng lặp và hành vi validation trước khi mở rộng phạm vi.
Đối với các cấu hình huấn luyện Ultralytics YOLO được tài liệu hóa, các biện pháp kiểm soát bộ nhớ hàng đầu gồm automatic mixed precision, kích thước ảnh và kích thước batch vật lý. Tài liệu tham khảo Ultralytics AutoBatch giải thích cách tự động chọn batch dựa trên bộ nhớ GPU khả dụng. Khi phần cứng cục bộ vẫn không đủ, huấn luyện trên cloud bằng Ultralytics Platform cung cấp GPU cloud có thể cấu hình cho các lượt huấn luyện được quản lý. Gradient checkpointing có thể bổ trợ các biện pháp này khi kiến trúc PyTorch tùy chỉnh cần quản lý bộ nhớ activation chi tiết hơn.









