Skip to content

Commit 03d15b2

Browse files
committed
update: codes
1 parent 359c3a3 commit 03d15b2

1 file changed

Lines changed: 23 additions & 0 deletions

File tree

tests/test_optimizers.py

Lines changed: 23 additions & 0 deletions
Original file line numberDiff line numberDiff line change
@@ -770,6 +770,29 @@ def test_flash_adamw_loads_compressed_state_as_uncompressed_state():
770770
assert 'exp_avg' in raw_optimizer.state_dict()['state'][0]
771771

772772

773+
def test_flash_adamw_uncompressed_state_dict_reloads_as_quantized_state():
774+
param = nn.Parameter(torch.tensor([1.0, -2.0, 3.0]))
775+
param.grad = torch.tensor([0.2, -0.3, 0.4])
776+
777+
optimizer = load_optimizer('flashadamw')([param], lr=1e-2, weight_decay=0.0, compress_state_dict=False)
778+
optimizer.step()
779+
780+
state_dict = optimizer.state_dict()
781+
saved_state = state_dict['state'][0]
782+
assert 'exp_avg' in saved_state
783+
assert 'exp_avg::quantized' not in saved_state
784+
785+
new_optimizer = load_optimizer('flashadamw')(
786+
[nn.Parameter(param.detach().clone())],
787+
lr=1e-2,
788+
weight_decay=0.0,
789+
)
790+
new_optimizer.load_state_dict(state_dict)
791+
new_state = next(iter(new_optimizer.state.values()))
792+
assert 'exp_avg::quantized' in new_state
793+
assert 'exp_avg' not in new_state
794+
795+
773796
@pytest.mark.parametrize(('master_weight_bits', 'error_dtype'), [(24, torch.int8), (32, torch.int16)])
774797
def test_flash_adamw_master_weight_bits(master_weight_bits, error_dtype):
775798
model = nn.Linear(2, 1).bfloat16()

0 commit comments

Comments
 (0)