Skip to content
Open
Show file tree
Hide file tree
Changes from all commits
Commits
File filter

Filter by extension

Filter by extension

Conversations
Failed to load comments.
Loading
Jump to
Jump to file
Failed to load files.
Loading
Diff view
Diff view
3 changes: 3 additions & 0 deletions colossalai/tensor/padded_tensor/api.py
Original file line number Diff line number Diff line change
Expand Up @@ -14,18 +14,21 @@ def _hijack_detach_and_clone(ptensor: torch.Tensor) -> torch.Tensor:
ptensor._unpad_detach = ptensor.detach
ptensor._unpad_clone = ptensor.clone

# the copies must be full padded tensors too, otherwise to_unpadded_tensor() fails on them
def new_detach(self):
t_ = self._unpad_detach()
t_._padding_dim = self._padding_dim
t_._origin_length = self._origin_length
t_._current_length = self._current_length
_hijack_detach_and_clone(t_)
return t_

def new_clone(self, *args, **kwargs):
t_ = self._unpad_clone(*args, **kwargs)
t_._padding_dim = self._padding_dim
t_._origin_length = self._origin_length
t_._current_length = self._current_length
_hijack_detach_and_clone(t_)
return t_

# bind the new methods to the tensor
Expand Down
Original file line number Diff line number Diff line change
@@ -1,8 +1,11 @@
import glob

import pytest
import torch
import torch.distributed as dist
from packaging.version import Version
from torch.optim import Adam
from transformers import GPT2Config, GPT2LMHeadModel
from utils import shared_tempdir

import colossalai
Expand Down Expand Up @@ -143,9 +146,39 @@ def _preprocess_data(data):
clear_layout_converter()


@clear_cache_before_run()
@parameterize("shard", [True, False])
@parameterize("test_config", [{"tp_size": 1, "pp_size": 1}, {"tp_size": 1, "pp_size": 2, "num_microbatches": 2}])
def exam_state_dict_padded_vocab(shard: bool, test_config: dict):
# a vocab size that is not a multiple of make_vocab_size_divisible_by (64), so the embedding and
# lm_head weights are padded tensors, which must be unpadded when saving (#6253)
config = GPT2Config(n_layer=2, n_head=4, n_embd=64, vocab_size=1000, n_positions=64)
booster = Booster(plugin=HybridParallelPlugin(**test_config, precision="fp16", initial_scale=1))
model = GPT2LMHeadModel(config).cuda()
model, *_ = booster.boost(model, Adam(model.parameters(), lr=1e-3))

with shared_tempdir() as tempdir:
model_ckpt_path = f"{tempdir}/model"
booster.save_model(model, model_ckpt_path, shard=shard)
dist.barrier()
if dist.get_rank() == 0:
# the checkpoint must hold the original (unpadded) vocab size
files = glob.glob(f"{model_ckpt_path}/*.bin") if shard else [model_ckpt_path]
saved = {}
for f in files:
saved.update(torch.load(f, map_location="cpu"))
assert saved["transformer.wte.weight"].shape == (config.vocab_size, config.n_embd)
new_model = GPT2LMHeadModel(config).cuda()
new_model, *_ = booster.boost(new_model, Adam(new_model.parameters(), lr=1e-3))
booster.load_model(new_model, model_ckpt_path)
check_state_dict_equal(model.unwrap().state_dict(), new_model.unwrap().state_dict())
dist.barrier()


def run_dist(rank, world_size, port):
colossalai.launch(rank=rank, world_size=world_size, host="localhost", port=port, backend="nccl")
exam_state_dict()
exam_state_dict_padded_vocab()


@pytest.mark.dist
Expand Down
15 changes: 15 additions & 0 deletions tests/test_tensor/test_padded_tensor.py
Original file line number Diff line number Diff line change
Expand Up @@ -36,6 +36,21 @@ def check_padded_tensor(rank, world_size, port):
assert global_tensor.shape == original_tensor.shape


def test_unpad_detached_and_cloned_tensor():
# detach() / clone() of a padded tensor must give a padded tensor that can be unpadded on its own,
# this is what checkpoint saving does before writing a parameter
original = torch.rand(10, 4)
padded = to_padded_tensor(original.clone(), current_length=16, padding_dim=0)
for tensor_copy in (padded.detach(), padded.clone(), padded.detach().clone()):
assert is_padded_tensor(tensor_copy)
assert tensor_copy.shape == (16, 4)
unpadded = to_unpadded_tensor(tensor_copy)
assert not is_padded_tensor(unpadded)
assert torch.equal(unpadded, original)
# the source tensor is still padded
assert is_padded_tensor(padded) and padded.shape == (16, 4)


@rerun_if_address_is_in_use()
def test_padded_tensor():
world_size = 4
Expand Down