name: Validate a TestPyPI Release Candidate

on:
  workflow_call:
  workflow_dispatch:

permissions:
  contents: read

jobs:
  validate-rc:
    runs-on: ubuntu-latest
    strategy:
      fail-fast: false
      matrix:
        python-version: ["3.11", "3.12"]

    steps:
      - uses: actions/checkout@v7

      - uses: actions/setup-python@v7
        with:
          python-version: ${{ matrix.python-version }}

      - name: Read RC version
        run: |
          version=$(tr -d '\n' < assets/version.txt)
          case "$version" in
            *rc*) echo "RC_VERSION=$version" >> "$GITHUB_ENV" ;;
            *) echo "Expected an RC version, got: $version"; exit 1 ;;
          esac

      - name: Create clean virtual environment
        run: python -m venv "$RUNNER_TEMP/validate-rc-venv"

      - name: Install release PyTorch packages
        run: |
          "$RUNNER_TEMP/validate-rc-venv/bin/python" -m pip install --upgrade pip
          "$RUNNER_TEMP/validate-rc-venv/bin/python" -m pip install \
            --pre \
            --index-url https://download.pytorch.org/whl/test/cpu \
            "torch==2.14.0" \
            "torchvision==0.29.0" \
            "torchao==0.18.0" \
            "triton==3.8.0"

      - name: Install RC from TestPyPI
        run: |
          for attempt in 1 2 3 4 5; do
            if "$RUNNER_TEMP/validate-rc-venv/bin/python" -m pip install \
              --index-url https://test.pypi.org/simple/ \
              --extra-index-url https://pypi.org/simple/ \
              "torchtitan==$RC_VERSION"; then
              exit 0
            fi
            if [ "$attempt" -eq 5 ]; then
              exit 1
            fi
            sleep 15
          done

      - name: Install CPU unit test dependencies
        run: |
          "$RUNNER_TEMP/validate-rc-venv/bin/python" -m pip install \
            -r requirements-dev.txt

      - name: Verify installed package
        working-directory: ${{ runner.temp }}
        run: |
          "$RUNNER_TEMP/validate-rc-venv/bin/python" - <<'PY'
          import os
          import torchtitan
          import torch
          import torchao
          import torchvision
          import triton

          assert torchtitan.__version__ == os.environ["RC_VERSION"]
          assert "site-packages" in torchtitan.__file__
          print(f"torchtitan={torchtitan.__version__} ({torchtitan.__file__})")
          print(f"torch={torch.__version__}")
          print(f"torchvision={torchvision.__version__}")
          print(f"torchao={torchao.__version__}")
          print(f"triton={triton.__version__}")
          PY

      - name: Run CPU debug-model training smoke test
        working-directory: ${{ runner.temp }}
        run: |
          "$RUNNER_TEMP/validate-rc-venv/bin/python" - <<'PY'
          import torch
          import torch.nn.functional as F
          from torchtitan.models.common.attention import ScaledDotProductInnerAttention
          from torchtitan.models.llama3 import build_model_config

          torch.manual_seed(42)
          model_config = build_model_config("debugmodel")
          for layer_config in model_config.layers:
              layer_config.attention.inner_attention = ScaledDotProductInnerAttention.Config()

          model = model_config.build()
          model.init_states(buffer_device=torch.device("cpu"))
          model.train()
          optimizer = torch.optim.AdamW(model.parameters(), lr=1e-3)
          tokens = torch.randint(0, 2048, (1, 16))
          positions = torch.arange(16).unsqueeze(0)
          losses = []

          for _ in range(3):
              optimizer.zero_grad(set_to_none=True)
              logits = model(tokens, positions, None)
              loss = F.cross_entropy(
                  logits[:, :-1].reshape(-1, logits.size(-1)),
                  tokens[:, 1:].reshape(-1),
              )
              assert torch.isfinite(loss), loss
              loss.backward()
              optimizer.step()
              losses.append(loss.item())

          assert losses[-1] < losses[0], losses
          print(f"losses={losses}")
          PY

      - name: Run CPU unit tests against installed RC
        run: |
          test_root="$RUNNER_TEMP/validate-rc-tests"
          mkdir -p "$test_root"
          cp -R tests "$test_root/tests"
          cp -R scripts "$test_root/scripts"
          cp pyproject.toml "$test_root/pyproject.toml"
          cd "$test_root"
          unset PYTHONPATH
          "$RUNNER_TEMP/validate-rc-venv/bin/python" -m pytest \
            tests/unit_tests/cpu \
            --durations=20 \
            -vv
