| """ |
| Derived from Andrej Karpathy's nanochat project. |
| |
| MIT License |
| |
| Copyright (c) 2025 Andrej Karpathy |
| |
| Permission is hereby granted, free of charge, to any person obtaining a copy |
| of this software and associated documentation files (the "Software"), to deal |
| in the Software without restriction, including without limitation the rights |
| to use, copy, modify, merge, publish, distribute, sublicense, and/or sell |
| copies of the Software, and to permit persons to whom the Software is |
| furnished to do so, subject to the following conditions: |
| |
| The above copyright notice and this permission notice shall be included in all |
| copies or substantial portions of the Software. |
| """ |
|
|
| from __future__ import annotations |
|
|
| import os |
|
|
| import torch |
|
|
|
|
| def assert_mps_only() -> torch.device: |
| if os.environ.get("PYTORCH_ENABLE_MPS_FALLBACK") == "1": |
| raise SystemExit( |
| "Refusing to run: PYTORCH_ENABLE_MPS_FALLBACK=1 could execute CPU fallbacks." |
| ) |
| if not torch.backends.mps.is_built(): |
| raise SystemExit("PyTorch was not built with MPS. Stopping.") |
| if not torch.backends.mps.is_available(): |
| raise SystemExit("MPS is not available. Stopping before any Torch experiment work.") |
| device = torch.device("mps") |
| torch.set_default_device(device) |
| return device |
|
|