name: GPU tests

on:
  schedule:
    # Runs at midnight every day
    - cron:  '0 0 * * *'
  push:
    branches: [ main ]
  pull_request:
  workflow_dispatch:

concurrency:
  group: gpu-test-${{ github.workflow }}-${{ github.ref == 'refs/heads/main' && github.run_number || github.ref }}
  cancel-in-progress: true

permissions:
  id-token: write
  contents: read

defaults:
  run:
    shell: bash -l -eo pipefail {0}

jobs:
  gpu_test:
    if: github.repository_owner == 'pytorch'
    runs-on: linux.g5.12xlarge.nvidia.gpu
    strategy:
      matrix:
        python-version: ['3.9', '3.10', '3.11']
        torch-version: ["stable", "nightly"]
        # Do not run against nightlies on PR
        exclude:
          - torch-version: ${{ github.event_name == 'pull_request' && 'nightly' }}
    steps:
      - name: Check out repo
        uses: actions/checkout@v4
      - name: Setup conda env
        uses: conda-incubator/setup-miniconda@v2
        with:
          auto-update-conda: true
          miniconda-version: "latest"
          activate-environment: test
          python-version: ${{ matrix.python-version }}
      - name: Update pip
        run: python -m pip install --upgrade pip
      - name: Install nightly versions of PyTorch packages (if applicable)
        if: ${{ matrix.torch-version == 'nightly' }}
        run: python -m pip install --pre torch torchvision torchao --index-url https://download.pytorch.org/whl/nightly/cu126
      - name: Install torch stable (if applicable)
        if: ${{ matrix.torch-version == 'stable' }}
        run: python -m pip install torch torchvision torchao
      - name: Install recipe-specific dependencies
        run: python -m pip install lm-eval==0.4.8
      - name: Install the torchtune library with dev options
        run: python -m pip install -e ".[dev]"
      - name: Run recipe and unit tests with coverage
        run: pytest tests --ignore tests/torchtune/modules/_export --with-integration --cov=. --cov-report=xml --durations=20 -vv
      - name: Upload coverage to Codecov
        uses: codecov/codecov-action@v3
