starting to structure out test suite
This commit is contained in:
parent
b60c138131
commit
598dda3a5b
4 changed files with 74 additions and 0 deletions
0
tests/__init__.py
Normal file
0
tests/__init__.py
Normal file
74
tests/conftest.py
Normal file
74
tests/conftest.py
Normal file
|
|
@ -0,0 +1,74 @@
|
|||
# Copyright 2023-present Daniel Han-Chen & the Unsloth team. All rights reserved.
|
||||
#
|
||||
# Licensed under the Apache License, Version 2.0 (the "License");
|
||||
# you may not use this file except in compliance with the License.
|
||||
# You may obtain a copy of the License at
|
||||
#
|
||||
# http://www.apache.org/licenses/LICENSE-2.0
|
||||
#
|
||||
# Unless required by applicable law or agreed to in writing, software
|
||||
# distributed under the License is distributed on an "AS IS" BASIS,
|
||||
# WITHOUT WARRANTIES OR CONDITIONS OF ANY KIND, either express or implied.
|
||||
# See the License for the specific language governing permissions and
|
||||
# limitations under the License.
|
||||
|
||||
|
||||
import gc
|
||||
import random
|
||||
from contextlib import contextmanager
|
||||
|
||||
import os
|
||||
import pytest
|
||||
import numpy as np
|
||||
import torch
|
||||
import torch._dynamo as dynamo
|
||||
|
||||
|
||||
@contextmanager
|
||||
def set_seed(seed: int = 0):
|
||||
np.random.seed(42)
|
||||
torch.manual_seed(42)
|
||||
torch.backends.cudnn.deterministic = True
|
||||
torch.backends.cudnn.benchmark = False
|
||||
torch.use_deterministic_algorithms(True)
|
||||
os.environ['CUBLAS_WORKSPACE_CONFIG'] = ":4096:8"
|
||||
yield
|
||||
|
||||
@pytest.fixture(autouse=True)
|
||||
def reset_dyno_state():
|
||||
cache_limit = dynamo.config.cache_size_limit
|
||||
try:
|
||||
dynamo.config.cache_size_limit = 512
|
||||
dynamo.reset()
|
||||
torch.cuda.synchronize()
|
||||
gc.collect()
|
||||
torch.cuda.empty_cache()
|
||||
yield {}
|
||||
except RuntimeError as err:
|
||||
raise err
|
||||
finally:
|
||||
dynamo.config.cache_size_limit = cache_limit
|
||||
torch.cuda.synchronize()
|
||||
gc.collect()
|
||||
torch.cuda.empty_cache()
|
||||
|
||||
|
||||
def assert_all_close(a: torch.Tensor, b: torch.Tensor, rtol=0, atol=1e-1) -> None:
|
||||
"""
|
||||
Check that all elements of tensors a and b are within provided thresholds.
|
||||
"""
|
||||
assert a.shape == b.shape, f"Shapes don't match: {a.shape} != {b.shape}"
|
||||
assert a.dtype == b.dtype, f"Dtypes don't match: {a.dtype} != {b.dtype}"
|
||||
assert a.device == b.device, f"Devices don't match: {a.device} != {b.device}"
|
||||
max_abs_diff = torch.max(torch.abs(a - b))
|
||||
rel_diff = torch.abs(a / b)
|
||||
max_rel_diff = torch.max(rel_diff)
|
||||
mismatch_elements = torch.sum(torch.abs(a - b) > atol + rtol * torch.abs(b))
|
||||
nb_elements = torch.numel(a)
|
||||
msg = (
|
||||
f"Differences: "
|
||||
f"{max_abs_diff:.3f} (max abs), "
|
||||
f"{max_rel_diff:.3f} (max rel), "
|
||||
f"{mismatch_elements}/{nb_elements} (mismatch elements)"
|
||||
)
|
||||
assert torch.allclose(a, b, rtol=rtol, atol=atol), msg
|
||||
0
tests/kernels/__init__.py
Normal file
0
tests/kernels/__init__.py
Normal file
0
tests/kernels/test_rmsprop.py
Normal file
0
tests/kernels/test_rmsprop.py
Normal file
Loading…
Add table
Add a link
Reference in a new issue