Merge nucleic/sleek-ember-seal-uady into dev
This commit is contained in:
@@ -15,6 +15,38 @@ import train
|
||||
import train_mlx
|
||||
|
||||
|
||||
class MLXDeviceTests(unittest.TestCase):
|
||||
class Metal:
|
||||
def __init__(self, available):
|
||||
self.available = available
|
||||
|
||||
def is_available(self):
|
||||
return self.available
|
||||
|
||||
class MLX:
|
||||
cpu = "cpu"
|
||||
gpu = "gpu"
|
||||
|
||||
def __init__(self, metal_available):
|
||||
self.metal = MLXDeviceTests.Metal(metal_available)
|
||||
self.selected = None
|
||||
|
||||
def set_default_device(self, device):
|
||||
self.selected = device
|
||||
|
||||
def test_cpu_is_an_explicit_fallback(self):
|
||||
mlx = self.MLX(metal_available=False)
|
||||
train_mlx._configure_mlx_device(mlx, "cpu")
|
||||
self.assertEqual("cpu", mlx.selected)
|
||||
|
||||
def test_metal_fails_closed_when_unavailable(self):
|
||||
with self.assertRaisesRegex(train.DataError, "requires Apple Silicon"):
|
||||
train_mlx._configure_mlx_device(
|
||||
self.MLX(metal_available=False),
|
||||
"metal",
|
||||
)
|
||||
|
||||
|
||||
class FixedShapeTokenizerTests(unittest.TestCase):
|
||||
class Tokenizer:
|
||||
pad_token_id = 0
|
||||
|
||||
Reference in New Issue
Block a user