Gradient Checkpointing
Tìm hiểu cách gradient checkpointing giúp giảm bộ nhớ GPU bằng cách tính toán lại các activation trong quá trình backpropagation, kèm theo các ví dụ PyTorch, đánh đổi và hướng dẫn huấn luyện thực tế.
Gradient checkpointing là một kỹ thuật huấn luyện tiết kiệm bộ nhớ, chỉ lưu trữ các activation trung gian được chọn từ luồng tiến (forward pass) và tính toán lại các activation còn lại trong quá trình backpropagation. Còn được gọi là activation checkpointing, kỹ thuật này đánh đổi thêm thời gian tính toán để sử dụng bộ nhớ đỉnh (peak memory) thấp hơn. Mặc dù có tên gọi như vậy, kỹ thuật này thực hiện checkpoint các activation chứ không phải gradient của tham số hay tệp mô hình, điều này làm cho nó đặc biệt có giá trị khi các tensor activation ngăn cản mạng nơ-ron vừa vặn trong GPU memory khả dụng.
Cách thức hoạt động của Gradient Checkpointing#
Trong một luồng tiến tiêu chuẩn, mạng nơ-ron tính toán các tensor trung gian gọi là activations. Hệ thống vi phân tự động giữ lại các activation cần thiết để tính toán gradient sau đó, như được mô tả trong PyTorch autograd mechanics. Các 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 trữ này tiêu thụ lượng bộ nhớ đáng kể.
Gradient checkpointing chia mạng thành các đoạn:
- Luồng tiến lưu trữ các đầu vào hoặc activation ranh giới cho các đoạn được chọn.
- Các activation trung gian khác bên trong các đoạn đó bị loại bỏ.
- Trong luồng lùi (backward pass), mỗi đoạn được checkpoint sẽ chạy tiến lại một lần nữa để tái tạo các giá trị bị thiếu.
- Các activation được tái tạo được sử dụng ngay lập tức để tính toán gradient.
PyTorch activation checkpointing API hiển thị 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 JAX gradient checkpointing and rematerialization và TensorFlow tape checkpointing.
Việc đặt 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 toán hơn. Việc đặt có chọn lọc xung quanh các khối nặng activation có thể mang lại sự cân bằng tốt hơn so với việc tính toán lại toàn bộ mạng.
Đánh đổi và các kỹ thuật liên quan#
Gradient checkpointing thường giữ nguyên các tham số mô hình, gradient và trạng thái bộ tối ưu hóa (optimizer states). Mục tiêu chính của nó là bộ nhớ activation, và việc huấn luyện sẽ trở nên chậm hơn vì một số phép tính tiến chạy hai lần. Kết quả chính xác phụ thuộc vào kiến trúc, ranh giới checkpoint, hình dạng batch và phần cứng.
Nó khác biệt với một số kỹ thuật liên quan:
- Gradient accumulation xử lý nhiều microbatch trước khi cập nhật trọng số mô hình, tạo ra một batch hiệu quả lớn hơn mà không cần tải mọi mẫu cùng một lúc. Thay vào đó, gradient checkpointing làm giảm các activation được giữ lại cho mỗi microbatch.
- Mixed precision sử dụng các kiểu dữ liệu độ precision thấp hơn cho các phép toán được chọn. PyTorch automatic mixed precision workflow có thể giảm bộ nhớ và tăng tốc các phép toán tương thích, trong khi checkpointing cố ý bổ sung thêm tính toán.
- Giảm Batch size làm giảm bộ nhớ bằng cách xử lý ít mẫu hơn cùng nhau. Checkpointing có thể cho phép một batch lớn hơn hoặc độ phân giải đầu vào lớn hơn tiếp tục khả thi.
- Training checkpoints lưu trọng số và trạng thái bộ tối ưu hóa để khôi phục hoặc suy luận sau đó. PyTorch model checkpoint guide mô tả cơ chế duy trì này, vốn không liên quan đến việc tính toán lại activation.
Mã được checkpoint nên nhất quán về mặt chức năng giữa các lần thực thi tiến ban đầu và được tính toán lại. Trạng thái biến đổi (mutable state), việc truyền thiết bị hoặc tính ngẫu nhiên không được kiểm soát bên trong một vùng được checkpoint có thể gây ra lỗi hoặc gradient không chính xác. Trong các transformer kiểu decoder, key-value cache cũng có thể cần được vô hiệu hóa trong quá trình huấn luyện được checkpoint vì trạng thái suy luận được lưu trong cache có thể xung đột với việc xây dựng lại đồ thị tiến.
Ví dụ với PyTorch#
Ví dụ sau đây thực hiện checkpoint một khối 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 bắt buộc được giữ lại; các activation bên trong của nó được tái tạo trong loss.backward(). Các dự án thực tế nên so sánh các lần chạy có và không có checkpoint bằng cách sử dụng một phép đo như PyTorch peak allocated GPU memory.
Các ứng dụng trong 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 không và kiểm tra công nghiệp có thể huấn luyện trên các hình ảnh lớn có bản đồ đặc trưng tiêu thụ nhiều bộ nhớ hơn trọng số mô hình. Checkpoint các giai đoạn backbone được chọn có thể giữ nguyên độ phân giải hình ảnh hoặc cho phép các mẫu bổ sung trên mỗi batch thay vì thu nhỏ đầu vào một cách quá mức.
-
Huấn luyện transformer chuỗi dài: Lưu trữ activation tăng trưởng nhanh chóng khi độ dài chuỗi và số lớp tăng lên. Tính toán lại các khối transformer có thể làm cho các ngữ cảnh dài hơn hoặc các microbatch lớn hơn vừa vặn trên cùng một bộ gia tốc. NVIDIA activation recomputation guide minh họa việc tính toán lại toàn bộ và có chọn lọc cho các lớp transformer.
Hướng dẫn thực tế#
Sử dụng gradient checkpointing khi profiling cho thấy activations, chứ không phải các tham số hay trạng thái bộ tối ưu hóa, chiếm ưu thế về bộ nhớ. Bắt đầu với các khối lặp lại lớn và benchmark bộ nhớ đỉnh, thời gian lặp và hành vi xác thực trước khi mở rộng phạm vi.
Đối với Ultralytics YOLO training configurations được tài liệu hóa, các điều khiển bộ nhớ hàng đầu bao gồm mixed precision tự động, kích thước hình ảnh và kích thước batch vật lý. Ultralytics AutoBatch reference giải thích việc lựa chọn batch tự động dựa trên bộ nhớ GPU khả dụng. Khi phần cứng cục bộ vẫn không đủ, Ultralytics Platform cloud training cung cấp các GPU đám mây có thể cấu hình cho các lượt huấn luyện được quản lý. Gradient checkpointing có thể bổ sung cho các điều khiển này khi kiến trúc PyTorch tùy chỉnh yêu cầu quản lý bộ nhớ activation chi tiết hơn.






