fix: normalize CUDA device requests #31
This file contains hidden or bidirectional Unicode text that may be interpreted or compiled differently than what appears below. To review, open the file in an editor that reveals hidden Unicode characters.
Learn more about bidirectional Unicode characters
| name: CI | |
| on: | |
| push: | |
| branches: [main, develop] | |
| pull_request: | |
| branches: [main] | |
| env: | |
| PYTHON_VERSION: "3.12" | |
| UV_NO_SOURCES: "1" | |
| jobs: | |
| lint: | |
| name: Lint | |
| runs-on: ubuntu-latest | |
| steps: | |
| - uses: actions/checkout@v6.0.3 | |
| - name: Install uv | |
| uses: astral-sh/setup-uv@v8.2.0 | |
| with: | |
| version: "latest" | |
| - name: Set up Python | |
| run: uv python install ${{ env.PYTHON_VERSION }} | |
| - name: Install dependencies | |
| run: uv sync --dev | |
| - name: Run ruff check | |
| run: uv run ruff check src/ tests/ | |
| - name: Run ruff format check | |
| run: uv run ruff format --check src/ tests/ | |
| typecheck: | |
| name: Type Check | |
| runs-on: ubuntu-latest | |
| steps: | |
| - uses: actions/checkout@v6.0.3 | |
| - name: Install uv | |
| uses: astral-sh/setup-uv@v8.2.0 | |
| with: | |
| version: "latest" | |
| - name: Set up Python | |
| run: uv python install ${{ env.PYTHON_VERSION }} | |
| - name: Install dependencies | |
| run: uv sync --dev | |
| - name: Run ty check | |
| run: uv run ty check | |
| test: | |
| name: Test (Python ${{ matrix.python-version }}) | |
| runs-on: ubuntu-latest | |
| strategy: | |
| fail-fast: false | |
| matrix: | |
| python-version: ["3.12", "3.13", "3.14"] | |
| steps: | |
| - uses: actions/checkout@v6.0.3 | |
| - name: Install uv | |
| uses: astral-sh/setup-uv@v8.2.0 | |
| with: | |
| version: "latest" | |
| - name: Set up Python ${{ matrix.python-version }} | |
| run: uv python install ${{ matrix.python-version }} | |
| - name: Install dependencies | |
| run: uv sync --dev | |
| - name: Run tests | |
| run: uv run pytest tests/ -v --tb=short | |
| test-package-smoke: | |
| name: Test Package Smoke | |
| runs-on: ubuntu-latest | |
| needs: [lint, typecheck] | |
| steps: | |
| - uses: actions/checkout@v6.0.3 | |
| - name: Install uv | |
| uses: astral-sh/setup-uv@v8.2.0 | |
| with: | |
| version: "latest" | |
| - name: Set up Python | |
| run: uv python install ${{ env.PYTHON_VERSION }} | |
| - name: Build wheel | |
| run: uv build | |
| - name: Create isolated smoke environment | |
| run: uv venv .smoke-venv --python ${{ env.PYTHON_VERSION }} | |
| - name: Install built wheel | |
| run: uv pip install --python .smoke-venv/bin/python dist/*.whl | |
| - name: Verify import surface | |
| run: | | |
| .smoke-venv/bin/python -c "import ml4t.models as m; print(m.__version__)" | |
| .smoke-venv/bin/python -c "from ml4t.models import PCAModel, SAEModel, StochasticDiscountFactorModel" | |
| mps-smoke: | |
| name: MPS Smoke | |
| runs-on: macos-15 | |
| timeout-minutes: 10 | |
| steps: | |
| - uses: actions/checkout@v6.0.3 | |
| - name: Install uv | |
| uses: astral-sh/setup-uv@v8.2.0 | |
| with: | |
| version: "latest" | |
| - name: Set up Python | |
| run: uv python install ${{ env.PYTHON_VERSION }} | |
| - name: Install dependencies | |
| run: uv sync --dev | |
| - name: Verify PyTorch MPS | |
| run: | | |
| uv run python - <<'PY' | |
| import platform | |
| import torch | |
| from ml4t.models._internal.torch_runtime import resolve_device, seed_torch | |
| print(platform.platform()) | |
| print("machine:", platform.machine()) | |
| print("torch:", torch.__version__) | |
| print("mps built:", torch.backends.mps.is_built()) | |
| print("mps available:", torch.backends.mps.is_available()) | |
| assert platform.machine() == "arm64" | |
| assert torch.backends.mps.is_built() | |
| assert torch.backends.mps.is_available() | |
| device = resolve_device(torch, "mps") | |
| assert device.type == "mps" | |
| seed_torch(torch, 42, device) | |
| x = torch.ones(4, device=device) | |
| y = x * 2 | |
| assert y.device.type == "mps" | |
| assert y.cpu().tolist() == [2.0, 2.0, 2.0, 2.0] | |
| PY | |
| build: | |
| name: Build Package | |
| runs-on: ubuntu-latest | |
| needs: [lint, typecheck, test, test-package-smoke, mps-smoke] | |
| steps: | |
| - uses: actions/checkout@v6.0.3 | |
| with: | |
| fetch-depth: 0 | |
| - name: Install uv | |
| uses: astral-sh/setup-uv@v8.2.0 | |
| with: | |
| version: "latest" | |
| - name: Set up Python | |
| run: uv python install ${{ env.PYTHON_VERSION }} | |
| - name: Build package | |
| run: uv build | |
| - name: Upload build artifacts | |
| uses: actions/upload-artifact@v7.0.1 | |
| with: | |
| name: dist | |
| path: dist/ |