项目文件夹

文件
wehub-resource-sync 78ec5d9290
ci-cd / build-axolotl-uv (<nil>, 130, 13.0.0, linux/amd64,linux/arm64, 3.12, 2.11.0) (push) Has been cancelled
ci-cd / build-axolotl-uv (<nil>, 130, 13.0.0, true, linux/amd64,linux/arm64, 3.12, 2.12.0) (push) Has been cancelled
ci-cd / build-axolotl-cloud-uv (<nil>, 130, 13.0.0, linux/amd64,linux/arm64, 3.12, 2.11.0) (push) Has been cancelled
ci-cd / build-axolotl-cloud-uv (<nil>, 130, 13.0.0, true, linux/amd64,linux/arm64, 3.12, 2.12.0) (push) Has been cancelled
ci-cd / build-axolotl-cloud-no-tmux-uv (<nil>, 130, 13.0.0, linux/amd64,linux/arm64, 3.12, 2.11.0) (push) Has been cancelled
ci-cd / build-axolotl-cloud-no-tmux-uv (<nil>, 130, 13.0.0, true, linux/amd64,linux/arm64, 3.12, 2.12.0) (push) Has been cancelled
ci-cd-base / build-base-uv (130, 13.0.0, , Dockerfile-uv-base, linux/amd64,linux/arm64, 3.12, 2.12.0, 9.0 10.0 10.3 12.0+PTX) (push) Has been cancelled
ci-cd-base / build-base-uv (130, 13.0.0, , Dockerfile-uv-base, linux/amd64,linux/arm64, 3.12, 2.12.1, 9.0 10.0 10.3 12.0+PTX) (push) Has been cancelled
ci-cd-base / build-base-uv (130, 13.0.0, , Dockerfile-uv-base, linux/amd64,linux/arm64, 3.12, 2.13.0, 9.0 10.0 10.3 12.0+PTX) (push) Has been cancelled
ci-cd-base / build-base-uv (132, 13.2.1, , Dockerfile-uv-base, linux/amd64,linux/arm64, 3.12, 2.13.0, 9.0 10.0 10.3 12.0+PTX, https://download.pytorch.org/whl/cu132) (push) Has been cancelled
ci-cd-base / build-base-uv (130, 13.0.0, , Dockerfile-uv-base, linux/amd64,linux/arm64, 3.12, 2.11.0, 9.0 10.0 10.3 12.0+PTX) (push) Has been cancelled
Tests / PyTest (3.12, 2.12.1) (push) Has been cancelled
Tests / PyTest (3.12, 2.13.0) (push) Has been cancelled
docker-e2e-tests / gate-skip-e2e (push) Has been cancelled
docker-e2e-tests / docker-e2e-tests-1st (<nil>, 130, 13.0.0, 1, 3.12, 2.12.1) (push) Has been cancelled
docker-e2e-tests / docker-e2e-tests (<nil>, 130, 13.0.0, 1, 3.12, 2.11.0) (push) Has been cancelled
docker-e2e-tests / docker-e2e-kernel-tests (<nil>, 130, 13.0.0, 1, 3.12, 2.11.0) (push) Has been cancelled
docker-e2e-tests / docker-e2e-kernel-tests (<nil>, 130, 13.0.0, 1, 3.12, 2.12.1) (push) Has been cancelled
docker-e2e-tests / docker-e2e-cleanup (<nil>, 130, 13.0.0, 1, 3.12, 2.12.1) (push) Has been cancelled
Publish Docs / build-deploy (push) Has been cancelled
Tests / PyTest from Source Dist (3.12, 2.11.0) (push) Has been cancelled
Tests / PyTest from Source Dist (3.12, 2.12.1) (push) Has been cancelled
Tests / PyTest from Source Dist (3.12, 2.13.0) (push) Has been cancelled
Tests / pre-commit (push) Has been cancelled
Tests / Prefetch S3 once to prime the CDN cache (push) Has been cancelled
Tests / PyTest (3.12, 2.11.0) (push) Has been cancelled
chore: import upstream snapshot with attribution
2026-07-13 12:48:45 +08:00

56 行
1.9 KiB
Python

from functools import partial
import torch
from torch import nn
from torch.utils.checkpoint import checkpoint
from transformers import GradientCheckpointingLayer
from axolotl.monkeypatch.checkpoint_activation_offload import (
CheckpointHiddenStatesOffload,
)
class TinyCheckpointLayer(GradientCheckpointingLayer):
def forward(self, hidden_states):
hidden_states = hidden_states.sin()
return hidden_states * hidden_states
def test_checkpoint_offload_marks_non_reentrant_checkpoint_input():
device = "cuda" if torch.cuda.is_available() else "cpu"
layer = TinyCheckpointLayer()
layer.gradient_checkpointing = True
layer._gradient_checkpointing_func = partial(checkpoint, use_reentrant=False)
layer.to(device)
layer.train()
hidden_states = torch.randn(4, 8, device=device, requires_grad=True)
offload = CheckpointHiddenStatesOffload(use_streams=False, min_offload_size=0)
with offload:
loss = layer(hidden_states).sum()
loss.backward()
assert hidden_states.grad is not None
assert offload.stats.marked_tensors == 1
assert offload.stats.saved_tensors_seen >= offload.stats.marked_tensors
if hidden_states.device.type == "cuda":
assert offload.stats.offloaded_tensors == 1
assert offload.stats.restored_tensors == 1
else:
assert offload.stats.skipped_marked_tensors == 1
def test_checkpoint_offload_ignores_unmarked_saved_tensors():
hidden_states = torch.randn(4, 8, requires_grad=True)
linear = nn.Linear(8, 8)
offload = CheckpointHiddenStatesOffload(use_streams=False, min_offload_size=0)
with offload:
loss = linear(hidden_states).square().sum()
loss.backward()
assert hidden_states.grad is not None
assert offload.stats.saved_tensors_seen > 0
assert offload.stats.marked_tensors == 0
assert offload.stats.offloaded_tensors == 0