diff --git a/.github/workflows/release-desktop.yml b/.github/workflows/release-desktop.yml index 2d1c6f51f2..ea82739968 100644 --- a/.github/workflows/release-desktop.yml +++ b/.github/workflows/release-desktop.yml @@ -54,10 +54,43 @@ jobs: with: node-version: 24 + - name: Install pinned Tauri CLI + run: npm install --save-dev --prefix studio @tauri-apps/cli@2.10.1 + + - name: Verify pinned Tauri CLI + shell: bash + run: | + out="$(npx --prefix studio tauri --version)" + echo "$out" + if [ "$out" != "tauri-cli 2.10.1" ]; then + echo "Expected tauri-cli 2.10.1, got $out" >&2 + exit 1 + fi + - name: Install frontend dependencies working-directory: studio/frontend run: npm install + - name: Verify backend package is published + shell: bash + run: | + node <<'JS' + const { readFileSync } = require('node:fs'); + + (async () => { + const cargo = readFileSync('studio/src-tauri/Cargo.toml', 'utf8'); + const match = cargo.match(/^version\s*=\s*"([^"]+)"/m); + if (!match) throw new Error('Could not read desktop app version'); + + const appVersion = match[1]; + const response = await fetch(`https://pypi.org/pypi/unsloth/${appVersion}/json`); + if (!response.ok) { + const message = 'Publish unsloth=={app_version} to PyPI before the desktop release'; + throw new Error(`${message.replace('{app_version}', appVersion)} (HTTP ${response.status})`); + } + })(); + JS + # ── Rust ── - name: Install Rust stable uses: dtolnay/rust-toolchain@stable @@ -105,13 +138,14 @@ jobs: # ── Linux: build + sign + upload ── - name: Build Linux app if: matrix.platform == 'ubuntu-22.04' - uses: tauri-apps/tauri-action@fce9c6108b31ea247710505d3aaaa893ee6768d4 + uses: tauri-apps/tauri-action@84b9d35b5fc46c1e45415bdb6144030364f7ebc5 env: GITHUB_TOKEN: ${{ secrets.GITHUB_TOKEN }} TAURI_SIGNING_PRIVATE_KEY: ${{ secrets.TAURI_SIGNING_PRIVATE_KEY }} TAURI_SIGNING_PRIVATE_KEY_PASSWORD: ${{ secrets.TAURI_SIGNING_PRIVATE_KEY_PASSWORD }} with: projectPath: studio + tauriScript: npx --prefix . tauri tagName: desktop-v__VERSION__ releaseName: 'Unsloth Studio (Desktop) v__VERSION__' releaseBody: | @@ -123,6 +157,7 @@ jobs: > Linux in-app updates are AppImage-oriented. Package installs should update by downloading a new package. > Linux AppImage on Ubuntu 24.04+ may require: `sudo apt install libfuse2t64` + > First-run system dependency elevation is supported on Ubuntu/Debian. Other Linux distributions should install system packages manually. releaseDraft: ${{ inputs.draft }} prerelease: false args: -v ${{ matrix.args }} @@ -130,7 +165,7 @@ jobs: # ── macOS: build + sign + notarize + upload ── - name: Build macOS app if: matrix.platform == 'macos-latest' - uses: tauri-apps/tauri-action@fce9c6108b31ea247710505d3aaaa893ee6768d4 + uses: tauri-apps/tauri-action@84b9d35b5fc46c1e45415bdb6144030364f7ebc5 env: GITHUB_TOKEN: ${{ secrets.GITHUB_TOKEN }} TAURI_SIGNING_PRIVATE_KEY: ${{ secrets.TAURI_SIGNING_PRIVATE_KEY }} @@ -141,6 +176,7 @@ jobs: APPLE_TEAM_ID: ${{ secrets.APPLE_TEAM_ID }} with: projectPath: studio + tauriScript: npx --prefix . tauri tagName: desktop-v__VERSION__ releaseName: 'Unsloth Studio (Desktop) v__VERSION__' releaseBody: | @@ -152,6 +188,7 @@ jobs: > Linux in-app updates are AppImage-oriented. Package installs should update by downloading a new package. > Linux AppImage on Ubuntu 24.04+ may require: `sudo apt install libfuse2t64` + > First-run system dependency elevation is supported on Ubuntu/Debian. Other Linux distributions should install system packages manually. releaseDraft: ${{ inputs.draft }} prerelease: false args: -v ${{ matrix.args }} @@ -159,7 +196,7 @@ jobs: # ── Windows: build + sign + upload ── - name: Build Windows app if: matrix.platform == 'windows-latest' - uses: tauri-apps/tauri-action@fce9c6108b31ea247710505d3aaaa893ee6768d4 + uses: tauri-apps/tauri-action@84b9d35b5fc46c1e45415bdb6144030364f7ebc5 env: GITHUB_TOKEN: ${{ secrets.GITHUB_TOKEN }} TAURI_SIGNING_PRIVATE_KEY: ${{ secrets.TAURI_SIGNING_PRIVATE_KEY }} @@ -171,6 +208,7 @@ jobs: AZURE_CERTIFICATE_PROFILE_NAME: ${{ secrets.AZURE_CERTIFICATE_PROFILE_NAME }} with: projectPath: studio + tauriScript: npx --prefix . tauri tagName: desktop-v__VERSION__ releaseName: 'Unsloth Studio (Desktop) v__VERSION__' releaseBody: | @@ -182,6 +220,7 @@ jobs: > Linux in-app updates are AppImage-oriented. Package installs should update by downloading a new package. > Linux AppImage on Ubuntu 24.04+ may require: `sudo apt install libfuse2t64` + > First-run system dependency elevation is supported on Ubuntu/Debian. Other Linux distributions should install system packages manually. releaseDraft: ${{ inputs.draft }} prerelease: false args: -v ${{ matrix.args }} diff --git a/.pre-commit-config.yaml b/.pre-commit-config.yaml index 1d3137b373..a2a4995d62 100644 --- a/.pre-commit-config.yaml +++ b/.pre-commit-config.yaml @@ -1,6 +1,6 @@ repos: - repo: https://github.com/astral-sh/ruff-pre-commit - rev: v0.15.11 + rev: v0.15.12 hooks: - id: ruff args: diff --git a/README.md b/README.md index e1b7e448ef..6c12b28096 100644 --- a/README.md +++ b/README.md @@ -79,8 +79,9 @@ irm https://unsloth.ai/install.ps1 | iex #### Launch ```bash -unsloth studio -H 0.0.0.0 -p 8888 +unsloth studio -p 8888 ``` +> For cloud VMs or LAN access, add `-H 0.0.0.0` to bind on all interfaces. #### Update To update, use the same install commands as above. Or run (does not work on Windows): @@ -167,7 +168,7 @@ The below advanced instructions are for Unsloth Studio. For Unsloth Core advance git clone https://github.com/unslothai/unsloth cd unsloth ./install.sh --local -unsloth studio -H 0.0.0.0 -p 8888 +unsloth studio -p 8888 ``` Then to update : ```bash @@ -180,7 +181,7 @@ git clone https://github.com/unslothai/unsloth.git cd unsloth Set-ExecutionPolicy -Scope Process -ExecutionPolicy Bypass .\install.ps1 --local -unsloth studio -H 0.0.0.0 -p 8888 +unsloth studio -p 8888 ``` Then to update : ```bash @@ -193,11 +194,11 @@ git clone https://github.com/unslothai/unsloth cd unsloth git checkout nightly ./install.sh --local -unsloth studio -H 0.0.0.0 -p 8888 +unsloth studio -p 8888 ``` Then to launch every time: ```bash -unsloth studio -H 0.0.0.0 -p 8888 +unsloth studio -p 8888 ``` #### Nightly: Windows: @@ -208,11 +209,11 @@ cd unsloth git checkout nightly Set-ExecutionPolicy -Scope Process -ExecutionPolicy Bypass .\install.ps1 --local -unsloth studio -H 0.0.0.0 -p 8888 +unsloth studio -p 8888 ``` Then to launch every time: ```bash -unsloth studio -H 0.0.0.0 -p 8888 +unsloth studio -p 8888 ``` #### Uninstall diff --git a/install.ps1 b/install.ps1 index 44464101f3..7dc5a50250 100644 --- a/install.ps1 +++ b/install.ps1 @@ -8,6 +8,79 @@ function Install-UnslothStudio { $ErrorActionPreference = "Stop" $script:UnslothVerbose = ($env:UNSLOTH_VERBOSE -eq "1") + # ── Tauri structured output ── + function Write-TauriLog { + param([string]$Tag, [string]$Message) + if ($TauriMode) { + Write-Host "[TAURI:$Tag] $Message" + } + } + + function Format-TauriDiagBool { + param([bool]$Value) + if ($Value) { return "true" } + return "false" + } + + function Get-TauriDiagArch { + $arch = [string]$env:PROCESSOR_ARCHITECTURE + if ([string]::IsNullOrWhiteSpace($arch)) { + try { $arch = [System.Runtime.InteropServices.RuntimeInformation]::OSArchitecture.ToString() } catch { $arch = "unknown" } + } + $arch = $arch.ToLowerInvariant() + switch ($arch) { + "amd64" { return "x86_64" } + "x64" { return "x86_64" } + "arm64" { return "arm64" } + "x86" { return "x86" } + default { return ($arch -replace '[^a-z0-9_.-]', '_') } + } + } + + function Get-TauriTorchIndexFamily { + param([string]$TorchIndexUrl) + if ($SkipTorch) { return "none" } + if ([string]::IsNullOrWhiteSpace($TorchIndexUrl)) { return "none" } + $leaf = ($TorchIndexUrl.TrimEnd('/') -split '/')[-1].ToLowerInvariant() + if (@("cpu", "cu118", "cu124", "cu126", "cu128", "cu130") -contains $leaf) { return $leaf } + if ($leaf -match '^rocm[0-9]+\.[0-9]+$') { return $leaf } + return "auto" + } + + function Get-TauriGpuBranch { + param([string]$TorchIndexFamily) + if ($SkipTorch) { return "no_torch" } + if ($TorchIndexFamily -like "cu*") { return "cuda" } + if ($TorchIndexFamily -like "rocm*") { return "rocm" } + if ($TorchIndexFamily -eq "cpu") { return "cpu" } + return "unknown" + } + + function Write-TauriDiag { + param( + [string]$GpuBranch = "unknown", + [string]$TorchIndexFamily = "none", + [string]$PythonVersionForDiag = $PythonVersion + ) + if ([string]::IsNullOrWhiteSpace($PythonVersionForDiag)) { $PythonVersionForDiag = "unknown" } + Write-TauriLog "DIAG" "diag_schema=1 platform=windows arch=$(Get-TauriDiagArch) python_version=$($PythonVersionForDiag.ToLowerInvariant()) skip_torch=$(Format-TauriDiagBool $SkipTorch) mac_intel=false gpu_branch=$GpuBranch torch_index_family=$TorchIndexFamily" + } + + function Exit-InstallFailure { + param( + [Parameter(Mandatory = $true)][string]$Message, + [int]$Code = 1 + ) + if ($Code -eq 0) { $Code = 1 } + Write-TauriLog "ERROR" $Message + if (Get-Command Restore-StudioVenvRollback -CommandType Function -ErrorAction SilentlyContinue) { + Restore-StudioVenvRollback + } + if ($TauriMode) { + exit $Code + } + } + # ── Parse flags ── $StudioLocalInstall = $false $PackageName = "unsloth" @@ -26,7 +99,7 @@ function Install-UnslothStudio { $i++ if ($i -ge $argList.Count) { Write-Host "[ERROR] --package requires an argument." -ForegroundColor Red - return + return (Exit-InstallFailure "--package requires an argument.") } $PackageName = $argList[$i] } @@ -42,22 +115,14 @@ function Install-UnslothStudio { $RepoRoot = (Resolve-Path (Split-Path -Parent $PSCommandPath)).Path if (-not (Test-Path (Join-Path $RepoRoot "pyproject.toml"))) { Write-Host "[ERROR] --local must be run from the unsloth repo root (pyproject.toml not found at $RepoRoot)" -ForegroundColor Red - return + return (Exit-InstallFailure "--local must be run from the unsloth repo root") } } # Validate --package to prevent injection into shell/Python commands if ($PackageName -notmatch '^[a-zA-Z0-9][a-zA-Z0-9._-]*$') { Write-Host "[ERROR] --package name contains invalid characters (allowed: a-z A-Z 0-9 . _ -)" -ForegroundColor Red - return - } - - # ── Tauri structured output ── - function Write-TauriLog { - param([string]$Tag, [string]$Message) - if ($TauriMode) { - Write-Host "[TAURI:$Tag] $Message" - } + return (Exit-InstallFailure "--package name contains invalid characters") } $PythonVersion = "3.13" @@ -487,7 +552,7 @@ try { } catch {} exit 1 } - `$studioCommand = '& "' + `$studioExe + '" studio -H 0.0.0.0 -p ' + `$launchPort + `$studioCommand = '& "' + `$studioExe + '" studio -p ' + `$launchPort `$launchArgs = @( '-NoExit', '-NoProfile', @@ -630,7 +695,7 @@ shell.Run cmd, 0, False step "winget" "not available" "Red" substep "Install it from https://aka.ms/getwinget" "Yellow" substep "or install Python $PythonVersion and uv manually, then re-run." "Yellow" - return + return (Exit-InstallFailure "winget is not available") } # ── Helper: detect a working Python 3.11-3.13 on the system ── @@ -749,9 +814,14 @@ shell.Run cmd, 0, False Write-Host " Please install Python $PythonVersion manually from https://www.python.org/downloads/" -ForegroundColor Yellow Write-Host " Make sure to check 'Add Python to PATH' during installation." -ForegroundColor Yellow Write-Host " Then re-run this installer." -ForegroundColor Yellow - return + return (Exit-InstallFailure "Python installation failed") } } + $DiagPythonVersion = $PythonVersion + if ($DetectedPython) { $DiagPythonVersion = $DetectedPython.Version } + $InitialGpuBranch = "unknown" + if ($SkipTorch) { $InitialGpuBranch = "no_torch" } + Write-TauriDiag -GpuBranch $InitialGpuBranch -TorchIndexFamily "none" -PythonVersionForDiag $DiagPythonVersion # ── Install uv if not present ── Write-TauriLog "STEP" "Installing uv package manager" @@ -773,7 +843,7 @@ shell.Run cmd, 0, False if (-not (Get-Command uv -ErrorAction SilentlyContinue)) { step "uv" "could not be installed" "Red" substep "Install it from https://docs.astral.sh/uv/" "Yellow" - return + return (Exit-InstallFailure "uv could not be installed") } # ── Create venv (migrate old layout if possible, otherwise fresh) ── @@ -786,11 +856,68 @@ shell.Run cmd, 0, False $VenvPython = Join-Path $VenvDir "Scripts\python.exe" $_Migrated = $false + $script:StudioVenvRollbackDir = $null + $script:StudioVenvRollbackTarget = $VenvDir + $script:StudioVenvRollbackActive = $false + + function Start-StudioVenvRollback { + param([Parameter(Mandatory = $true)][string]$ExistingDir) + $stamp = Get-Date -Format "yyyyMMddHHmmss" + $candidate = Join-Path $StudioHome "unsloth_studio.rollback.$stamp.$PID" + $suffix = 0 + while (Test-Path $candidate) { + $suffix++ + $candidate = Join-Path $StudioHome "unsloth_studio.rollback.$stamp.$PID.$suffix" + } + Move-Item -Path $ExistingDir -Destination $candidate -ErrorAction Stop + $script:StudioVenvRollbackDir = $candidate + $script:StudioVenvRollbackTarget = $ExistingDir + $script:StudioVenvRollbackActive = $true + substep "previous environment preserved for rollback" + } + + function Restore-StudioVenvRollback { + if (-not $script:StudioVenvRollbackActive) { return } + $backup = $script:StudioVenvRollbackDir + $target = $script:StudioVenvRollbackTarget + if (-not $backup -or -not (Test-Path $backup)) { + $script:StudioVenvRollbackActive = $false + return + } + substep "restoring previous environment after failed install..." "Yellow" + try { + if (Test-Path $target) { + Remove-Item -Recurse -Force $target -ErrorAction SilentlyContinue + } + Move-Item -Path $backup -Destination $target -Force -ErrorAction Stop + substep "restored previous environment" + $script:StudioVenvRollbackActive = $false + $script:StudioVenvRollbackDir = $null + } catch { + Write-Host "[WARN] Could not restore previous environment from $backup to $target" -ForegroundColor Yellow + Write-Host " $($_.Exception.Message)" -ForegroundColor Yellow + } + } + + function Complete-StudioVenvRollback { + if (-not $script:StudioVenvRollbackActive) { return } + $backup = $script:StudioVenvRollbackDir + if ($backup -and (Test-Path $backup)) { + Remove-Item -Recurse -Force $backup -ErrorAction SilentlyContinue + } + $script:StudioVenvRollbackActive = $false + $script:StudioVenvRollbackDir = $null + } if (Test-Path $VenvPython) { - # New layout already exists -- nuke for fresh install - substep "removing existing environment for fresh install..." - Remove-Item -Recurse -Force $VenvDir + # New layout already exists -- replace only after preserving rollback copy. + substep "preserving existing environment for rollback..." + try { + Start-StudioVenvRollback -ExistingDir $VenvDir + } catch { + Write-Host "[ERROR] Could not prepare existing environment for reinstall: $($_.Exception.Message)" -ForegroundColor Red + return (Exit-InstallFailure "Could not prepare existing environment for reinstall") + } } elseif (Test-Path (Join-Path $StudioHome ".venv\Scripts\python.exe")) { # Old layout (~/.unsloth/studio/.venv) exists -- validate before migrating $OldVenv = Join-Path $StudioHome ".venv" @@ -799,18 +926,23 @@ shell.Run cmd, 0, False $prevEAP2 = $ErrorActionPreference $ErrorActionPreference = "Continue" try { - & $OldPy -c "import torch; A = torch.ones((2,2)); B = A + A" 2>$null | Out-Null - $torchOk = ($LASTEXITCODE -eq 0) - } catch { $torchOk = $false } + if ($SkipTorch) { + & $OldPy -c "import sys; print(sys.executable)" 2>$null | Out-Null + } else { + & $OldPy -c "import torch; A = torch.ones((2,2)); B = A + A" 2>$null | Out-Null + } + $legacyOk = ($LASTEXITCODE -eq 0) + } catch { $legacyOk = $false } $ErrorActionPreference = $prevEAP2 - if ($torchOk) { + if ($legacyOk) { substep "legacy environment is healthy -- migrating..." Move-Item -Path $OldVenv -Destination $VenvDir -Force substep "moved .venv -> unsloth_studio" $_Migrated = $true } else { substep "legacy environment failed validation -- creating fresh environment" "Yellow" - Remove-Item -Recurse -Force $OldVenv -ErrorAction SilentlyContinue + $invalidVenv = Join-Path $StudioHome (".venv.invalid.{0}.{1}" -f (Get-Date -Format "yyyyMMddHHmmss"), $PID) + Move-Item -Path $OldVenv -Destination $invalidVenv -Force -ErrorAction SilentlyContinue } } elseif (Test-Path (Join-Path $env:USERPROFILE "unsloth_studio\Scripts\python.exe")) { # CWD-relative venv from old install.ps1 -- migrate to absolute path @@ -826,9 +958,8 @@ shell.Run cmd, 0, False substep "$VenvDir" $venvExit = Invoke-InstallCommand { uv venv $VenvDir --python "$($DetectedPython.Path)" } if ($venvExit -ne 0) { - Write-TauriLog "ERROR" "Failed to create virtual environment (exit code $venvExit)" Write-Host "[ERROR] Failed to create virtual environment (exit code $venvExit)" -ForegroundColor Red - return + return (Exit-InstallFailure "Failed to create virtual environment (exit code $venvExit)" $venvExit) } } else { step "venv" "using migrated environment" @@ -886,6 +1017,9 @@ shell.Run cmd, 0, False return "$baseUrl/cu126" } $TorchIndexUrl = Get-TorchIndexUrl + $TorchIndexFamily = Get-TauriTorchIndexFamily $TorchIndexUrl + $GpuBranch = Get-TauriGpuBranch $TorchIndexFamily + Write-TauriDiag -GpuBranch $GpuBranch -TorchIndexFamily $TorchIndexFamily -PythonVersionForDiag $DetectedPython.Version # ── Print CPU-only hint when no GPU detected ── if (-not $SkipTorch -and $TorchIndexUrl -like "*/cpu") { @@ -934,7 +1068,7 @@ shell.Run cmd, 0, False if ($SkipTorch) { # No-torch: install unsloth + unsloth-zoo with --no-deps, then # runtime deps (typer, safetensors, transformers, etc.) with --no-deps. - $baseInstallExit = Invoke-InstallCommand { uv pip install --python $VenvPython --no-deps --reinstall-package unsloth --reinstall-package unsloth-zoo "unsloth>=2026.4.7" unsloth-zoo } + $baseInstallExit = Invoke-InstallCommand { uv pip install --python $VenvPython --no-deps --reinstall-package unsloth --reinstall-package unsloth-zoo "unsloth>=2026.4.8" unsloth-zoo } if ($baseInstallExit -eq 0) { $NoTorchReq = Find-NoTorchRuntimeFile if ($NoTorchReq) { @@ -942,18 +1076,24 @@ shell.Run cmd, 0, False } } } else { - $baseInstallExit = Invoke-InstallCommand { uv pip install --python $VenvPython --reinstall-package unsloth --reinstall-package unsloth-zoo "unsloth>=2026.4.7" unsloth-zoo } + $baseInstallExit = Invoke-InstallCommand { uv pip install --python $VenvPython --reinstall-package unsloth --reinstall-package unsloth-zoo "unsloth>=2026.4.8" unsloth-zoo } } if ($baseInstallExit -ne 0) { Write-Host "[ERROR] Failed to install unsloth (exit code $baseInstallExit)" -ForegroundColor Red - return + return (Exit-InstallFailure "Failed to install unsloth (exit code $baseInstallExit)" $baseInstallExit) } if ($StudioLocalInstall) { substep "overlaying local repo (editable)..." $overlayExit = Invoke-InstallCommand { uv pip install --python $VenvPython -e $RepoRoot --no-deps } if ($overlayExit -ne 0) { Write-Host "[ERROR] Failed to overlay local repo (exit code $overlayExit)" -ForegroundColor Red - return + return (Exit-InstallFailure "Failed to overlay local repo (exit code $overlayExit)" $overlayExit) + } + substep "overlaying unsloth-zoo from git main..." + $zooOverlayExit = Invoke-InstallCommand { uv pip install --python $VenvPython --no-deps --reinstall-package unsloth-zoo "unsloth-zoo @ git+https://github.com/unslothai/unsloth-zoo" } + if ($zooOverlayExit -ne 0) { + Write-Host "[ERROR] Failed to overlay unsloth-zoo (exit code $zooOverlayExit)" -ForegroundColor Red + return (Exit-InstallFailure "Failed to overlay unsloth-zoo (exit code $zooOverlayExit)" $zooOverlayExit) } } } elseif ($TorchIndexUrl) { @@ -964,9 +1104,8 @@ shell.Run cmd, 0, False substep "installing PyTorch ($TorchIndexUrl)..." $torchInstallExit = Invoke-InstallCommand { uv pip install --python $VenvPython "torch>=2.4,<2.11.0" torchvision torchaudio --index-url $TorchIndexUrl } if ($torchInstallExit -ne 0) { - Write-TauriLog "ERROR" "Failed to install PyTorch (exit code $torchInstallExit)" Write-Host "[ERROR] Failed to install PyTorch (exit code $torchInstallExit)" -ForegroundColor Red - return + return (Exit-InstallFailure "Failed to install PyTorch (exit code $torchInstallExit)" $torchInstallExit) } } @@ -975,7 +1114,7 @@ shell.Run cmd, 0, False if ($SkipTorch) { # No-torch: install unsloth + unsloth-zoo with --no-deps, then # runtime deps (typer, safetensors, transformers, etc.) with --no-deps. - $baseInstallExit = Invoke-InstallCommand { uv pip install --python $VenvPython --no-deps --upgrade-package unsloth --upgrade-package unsloth-zoo "unsloth>=2026.4.7" unsloth-zoo } + $baseInstallExit = Invoke-InstallCommand { uv pip install --python $VenvPython --no-deps --upgrade-package unsloth --upgrade-package unsloth-zoo "unsloth>=2026.4.8" unsloth-zoo } if ($baseInstallExit -eq 0) { $NoTorchReq = Find-NoTorchRuntimeFile if ($NoTorchReq) { @@ -983,14 +1122,13 @@ shell.Run cmd, 0, False } } } elseif ($StudioLocalInstall) { - $baseInstallExit = Invoke-InstallCommand { uv pip install --python $VenvPython --upgrade-package unsloth "unsloth>=2026.4.7" unsloth-zoo } + $baseInstallExit = Invoke-InstallCommand { uv pip install --python $VenvPython --upgrade-package unsloth "unsloth>=2026.4.8" unsloth-zoo } } else { $baseInstallExit = Invoke-InstallCommand { uv pip install --python $VenvPython --upgrade-package unsloth -- "$PackageName" } } if ($baseInstallExit -ne 0) { - Write-TauriLog "ERROR" "Failed to install unsloth (exit code $baseInstallExit)" Write-Host "[ERROR] Failed to install unsloth (exit code $baseInstallExit)" -ForegroundColor Red - return + return (Exit-InstallFailure "Failed to install unsloth (exit code $baseInstallExit)" $baseInstallExit) } if ($StudioLocalInstall) { @@ -998,7 +1136,13 @@ shell.Run cmd, 0, False $overlayExit = Invoke-InstallCommand { uv pip install --python $VenvPython -e $RepoRoot --no-deps } if ($overlayExit -ne 0) { Write-Host "[ERROR] Failed to overlay local repo (exit code $overlayExit)" -ForegroundColor Red - return + return (Exit-InstallFailure "Failed to overlay local repo (exit code $overlayExit)" $overlayExit) + } + substep "overlaying unsloth-zoo from git main..." + $zooOverlayExit = Invoke-InstallCommand { uv pip install --python $VenvPython --no-deps --reinstall-package unsloth-zoo "unsloth-zoo @ git+https://github.com/unslothai/unsloth-zoo" } + if ($zooOverlayExit -ne 0) { + Write-Host "[ERROR] Failed to overlay unsloth-zoo (exit code $zooOverlayExit)" -ForegroundColor Red + return (Exit-InstallFailure "Failed to overlay unsloth-zoo (exit code $zooOverlayExit)" $zooOverlayExit) } } } else { @@ -1006,53 +1150,72 @@ shell.Run cmd, 0, False Write-TauriLog "STEP" "Installing unsloth" substep "installing unsloth (this may take a few minutes)..." if ($StudioLocalInstall) { - $baseInstallExit = Invoke-InstallCommand { uv pip install --python $VenvPython unsloth-zoo "unsloth>=2026.4.7" --torch-backend=auto } + $baseInstallExit = Invoke-InstallCommand { uv pip install --python $VenvPython unsloth-zoo "unsloth>=2026.4.8" --torch-backend=auto } if ($baseInstallExit -ne 0) { Write-Host "[ERROR] Failed to install unsloth (exit code $baseInstallExit)" -ForegroundColor Red - return + return (Exit-InstallFailure "Failed to install unsloth (exit code $baseInstallExit)" $baseInstallExit) } substep "overlaying local repo (editable)..." $overlayExit = Invoke-InstallCommand { uv pip install --python $VenvPython -e $RepoRoot --no-deps } if ($overlayExit -ne 0) { Write-Host "[ERROR] Failed to overlay local repo (exit code $overlayExit)" -ForegroundColor Red - return + return (Exit-InstallFailure "Failed to overlay local repo (exit code $overlayExit)" $overlayExit) + } + substep "overlaying unsloth-zoo from git main..." + $zooOverlayExit = Invoke-InstallCommand { uv pip install --python $VenvPython --no-deps --reinstall-package unsloth-zoo "unsloth-zoo @ git+https://github.com/unslothai/unsloth-zoo" } + if ($zooOverlayExit -ne 0) { + Write-Host "[ERROR] Failed to overlay unsloth-zoo (exit code $zooOverlayExit)" -ForegroundColor Red + return (Exit-InstallFailure "Failed to overlay unsloth-zoo (exit code $zooOverlayExit)" $zooOverlayExit) } } else { $baseInstallExit = Invoke-InstallCommand { uv pip install --python $VenvPython --torch-backend=auto -- "$PackageName" } if ($baseInstallExit -ne 0) { - Write-TauriLog "ERROR" "Failed to install unsloth (exit code $baseInstallExit)" Write-Host "[ERROR] Failed to install unsloth (exit code $baseInstallExit)" -ForegroundColor Red - return + return (Exit-InstallFailure "Failed to install unsloth (exit code $baseInstallExit)" $baseInstallExit) } } } - # Hotfix: patch install_python_stack.py for Windows GUI stdout - # The PyPI version crashes with OSError when stdout is piped from a GUI app. - # Copy our fixed version (bundled by Tauri) over the installed one. - # Remove this block once PyPI ships the fix from commit 18c5aae7. + # Overlay Tauri-bundled studio fixes that may be ahead of PyPI. Skipped + # for --local: the editable install above already makes _PACKAGE_ROOT in + # unsloth_cli/commands/studio.py resolve to the repo (PEP 660 __file__). + # Source paths match the Tauri bundle layout in studio/src-tauri/tauri.conf.json, + # which bundles install_python_stack.py at the bundle root next to install.ps1. if ($TauriMode) { $rawPath = if ($PSCommandPath) { $PSCommandPath } else { $MyInvocation.ScriptName } - $scriptDir = Split-Path -Parent ($rawPath -replace '^\\\\\?\\', '') - $fixedPy = Join-Path $scriptDir "install_python_stack.py" - $target = Join-Path $VenvDir "Lib\site-packages\studio\install_python_stack.py" - $sentinel = "# UNSLOTH_DESKTOP_HOTFIX_APPLIED_v1" - $sentinelPattern = [regex]::Escape($sentinel) - if ((Test-Path $fixedPy) -and (Test-Path $target)) { - $installed = Get-Content $target -Raw - if ($installed -notmatch $sentinelPattern) { - Copy-Item $fixedPy $target -Force - Add-Content -Path $target -Value "`n$sentinel" - substep "patched install_python_stack.py (stdout fix)" - } else { - substep "install_python_stack.py already has stdout fix" + if ($rawPath) { + # Strip leading \\?\ extended-length prefix if the launcher passed one. + $scriptDir = Split-Path -Parent ($rawPath -replace '^\\\\\?\\', '') + $overlayMap = [ordered]@{ + "install_python_stack.py" = "Lib\site-packages\studio\install_python_stack.py" + } + foreach ($rel in $overlayMap.Keys) { + $src = Join-Path $scriptDir $rel + $dst = Join-Path $VenvDir $overlayMap[$rel] + if (-not (Test-Path $src)) { continue } + $dstParent = Split-Path -Parent $dst + if (-not (Test-Path $dstParent)) { + Write-Host "[WARN] Overlay target dir missing: $dstParent; studio setup may use stale bundled file" -ForegroundColor Yellow + continue + } + try { + if (-not (Test-Path $dst)) { + # Backfill: target file missing but parent dir exists. + Copy-Item $src $dst -Force + substep ("backfilled bundled " + (Split-Path -Leaf $rel)) + } else { + # Hash-compare so re-runs are no-ops when files already match. + $srcHash = (Get-FileHash $src -Algorithm SHA256).Hash + $dstHash = (Get-FileHash $dst -Algorithm SHA256).Hash + if ($srcHash -ne $dstHash) { + Copy-Item $src $dst -Force + substep ("applied bundled " + (Split-Path -Leaf $rel)) + } + } + } catch { + Write-Host "[WARN] Could not overlay $($rel): $($_.Exception.Message); studio setup may use stale bundled file" -ForegroundColor Yellow + } } - } elseif ((Test-Path $fixedPy) -and (Test-Path (Split-Path $target))) { - Copy-Item $fixedPy $target -Force - Add-Content -Path $target -Value "`n$sentinel" - substep "patched install_python_stack.py (stdout fix)" - } else { - Write-Host "[WARN] Could not patch install_python_stack.py (bundled file or target dir missing)" -ForegroundColor Yellow } } @@ -1063,12 +1226,11 @@ shell.Run cmd, 0, False step "setup" "running unsloth studio setup..." $UnslothExe = Join-Path $VenvDir "Scripts\unsloth.exe" if (-not (Test-Path $UnslothExe)) { - Write-TauriLog "ERROR" "unsloth CLI was not installed correctly" Write-Host "[ERROR] unsloth CLI was not installed correctly." -ForegroundColor Red Write-Host " Expected: $UnslothExe" -ForegroundColor Yellow Write-Host " This usually means an older unsloth version was installed that does not include the Studio CLI." -ForegroundColor Yellow Write-Host " Try re-running the installer or see: https://github.com/unslothai/unsloth?tab=readme-ov-file#-quickstart" -ForegroundColor Yellow - return + return (Exit-InstallFailure "unsloth CLI was not installed correctly") } # Tell setup.ps1 to skip base package installation (install.ps1 already did it) $env:SKIP_STUDIO_BASE = "1" @@ -1090,12 +1252,16 @@ shell.Run cmd, 0, False # and bypass the fast-path version check from PR #4667. $studioArgs = @('studio', 'setup') if ($script:UnslothVerbose) { $studioArgs += '--verbose' } - & $UnslothExe @studioArgs - $setupExit = $LASTEXITCODE + $env:UNSLOTH_INSTALL_ROLLBACK_MANAGED = "1" + try { + & $UnslothExe @studioArgs + $setupExit = $LASTEXITCODE + } finally { + Remove-Item Env:UNSLOTH_INSTALL_ROLLBACK_MANAGED -ErrorAction SilentlyContinue + } if ($setupExit -ne 0) { - Write-TauriLog "ERROR" "unsloth studio setup failed (exit code $setupExit)" Write-Host "[ERROR] unsloth studio setup failed (exit code $setupExit)" -ForegroundColor Red - return + return (Exit-InstallFailure "unsloth studio setup failed (exit code $setupExit)" $setupExit) } # ── Expose `unsloth` via a shim dir containing only unsloth.exe ── @@ -1168,6 +1334,7 @@ shell.Run cmd, 0, False step "path" "added unsloth launcher to PATH" } Refresh-SessionPath # sync current session with registry + Complete-StudioVenvRollback # ── Tauri mode: done, skip shortcuts and auto-launch ── if ($TauriMode) { @@ -1177,15 +1344,25 @@ shell.Run cmd, 0, False New-StudioShortcuts -UnslothExePath $UnslothExe - # Launch studio automatically in interactive terminals; - # in non-interactive environments (CI, Docker) just print instructions. + # In interactive terminals, ask the user before starting Studio. + # In non-interactive environments (CI, Docker) just print instructions. $IsInteractive = [Environment]::UserInteractive -and (-not [Console]::IsInputRedirected) if ($IsInteractive) { - & $UnslothExe studio -H 0.0.0.0 -p 8888 + Write-Host "" + $reply = Read-Host " Start Unsloth Studio now? [Y/n]" + if ([string]::IsNullOrWhiteSpace($reply) -or $reply -match '^[Yy]') { + & $UnslothExe studio -p 8888 + } else { + step "launch" "to start later, run:" + substep "unsloth studio -p 8888" + substep "(add -H 0.0.0.0 to allow network / cloud access)" + Write-Host "" + } } else { step "launch" "manual commands:" substep "& `"$VenvDir\Scripts\Activate.ps1`"" - substep "unsloth studio -H 0.0.0.0 -p 8888" + substep "unsloth studio -p 8888" + substep "(add -H 0.0.0.0 to allow network / cloud access)" Write-Host "" } } diff --git a/install.sh b/install.sh index 07473e441d..7948170043 100755 --- a/install.sh +++ b/install.sh @@ -162,9 +162,119 @@ tauri_log() { fi } +tauri_diag_marker() { + _diag_gpu_branch="${1:-unknown}" + _diag_torch_index_family="${2:-none}" + tauri_log "DIAG" "diag_schema=1 platform=${OS:-unknown} arch=${_ARCH:-unknown} python_version=${PYTHON_VERSION:-unknown} skip_torch=${SKIP_TORCH:-false} mac_intel=${MAC_INTEL:-false} gpu_branch=${_diag_gpu_branch} torch_index_family=${_diag_torch_index_family}" +} + +_tauri_torch_index_family() { + if [ "${SKIP_TORCH:-false}" = true ]; then + echo "none" + return + fi + _diag_url="${1:-}" + case "$_diag_url" in + */cu118) echo "cu118" ;; + */cu124) echo "cu124" ;; + */cu126) echo "cu126" ;; + */cu128) echo "cu128" ;; + */cu130) echo "cu130" ;; + */cpu) echo "cpu" ;; + */rocm[0-9]*.[0-9]*) + _diag_family=${_diag_url##*/} + case "$_diag_family" in + rocm[0-9]*.[0-9]*) echo "$_diag_family" ;; + *) echo "auto" ;; + esac ;; + "") echo "none" ;; + *) echo "auto" ;; + esac +} + +_tauri_gpu_branch() { + _diag_family="${1:-unknown}" + _diag_radeon="${2:-false}" + if [ "${SKIP_TORCH:-false}" = true ]; then + echo "no_torch" + return + fi + if [ "${OS:-}" = "macos" ]; then + echo "mac" + return + fi + case "$_diag_family" in + cu*) echo "cuda" ;; + rocm*) + if [ "$_diag_radeon" = true ]; then + echo "rocm_radeon" + else + echo "rocm" + fi ;; + radeon) echo "rocm_radeon" ;; + cpu) echo "cpu" ;; + none) echo "no_torch" ;; + *) echo "unknown" ;; + esac +} + PYTHON_VERSION="" # resolved after platform detection STUDIO_HOME="$HOME/.unsloth/studio" VENV_DIR="$STUDIO_HOME/unsloth_studio" +_VENV_ROLLBACK_DIR="" +_VENV_ROLLBACK_TARGET="$VENV_DIR" +_VENV_ROLLBACK_ACTIVE=false + +_start_studio_venv_replacement() { + _existing_dir="$1" + _stamp=$(date +%Y%m%d%H%M%S 2>/dev/null || echo "time") + _candidate="$STUDIO_HOME/unsloth_studio.rollback.$_stamp.$$" + _suffix=0 + while [ -e "$_candidate" ]; do + _suffix=$((_suffix + 1)) + _candidate="$STUDIO_HOME/unsloth_studio.rollback.$_stamp.$$.$_suffix" + done + mv "$_existing_dir" "$_candidate" + _VENV_ROLLBACK_DIR="$_candidate" + _VENV_ROLLBACK_TARGET="$_existing_dir" + _VENV_ROLLBACK_ACTIVE=true + substep "previous environment preserved for rollback" +} + +_restore_studio_venv_replacement() { + [ "$_VENV_ROLLBACK_ACTIVE" = true ] || return 0 + [ -n "$_VENV_ROLLBACK_DIR" ] && [ -d "$_VENV_ROLLBACK_DIR" ] || { + _VENV_ROLLBACK_ACTIVE=false + return 0 + } + substep "restoring previous environment after failed install..." "$C_WARN" + rm -rf "$_VENV_ROLLBACK_TARGET" + if mv "$_VENV_ROLLBACK_DIR" "$_VENV_ROLLBACK_TARGET"; then + substep "restored previous environment" + _VENV_ROLLBACK_ACTIVE=false + _VENV_ROLLBACK_DIR="" + else + echo "⚠️ Could not restore previous environment from $_VENV_ROLLBACK_DIR to $_VENV_ROLLBACK_TARGET" >&2 + fi +} + +_commit_studio_venv_replacement() { + [ "$_VENV_ROLLBACK_ACTIVE" = true ] || return 0 + if [ -n "$_VENV_ROLLBACK_DIR" ] && [ -d "$_VENV_ROLLBACK_DIR" ]; then + rm -rf "$_VENV_ROLLBACK_DIR" || true + fi + _VENV_ROLLBACK_ACTIVE=false + _VENV_ROLLBACK_DIR="" +} + +_on_install_exit() { + _status=$? + if [ "$_status" -ne 0 ]; then + _restore_studio_venv_replacement + fi + exit "$_status" +} +trap _on_install_exit EXIT # ── Helper: download a URL to a file (supports curl and wget) ── download() { @@ -512,11 +622,11 @@ if [ -t 1 ]; then ) & # Clear traps so exec does not trigger _release_lock (the subshell owns it) trap - EXIT INT TERM - exec "$UNSLOTH_EXE" studio -H 0.0.0.0 -p "$_launch_port" + exec "$UNSLOTH_EXE" studio -p "$_launch_port" else # ── Background mode (no TTY) ── # Used by macOS .app and headless invocations. - _launch_cmd=$(printf '%q ' "$UNSLOTH_EXE" studio -H 0.0.0.0 -p "$_launch_port") + _launch_cmd=$(printf '%q ' "$UNSLOTH_EXE" studio -p "$_launch_port") _launch_cmd=${_launch_cmd% } _spawn_terminal "$_launch_cmd" @@ -828,6 +938,14 @@ if [ "$_NO_TORCH_FLAG" = true ] || [ "$MAC_INTEL" = true ]; then SKIP_TORCH=true fi +_TAURI_INITIAL_GPU_BRANCH="unknown" +if [ "$SKIP_TORCH" = true ]; then + _TAURI_INITIAL_GPU_BRANCH="no_torch" +elif [ "$OS" = "macos" ]; then + _TAURI_INITIAL_GPU_BRANCH="mac" +fi +tauri_diag_marker "$_TAURI_INITIAL_GPU_BRANCH" "none" + # ── Check system dependencies ── # cmake and git are needed by unsloth studio setup to build the GGUF inference # engine (llama.cpp). build-essential and libcurl-dev are also needed on Linux. @@ -856,9 +974,7 @@ case "$OS" in fi command -v gcc >/dev/null 2>&1 || MISSING="$MISSING build-essential" # libcurl dev headers for llama.cpp HTTPS support - if command -v dpkg >/dev/null 2>&1; then - dpkg -s libcurl4-openssl-dev >/dev/null 2>&1 || MISSING="$MISSING libcurl4-openssl-dev" - fi + command -v curl-config >/dev/null 2>&1 || MISSING="$MISSING libcurl4-openssl-dev" ;; esac @@ -883,9 +999,15 @@ if [ -n "$MISSING" ]; then if command -v apt-get >/dev/null 2>&1; then _smart_apt_install $MISSING else - echo " apt-get is not available. Please install with your package manager:" + echo " Automatic system package installation is supported on apt-based" + echo " Linux distributions (Ubuntu/Debian) only. Please install the" + echo " missing dependencies with your package manager, then re-run setup:" echo " $MISSING" - echo " Then re-run Unsloth Studio setup." + echo "" + echo " Examples:" + echo " Fedora/RHEL: sudo dnf install cmake git gcc gcc-c++ make libcurl-devel" + echo " Arch: sudo pacman -S --needed cmake git base-devel curl" + echo " openSUSE: sudo zypper install cmake git gcc gcc-c++ make libcurl-devel" exit 1 fi ;; @@ -957,12 +1079,19 @@ mkdir -p "$STUDIO_HOME" _MIGRATED=false if [ -x "$VENV_DIR/bin/python" ]; then - # New layout already exists — nuke for fresh install - rm -rf "$VENV_DIR" + # New layout already exists — replace only after preserving rollback copy. + substep "preserving existing environment for rollback..." + _start_studio_venv_replacement "$VENV_DIR" elif [ -x "$STUDIO_HOME/.venv/bin/python" ]; then - # Old layout exists — validate before migrating + # Old layout exists — validate before migrating. + # In no-torch mode, a missing torch package is expected; validate Python only. substep "found legacy Studio environment, validating..." - if "$STUDIO_HOME/.venv/bin/python" -c " + _legacy_ok=false + if [ "$SKIP_TORCH" = true ]; then + if "$STUDIO_HOME/.venv/bin/python" -c "import sys; print(sys.executable)" >/dev/null 2>&1; then + _legacy_ok=true + fi + elif "$STUDIO_HOME/.venv/bin/python" -c " import torch device = 'cuda' if torch.cuda.is_available() else 'cpu' A = torch.ones((10, 10), device=device) @@ -972,13 +1101,17 @@ D = A + B E = D @ C torch.testing.assert_close(torch.unique(E), torch.tensor((20,), device=E.device, dtype=E.dtype)) " >/dev/null 2>&1; then + _legacy_ok=true + fi + if [ "$_legacy_ok" = true ]; then echo "✅ Legacy environment is healthy — migrating..." mv "$STUDIO_HOME/.venv" "$VENV_DIR" echo " Moved ~/.unsloth/studio/.venv → $VENV_DIR" _MIGRATED=true else echo "⚠️ Legacy environment failed validation — creating fresh environment" - rm -rf "$STUDIO_HOME/.venv" + _invalid_venv="$STUDIO_HOME/.venv.invalid.$(date +%Y%m%d%H%M%S 2>/dev/null || echo time).$$" + mv "$STUDIO_HOME/.venv" "$_invalid_venv" 2>/dev/null || true fi fi @@ -1308,6 +1441,12 @@ case "$TORCH_INDEX_URL" in fi ;; esac +_TAURI_TORCH_INDEX_FAMILY=$(_tauri_torch_index_family "$TORCH_INDEX_URL") +if [ "$_amd_gpu_radeon" = true ] && [ "$SKIP_TORCH" = false ]; then + _TAURI_TORCH_INDEX_FAMILY="radeon" +fi +_TAURI_GPU_BRANCH=$(_tauri_gpu_branch "$_TAURI_TORCH_INDEX_FAMILY" "$_amd_gpu_radeon") +tauri_diag_marker "$_TAURI_GPU_BRANCH" "$_TAURI_TORCH_INDEX_FAMILY" # ── Print CPU-only hint when no GPU detected ── case "$TORCH_INDEX_URL" in @@ -1347,7 +1486,7 @@ if [ "$_MIGRATED" = true ]; then # to prevent transitive torch resolution. run_install_cmd "install unsloth (migrated no-torch)" uv pip install --python "$_VENV_PY" --no-deps \ --reinstall-package unsloth --reinstall-package unsloth-zoo \ - "unsloth>=2026.4.7" unsloth-zoo + "unsloth>=2026.4.8" unsloth-zoo _NO_TORCH_RT="$(_find_no_torch_runtime)" if [ -n "$_NO_TORCH_RT" ]; then run_install_cmd "install no-torch runtime deps" uv pip install --python "$_VENV_PY" --no-deps -r "$_NO_TORCH_RT" @@ -1355,11 +1494,15 @@ if [ "$_MIGRATED" = true ]; then else run_install_cmd "install unsloth (migrated)" uv pip install --python "$_VENV_PY" \ --reinstall-package unsloth --reinstall-package unsloth-zoo \ - "unsloth>=2026.4.7" unsloth-zoo + "unsloth>=2026.4.8" unsloth-zoo fi if [ "$STUDIO_LOCAL_INSTALL" = true ]; then substep "overlaying local repo (editable)..." run_install_cmd "overlay local repo" uv pip install --python "$_VENV_PY" -e "$_REPO_ROOT" --no-deps + substep "overlaying unsloth-zoo from git main..." + run_install_cmd "overlay unsloth-zoo (git main)" uv pip install --python "$_VENV_PY" \ + --no-deps --reinstall-package unsloth-zoo \ + "unsloth-zoo @ git+https://github.com/unslothai/unsloth-zoo" fi # AMD ROCm: install bitsandbytes even in migrated environments so # existing ROCm installs gain the AMD bitsandbytes build without a @@ -1519,7 +1662,7 @@ elif [ -n "$TORCH_INDEX_URL" ]; then # runtime deps (typer, safetensors, transformers, etc.) with --no-deps. run_install_cmd "install unsloth (no-torch)" uv pip install --python "$_VENV_PY" --no-deps \ --upgrade-package unsloth --upgrade-package unsloth-zoo \ - "unsloth>=2026.4.7" unsloth-zoo + "unsloth>=2026.4.8" unsloth-zoo _NO_TORCH_RT="$(_find_no_torch_runtime)" if [ -n "$_NO_TORCH_RT" ]; then run_install_cmd "install no-torch runtime deps" uv pip install --python "$_VENV_PY" --no-deps -r "$_NO_TORCH_RT" @@ -1527,12 +1670,20 @@ elif [ -n "$TORCH_INDEX_URL" ]; then if [ "$STUDIO_LOCAL_INSTALL" = true ]; then substep "overlaying local repo (editable)..." run_install_cmd "overlay local repo" uv pip install --python "$_VENV_PY" -e "$_REPO_ROOT" --no-deps + substep "overlaying unsloth-zoo from git main..." + run_install_cmd "overlay unsloth-zoo (git main)" uv pip install --python "$_VENV_PY" \ + --no-deps --reinstall-package unsloth-zoo \ + "unsloth-zoo @ git+https://github.com/unslothai/unsloth-zoo" fi elif [ "$STUDIO_LOCAL_INSTALL" = true ]; then run_install_cmd "install unsloth (local)" uv pip install --python "$_VENV_PY" \ - --upgrade-package unsloth "unsloth>=2026.4.7" unsloth-zoo + --upgrade-package unsloth "unsloth>=2026.4.8" unsloth-zoo substep "overlaying local repo (editable)..." run_install_cmd "overlay local repo" uv pip install --python "$_VENV_PY" -e "$_REPO_ROOT" --no-deps + substep "overlaying unsloth-zoo from git main..." + run_install_cmd "overlay unsloth-zoo (git main)" uv pip install --python "$_VENV_PY" \ + --no-deps --reinstall-package unsloth-zoo \ + "unsloth-zoo @ git+https://github.com/unslothai/unsloth-zoo" else run_install_cmd "install unsloth" uv pip install --python "$_VENV_PY" \ --upgrade-package unsloth -- "$PACKAGE_NAME" @@ -1558,9 +1709,13 @@ else tauri_log "STEP" "Installing Unsloth" substep "installing unsloth (this may take a few minutes)..." if [ "$STUDIO_LOCAL_INSTALL" = true ]; then - run_install_cmd "install unsloth (auto torch backend)" uv pip install --python "$_VENV_PY" unsloth-zoo "unsloth>=2026.4.7" --torch-backend=auto + run_install_cmd "install unsloth (auto torch backend)" uv pip install --python "$_VENV_PY" unsloth-zoo "unsloth>=2026.4.8" --torch-backend=auto substep "overlaying local repo (editable)..." run_install_cmd "overlay local repo" uv pip install --python "$_VENV_PY" -e "$_REPO_ROOT" --no-deps + substep "overlaying unsloth-zoo from git main..." + run_install_cmd "overlay unsloth-zoo (git main)" uv pip install --python "$_VENV_PY" \ + --no-deps --reinstall-package unsloth-zoo \ + "unsloth-zoo @ git+https://github.com/unslothai/unsloth-zoo" else run_install_cmd "install unsloth (auto torch backend)" uv pip install --python "$_VENV_PY" --torch-backend=auto -- "$PACKAGE_NAME" fi @@ -1679,6 +1834,8 @@ if [ "$_SETUP_EXIT" -ne 0 ]; then exit "$_SETUP_EXIT" fi +_commit_studio_venv_replacement + # ── Tauri mode: done, skip shortcuts and auto-launch ── if [ "$TAURI_MODE" = true ]; then tauri_log "DONE" "" @@ -1690,28 +1847,46 @@ printf " ${C_TITLE}%s${C_RST}\n" "Unsloth Studio installed!" printf " ${C_DIM}%s${C_RST}\n" "$RULE" echo "" -# Launch studio automatically in interactive terminals; -# in non-interactive environments (Docker, CI, cloud-init) just print instructions. +# In interactive terminals, ask the user before starting Studio. +# In non-interactive environments (Docker, CI, cloud-init) just print instructions. if [ -t 1 ]; then - step "launch" "starting Unsloth Studio..." - "$VENV_DIR/bin/unsloth" studio -H 0.0.0.0 -p 8888 - _LAUNCH_EXIT=$? - if [ "$_LAUNCH_EXIT" -ne 0 ] && [ "$_MIGRATED" = true ]; then - echo "" - echo "⚠️ Unsloth Studio failed to start after migration." - echo " Your migrated environment may be incompatible." - echo " To fix, remove the environment and reinstall:" - echo "" - echo " rm -rf $VENV_DIR" - echo " curl -fsSL https://unsloth.ai/install.sh | sh" - echo "" + echo "" + printf " Start Unsloth Studio now? [Y/n] " + if [ -r /dev/tty ]; then + read -r _reply sqlite3.Connection: created_at TEXT NOT NULL, last_used_at TEXT, expires_at TEXT, - is_active INTEGER NOT NULL DEFAULT 1 + is_active INTEGER NOT NULL DEFAULT 1, + is_internal INTEGER NOT NULL DEFAULT 0 ); """ ) + api_key_columns = { + row["name"] for row in conn.execute("PRAGMA table_info(api_keys)") + } + if "is_internal" not in api_key_columns: + conn.execute( + "ALTER TABLE api_keys ADD COLUMN is_internal INTEGER NOT NULL DEFAULT 0" + ) conn.execute( """ CREATE TABLE IF NOT EXISTS app_secrets ( @@ -592,11 +600,15 @@ def create_api_key( username: str, name: str, expires_at: Optional[str] = None, + internal: bool = False, ) -> Tuple[str, dict]: """Create a new API key for *username*. Returns ``(raw_key, row_dict)`` where *raw_key* is shown to the user - exactly once. The database only stores the SHA-256 hash. + exactly once. The database only stores the PBKDF2 hash. + + Pass ``internal=True`` for keys minted by workflows (e.g. data-recipe + runs) that should not appear in user-facing key listings. """ raw_key = API_KEY_PREFIX + secrets.token_hex(16) key_hash = _pbkdf2_api_key(raw_key) @@ -607,10 +619,18 @@ def create_api_key( try: conn.execute( """ - INSERT INTO api_keys (username, key_prefix, key_hash, name, created_at, expires_at) - VALUES (?, ?, ?, ?, ?, ?) + INSERT INTO api_keys (username, key_prefix, key_hash, name, created_at, expires_at, is_internal) + VALUES (?, ?, ?, ?, ?, ?, ?) """, - (username, key_prefix, key_hash, name, now, expires_at), + ( + username, + key_prefix, + key_hash, + name, + now, + expires_at, + 1 if internal else 0, + ), ) conn.commit() cur = conn.execute("SELECT * FROM api_keys WHERE key_hash = ?", (key_hash,)) @@ -620,19 +640,33 @@ def create_api_key( conn.close() -def list_api_keys(username: str) -> list: - """Return all API keys for *username* (never exposes ``key_hash``).""" +def list_api_keys(username: str, include_internal: bool = False) -> list: + """Return API keys for *username*. Internal workflow keys are hidden + by default so they do not clutter user-facing UIs.""" conn = get_connection() try: - cur = conn.execute( - """ - SELECT id, username, key_prefix, name, created_at, last_used_at, expires_at, is_active - FROM api_keys - WHERE username = ? - ORDER BY created_at DESC - """, - (username,), - ) + if include_internal: + cur = conn.execute( + """ + SELECT id, username, key_prefix, name, created_at, last_used_at, + expires_at, is_active, is_internal + FROM api_keys + WHERE username = ? + ORDER BY created_at DESC + """, + (username,), + ) + else: + cur = conn.execute( + """ + SELECT id, username, key_prefix, name, created_at, last_used_at, + expires_at, is_active, is_internal + FROM api_keys + WHERE username = ? AND is_internal = 0 + ORDER BY created_at DESC + """, + (username,), + ) return [dict(row) for row in cur.fetchall()] finally: conn.close() @@ -652,6 +686,24 @@ def revoke_api_key(username: str, key_id: int) -> bool: conn.close() +def revoke_internal_api_key(key_id: int) -> bool: + """Revoke an internal workflow-minted key without requiring a username. + + Used by the recipe runner to retire its sk-unsloth-* key once the job + terminates, shrinking the window a leaked key could be abused. + """ + conn = get_connection() + try: + cursor = conn.execute( + "UPDATE api_keys SET is_active = 0 WHERE id = ? AND is_internal = 1", + (key_id,), + ) + conn.commit() + return cursor.rowcount > 0 + finally: + conn.close() + + def validate_api_key(raw_key: str) -> Optional[str]: """Validate *raw_key* and return the owning username, or ``None``. diff --git a/studio/backend/core/data_recipe/jobs/constants.py b/studio/backend/core/data_recipe/jobs/constants.py index 08237326f8..0045276e20 100644 --- a/studio/backend/core/data_recipe/jobs/constants.py +++ b/studio/backend/core/data_recipe/jobs/constants.py @@ -9,6 +9,7 @@ STAGE_PREVIEW = "preview" STAGE_DAG = "dag" STAGE_HEALTHCHECK = "healthcheck" STAGE_SAMPLING = "sampling" +STAGE_SOURCE = "source" STAGE_COLUMN_CONFIG = "column_config" STAGE_GENERATING = "generating" STAGE_BATCH = "batch" diff --git a/studio/backend/core/data_recipe/jobs/manager.py b/studio/backend/core/data_recipe/jobs/manager.py index 3d7cf2dbe6..cdc28d9560 100644 --- a/studio/backend/core/data_recipe/jobs/manager.py +++ b/studio/backend/core/data_recipe/jobs/manager.py @@ -33,6 +33,60 @@ from .worker import run_job_process _CTX = mp.get_context("spawn") +def _github_source_estimated_total(recipe: dict) -> int | None: + seed_config = recipe.get("seed_config") + if not isinstance(seed_config, dict): + return None + source = seed_config.get("source") + if not isinstance(source, dict) or source.get("seed_type") != "github_repo": + return None + + repos_raw = source.get("repos") + repos = ( + [repo for repo in repos_raw if isinstance(repo, str) and repo.strip()] + if isinstance(repos_raw, list) + else [] + ) + item_types_raw = source.get("item_types") + item_types = ( + [ + item + for item in item_types_raw + if isinstance(item, str) and item in {"issues", "pulls", "commits"} + ] + if isinstance(item_types_raw, list) + else [] + ) + try: + limit = int(source.get("limit") or 0) + except (TypeError, ValueError): + return None + if not repos or not item_types or limit <= 0: + return None + return len(repos) * len(item_types) * limit + + +def _source_progress_status(job: Job) -> dict[str, Any] | None: + progress = job.source_progress + if progress is None: + return None + return { + "source": progress.source, + "status": progress.status, + "repo": progress.repo, + "resource": progress.resource, + "page": progress.page, + "page_items": progress.page_items, + "fetched_items": progress.fetched_items, + "estimated_total": progress.estimated_total, + "percent": progress.percent, + "rate_remaining": progress.rate_remaining, + "retry_after_sec": progress.retry_after_sec, + "message": progress.message, + "updated_at": progress.updated_at, + } + + @dataclass class Subscription: replay: list[dict] @@ -71,8 +125,20 @@ class JobManager: self._pump_thread: threading.Thread | None = None self._seq: int = 0 - def start(self, *, recipe: dict, run: dict) -> str: - """Spawn the job subprocess (one at a time, no cap).""" + def start( + self, + *, + recipe: dict, + run: dict, + internal_api_key_id: int | None = None, + ) -> str: + """Spawn the job subprocess (one at a time, no cap). + + ``internal_api_key_id`` is the row id of a workflow-scoped + sk-unsloth-* key minted by the route layer for local providers. + JobManager revokes it when the job reaches a terminal state so the + key's live window is no longer than the run. + """ llm_columns = recipe.get("columns") or [] llm_column_count = 0 if isinstance(llm_columns, list): @@ -92,18 +158,29 @@ class JobManager: job_id = uuid.uuid4().hex self._job = Job(job_id = job_id, status = "pending", started_at = time.time()) self._job.progress_columns_total = llm_column_count + self._job.source_progress_estimated_total = _github_source_estimated_total( + recipe + ) + self._job.internal_api_key_id = internal_api_key_id self._events.clear() self._seq = 0 run_payload = dict(run) run_payload["_job_id"] = job_id - mp_q = _CTX.Queue() - proc = _CTX.Process( - target = run_job_process, - kwargs = {"event_queue": mp_q, "recipe": recipe, "run": run_payload}, - daemon = True, + from utils.native_path_leases import ( + native_path_secret_removed_for_child_start, + run_without_native_path_secret, ) - proc.start() + + with native_path_secret_removed_for_child_start(): + mp_q = _CTX.Queue() + proc = _CTX.Process( + target = run_without_native_path_secret, + args = (run_job_process,), + kwargs = {"event_queue": mp_q, "recipe": recipe, "run": run_payload}, + daemon = True, + ) + proc.start() self._mp_q = mp_q self._proc = proc @@ -163,6 +240,7 @@ class JobManager: "ok": job.column_progress.ok, "failed": job.column_progress.failed, }, + "source_progress": _source_progress_status(job), "model_usage": { name: { "model": usage.model, @@ -405,6 +483,7 @@ class JobManager: for e in self._drain_queue(mp_q): self._handle_event(job, e) + retired_job: Job | None = None with self._lock: if self._job and self._job.status in { "pending", @@ -429,6 +508,9 @@ class JobManager: "job_id": self._job.job_id, } ) + retired_job = self._job + if retired_job is not None: + self._retire_workflow_key(retired_job) return def _handle_event(self, job: Job, event: dict) -> None: @@ -436,6 +518,7 @@ class JobManager: et = event.get("type") msg = event.get("message") if et == "log" else None + terminal = False with self._lock: if self._job is None or self._job.job_id != job.job_id: return @@ -452,18 +535,43 @@ class JobManager: if self._job.progress.total and self._job.progress.total > 0: self._job.progress.done = self._job.progress.total self._job.progress.percent = 100.0 + terminal = True if et == EVENT_JOB_ERROR: self._job.status = "error" self._job.finished_at = time.time() self._job.error = event.get("error") or "error" + terminal = True + if et == EVENT_JOB_CANCELLED: + terminal = True if msg: upd = parse_log_message(msg) if upd: apply_update(self._job, upd) + if terminal: + self._retire_workflow_key(job) + self._emit(event) + def _retire_workflow_key(self, job: Job) -> None: + """Revoke the workflow-scoped sk-unsloth-* key, if one was minted. + + Best-effort: revocation failures are swallowed. The key would + expire on its own after 24h, so a missed revoke is a latency + concern, not a correctness one. + """ + key_id = getattr(job, "internal_api_key_id", None) + if not key_id: + return + try: + from auth import storage # deferred: avoids circular import + + storage.revoke_internal_api_key(int(key_id)) + except Exception: + pass + job.internal_api_key_id = None + _JOB_MANAGER: JobManager | None = None diff --git a/studio/backend/core/data_recipe/jobs/parse.py b/studio/backend/core/data_recipe/jobs/parse.py index 324b62a92e..cea6d8ea64 100644 --- a/studio/backend/core/data_recipe/jobs/parse.py +++ b/studio/backend/core/data_recipe/jobs/parse.py @@ -4,6 +4,7 @@ from __future__ import annotations import re +import time from dataclasses import dataclass from typing import Any @@ -17,9 +18,10 @@ from .constants import ( STAGE_PREVIEW, STAGE_PROFILING, STAGE_SAMPLING, + STAGE_SOURCE, USAGE_RESET_STAGES, ) -from .types import Job, ModelUsage, Progress +from .types import Job, ModelUsage, Progress, SourceProgress @dataclass(frozen = True) @@ -41,6 +43,7 @@ class ParsedUpdate: usage_requests_total: int | None = None usage_rpm: float | None = None usage_section_start: bool | None = None + source_progress: SourceProgress | None = None # kinda of a bummber but currently only option, Best effort parser from data-designer logs -> structured status for UI. @@ -61,9 +64,165 @@ _RE_USAGE_TOKENS = re.compile( _RE_USAGE_REQUESTS = re.compile( r"requests:\s*success=(?P\d+),\s*failed=(?P\d+),\s*total=(?P\d+),\s*rpm=(?P[0-9.]+)" ) +_RE_GITHUB_PAGE = re.compile( + r"^\[(?P[^\]\s]+/[^\]\s]+)\]\s+" + r"(?Pissues|PRs|commits)\s+page\s+(?P\d+)\s+" + r"\(\+(?P\d+)\).*?\bremaining=(?P\d+)", + re.IGNORECASE, +) +_RE_GITHUB_RATE_LIMIT = re.compile( + r"Rate limit hit\. Sleeping (?P\d+)s until reset\.", + re.IGNORECASE, +) +_RE_GITHUB_SECONDARY_RATE_LIMIT = re.compile( + r"Secondary rate limit(?: on REST)?\. Sleep (?P\d+)s\.", + re.IGNORECASE, +) +_RE_GITHUB_REST_RATE_LIMIT = re.compile( + r"REST 403/429, sleep (?P\d+)", + re.IGNORECASE, +) +_RE_GITHUB_TRANSIENT = re.compile( + r"^(?PGraphQL|REST) (?P\d{3}) transient, retrying", + re.IGNORECASE, +) +_RE_GITHUB_NETWORK_RETRY = re.compile( + r"^(?PGraphQL|REST) network error: .* Retry\.", + re.IGNORECASE, +) +_RE_GITHUB_TRIAL_LIMIT = re.compile( + r"Trial limit reached for (?Pissues|PRs|commits) \((?P\d+)\)", + re.IGNORECASE, +) +_RE_GITHUB_COMPLETE = re.compile( + r"Scraper complete\. GraphQL calls=\d+ REST calls=\d+", + re.IGNORECASE, +) def parse_log_message(msg: str) -> ParsedUpdate | None: + m = _RE_GITHUB_PAGE.search(msg) + if m: + resource_raw = m.group("resource") + resource = "pulls" if resource_raw.lower() == "prs" else resource_raw.lower() + repo = m.group("repo") + page = int(m.group("page")) + page_items = int(m.group("items")) + return ParsedUpdate( + stage = STAGE_SOURCE, + source_progress = SourceProgress( + source = "github", + status = "fetching", + repo = repo, + resource = resource, + page = page, + page_items = page_items, + rate_remaining = int(m.group("remaining")), + message = ( + f"Scraping GitHub source: {repo} " + f"{resource} page {page} (+{page_items})" + ), + ), + ) + + m = _RE_GITHUB_RATE_LIMIT.search(msg) + if m: + seconds = int(m.group("seconds")) + return ParsedUpdate( + stage = STAGE_SOURCE, + source_progress = SourceProgress( + source = "github", + status = "rate_limited", + retry_after_sec = seconds, + message = ( + "Waiting for GitHub rate limit. " + "Studio will resume automatically." + ), + ), + ) + + m = _RE_GITHUB_SECONDARY_RATE_LIMIT.search(msg) + if m: + seconds = int(m.group("seconds")) + return ParsedUpdate( + stage = STAGE_SOURCE, + source_progress = SourceProgress( + source = "github", + status = "rate_limited", + retry_after_sec = seconds, + message = ( + "Waiting for GitHub secondary rate limit. " + "Studio will resume automatically." + ), + ), + ) + + m = _RE_GITHUB_REST_RATE_LIMIT.search(msg) + if m: + seconds = int(m.group("seconds")) + return ParsedUpdate( + stage = STAGE_SOURCE, + source_progress = SourceProgress( + source = "github", + status = "rate_limited", + retry_after_sec = seconds, + message = ( + "Waiting for GitHub rate limit. " + "Studio will resume automatically." + ), + ), + ) + + m = _RE_GITHUB_TRIAL_LIMIT.search(msg) + if m: + resource_raw = m.group("resource") + resource = "pulls" if resource_raw.lower() == "prs" else resource_raw.lower() + items = int(m.group("items")) + return ParsedUpdate( + stage = STAGE_SOURCE, + source_progress = SourceProgress( + source = "github", + status = "fetching", + resource = resource, + message = f"GitHub {resource} trial limit reached ({items}).", + ), + ) + + m = _RE_GITHUB_TRANSIENT.search(msg) + if m: + api = m.group("api") + code = m.group("code") + return ParsedUpdate( + stage = STAGE_SOURCE, + source_progress = SourceProgress( + source = "github", + status = "retrying", + message = f"GitHub {api} returned {code}; retrying automatically.", + ), + ) + + m = _RE_GITHUB_NETWORK_RETRY.search(msg) + if m: + api = m.group("api") + return ParsedUpdate( + stage = STAGE_SOURCE, + source_progress = SourceProgress( + source = "github", + status = "retrying", + message = f"GitHub {api} request failed; retrying automatically.", + ), + ) + + if _RE_GITHUB_COMPLETE.search(msg): + return ParsedUpdate( + stage = STAGE_SOURCE, + source_progress = SourceProgress( + source = "github", + status = "completed", + message = "GitHub source scrape complete.", + ), + ) + m = _RE_SAMPLERS.search(msg) if m: return ParsedUpdate( @@ -172,6 +331,8 @@ def apply_update(job: Job, update: ParsedUpdate) -> None: job.batch.idx = update.batch_idx if update.batch_total is not None: job.batch.total = update.batch_total + if update.source_progress is not None: + _apply_source_progress(job, update.source_progress) if update.stage in USAGE_RESET_STAGES: # usage summary is a short block so we reset once we move into the next stage. @@ -216,6 +377,67 @@ def apply_update(job: Job, update: ParsedUpdate) -> None: usage.rpm = update.usage_rpm +def _apply_source_progress(job: Job, progress: SourceProgress) -> None: + previous = job.source_progress + now = time.time() + + page_items = progress.page_items + if progress.repo and progress.resource and progress.page is not None: + page_key = f"{progress.repo}:{progress.resource}:{progress.page}" + count_key = f"{progress.repo}:{progress.resource}" + if page_key not in job._source_seen_pages: + job._source_seen_pages.add(page_key) + job._source_counts[count_key] = int( + job._source_counts.get(count_key, 0) + ) + int(page_items or 0) + + fetched_items = sum(job._source_counts.values()) + if fetched_items <= 0: + fetched_items = progress.fetched_items or ( + previous.fetched_items if previous else None + ) + + estimated_total = ( + progress.estimated_total + or job.source_progress_estimated_total + or (previous.estimated_total if previous else None) + ) + percent: float | None = progress.percent + if percent is None and estimated_total and fetched_items is not None: + raw_percent = (float(fetched_items) / float(max(1, estimated_total))) * 100.0 + percent = 100.0 if progress.status == "completed" else min(99.0, raw_percent) + if percent is None and previous is not None: + percent = previous.percent + + job.source_progress = SourceProgress( + source = "github", + status = progress.status or (previous.status if previous else None), + repo = progress.repo or (previous.repo if previous else None), + resource = progress.resource or (previous.resource if previous else None), + page = ( + progress.page + if progress.page is not None + else (previous.page if previous else None) + ), + page_items = ( + page_items + if page_items is not None + else (previous.page_items if previous else None) + ), + fetched_items = fetched_items, + estimated_total = estimated_total, + percent = percent, + rate_remaining = ( + progress.rate_remaining + if progress.rate_remaining is not None + else (previous.rate_remaining if previous else None) + ), + retry_after_sec = progress.retry_after_sec, + message = progress.message or (previous.message if previous else None), + updated_at = now, + ) + + def _compute_overall_progress(job: Job, column_progress: Progress) -> Progress: if not job.rows: return column_progress diff --git a/studio/backend/core/data_recipe/jobs/types.py b/studio/backend/core/data_recipe/jobs/types.py index 8d77903238..3d3ddb974e 100644 --- a/studio/backend/core/data_recipe/jobs/types.py +++ b/studio/backend/core/data_recipe/jobs/types.py @@ -35,6 +35,23 @@ class BatchProgress: total: int | None = None +@dataclass +class SourceProgress: + source: str = "github" + status: str | None = None + repo: str | None = None + resource: str | None = None + page: int | None = None + page_items: int | None = None + fetched_items: int | None = None + estimated_total: int | None = None + percent: float | None = None + rate_remaining: int | None = None + retry_after_sec: int | None = None + message: str | None = None + updated_at: float | None = None + + @dataclass class ModelUsage: model: str @@ -57,6 +74,7 @@ class Job: progress: Progress = field(default_factory = Progress) column_progress: Progress = field(default_factory = Progress) batch: BatchProgress = field(default_factory = BatchProgress) + source_progress: SourceProgress | None = None rows: int | None = None cols: int | None = None error: str | None = None @@ -70,8 +88,15 @@ class Job: processor_artifacts: dict[str, Any] | None = None model_usage: dict[str, ModelUsage] = field(default_factory = dict) progress_columns_total: int | None = None + source_progress_estimated_total: int | None = None completed_columns: list[str] = field(default_factory = list) + # Id of the internal sk-unsloth-* API key minted for a local-model + # workflow. Revoked when the job terminates so the key's live window + # matches the run rather than its 24h TTL. + internal_api_key_id: int | None = None _current_usage_model: str | None = None _in_usage_summary: bool = False _seen_generation_columns: list[str] = field(default_factory = list) _column_done: dict[str, int] = field(default_factory = dict) + _source_counts: dict[str, int] = field(default_factory = dict) + _source_seen_pages: set[str] = field(default_factory = set) diff --git a/studio/backend/core/data_recipe/jobs/worker.py b/studio/backend/core/data_recipe/jobs/worker.py index 63e38bd18d..8c5c7fe657 100644 --- a/studio/backend/core/data_recipe/jobs/worker.py +++ b/studio/backend/core/data_recipe/jobs/worker.py @@ -21,6 +21,15 @@ from ..service import build_config_builder, create_data_designer from utils.paths import ensure_dir, recipe_datasets_root _ARTIFACT_ROOT = recipe_datasets_root() +_RE_GITHUB_CURSOR = re.compile(r"\bcursor=[^\s,]+") +_RE_SECRET_TOKEN = re.compile( + r"\b(?:(?:ghp|gho|ghu|ghs|ghr|github_pat)_[A-Za-z0-9_]+|sk-unsloth-[A-Za-z0-9]+)" +) + + +def _sanitize_log_message(message: str) -> str: + message = _RE_GITHUB_CURSOR.sub("cursor=", message) + return _RE_SECRET_TOKEN.sub("", message) class _QueueLogHandler(logging.Handler): @@ -35,7 +44,7 @@ class _QueueLogHandler(logging.Handler): "ts": record.created, "level": record.levelname, "logger": record.name, - "message": record.getMessage(), + "message": _sanitize_log_message(record.getMessage()), } self._q.put(event) except (OSError, RuntimeError, ValueError): @@ -119,10 +128,16 @@ def run_job_process( # Attach queue logger directly to `data_designer` so parser events survive root resets. handler = _QueueLogHandler(event_queue) handler.setLevel(logging.INFO) - data_designer_logger = logging.getLogger("data_designer") - data_designer_logger.addHandler(handler) - data_designer_logger.setLevel(logging.INFO) - data_designer_logger.propagate = True + for logger_name in ( + "data_designer", + "scraper", + "gh_client", + "data_designer_github_repo_seed", + ): + logger = logging.getLogger(logger_name) + logger.addHandler(handler) + logger.setLevel(logging.INFO) + logger.propagate = True if run_config_raw: designer.set_run_config(RunConfig.model_validate(run_config_raw)) @@ -180,8 +195,8 @@ def run_job_process( { "type": EVENT_JOB_ERROR, "ts": time.time(), - "error": str(exc), - "stack": traceback.format_exc(limit = 20), + "error": _sanitize_log_message(str(exc)), + "stack": _sanitize_log_message(traceback.format_exc(limit = 20)), } ) diff --git a/studio/backend/core/data_recipe/local_callable_validators.py b/studio/backend/core/data_recipe/local_callable_validators.py index afd10b02d1..44459e88c5 100644 --- a/studio/backend/core/data_recipe/local_callable_validators.py +++ b/studio/backend/core/data_recipe/local_callable_validators.py @@ -33,6 +33,7 @@ _OXC_TOOL_DIR = Path(__file__).resolve().parent / "oxc-validator" _OXC_RUNNER_PATH = _OXC_TOOL_DIR / "validate.mjs" +from utils.native_path_leases import child_env_without_native_path_secret from utils.subprocess_compat import ( windows_hidden_subprocess_kwargs as _windows_hidden_subprocess_kwargs, ) @@ -248,7 +249,7 @@ def _run_oxc_batch( } try: tmp_dir = ensure_dir(oxc_validator_tmp_root()) - env = dict(os.environ) + env = child_env_without_native_path_secret() tmp_dir_str = str(tmp_dir) env["TMPDIR"] = tmp_dir_str env["TMP"] = tmp_dir_str diff --git a/studio/backend/core/export/export.py b/studio/backend/core/export/export.py index d8f2e8fa37..6fee5a38f7 100644 --- a/studio/backend/core/export/export.py +++ b/studio/backend/core/export/export.py @@ -28,6 +28,8 @@ from core.inference import get_inference_backend logger = get_logger(__name__) +_LLAMA_CPP_SCRIPTS_WARNING_EMITTED = False + def _is_wsl(): """Detect if running under Windows Subsystem for Linux.""" @@ -529,6 +531,31 @@ class ExportBackend: # Convert quantization method to lowercase for unsloth quant_method = quantization_method.lower() + # Pin convert_hf_to_gguf.py to the same llama.cpp ref as the + # llama-quantize binary (Studio installs at a tagged ref via + # setup.sh) so it can't drift past the pinned binary's gguf API. + # Set before both branches; hub-only export has save_directory == "". + global _LLAMA_CPP_SCRIPTS_WARNING_EMITTED + try: + from unsloth_zoo.llama_cpp import ( + LLAMA_CPP_DEFAULT_DIR, + _resolve_local_convert_script, # noqa: F401 + ) + + os.environ.setdefault( + "UNSLOTH_LLAMA_CPP_SCRIPTS_DIR", LLAMA_CPP_DEFAULT_DIR + ) + except ImportError: + if not _LLAMA_CPP_SCRIPTS_WARNING_EMITTED: + logger.warning( + "Unsloth: installed unsloth_zoo does not honor " + "UNSLOTH_LLAMA_CPP_SCRIPTS_DIR; convert_hf_to_gguf.py will " + "still be downloaded from llama.cpp master and may drift " + "past the pinned llama-quantize binary. Upgrade unsloth_zoo " + "to activate the local script pin." + ) + _LLAMA_CPP_SCRIPTS_WARNING_EMITTED = True + # Save locally if requested if save_directory: save_directory = str(resolve_export_dir(save_directory)) diff --git a/studio/backend/core/export/orchestrator.py b/studio/backend/core/export/orchestrator.py index 206dbd6dbb..82de925592 100644 --- a/studio/backend/core/export/orchestrator.py +++ b/studio/backend/core/export/orchestrator.py @@ -163,21 +163,28 @@ class ExportOrchestrator: def _spawn_subprocess(self, config: dict) -> None: """Spawn a new export subprocess.""" + from utils.native_path_leases import ( + native_path_secret_removed_for_child_start, + run_without_native_path_secret, + ) + from .worker import run_export_process - self._cmd_queue = _CTX.Queue() - self._resp_queue = _CTX.Queue() + with native_path_secret_removed_for_child_start(): + self._cmd_queue = _CTX.Queue() + self._resp_queue = _CTX.Queue() - self._proc = _CTX.Process( - target = run_export_process, - kwargs = { - "cmd_queue": self._cmd_queue, - "resp_queue": self._resp_queue, - "config": config, - }, - daemon = True, - ) - self._proc.start() + self._proc = _CTX.Process( + target = run_without_native_path_secret, + args = (run_export_process,), + kwargs = { + "cmd_queue": self._cmd_queue, + "resp_queue": self._resp_queue, + "config": config, + }, + daemon = True, + ) + self._proc.start() logger.info("Export subprocess started (pid=%s)", self._proc.pid) def _shutdown_subprocess(self, timeout: float = 10.0) -> None: diff --git a/studio/backend/core/inference/audio_codecs.py b/studio/backend/core/inference/audio_codecs.py index 895b112e85..df3bf27c16 100644 --- a/studio/backend/core/inference/audio_codecs.py +++ b/studio/backend/core/inference/audio_codecs.py @@ -17,6 +17,7 @@ from typing import Optional, Tuple import numpy as np import torch +from utils.native_path_leases import child_env_without_native_path_secret from utils.subprocess_compat import ( windows_hidden_subprocess_kwargs as _windows_hidden_subprocess_kwargs, ) @@ -105,6 +106,7 @@ class AudioCodecManager: spark_code_dir, ], check = True, + env = child_env_without_native_path_secret(), **_windows_hidden_subprocess_kwargs(), ) @@ -143,6 +145,7 @@ class AudioCodecManager: outetts_code_dir, ], check = True, + env = child_env_without_native_path_secret(), **_windows_hidden_subprocess_kwargs(), ) # Remove files that pull in heavy / incompatible dependencies diff --git a/studio/backend/core/inference/llama_cpp.py b/studio/backend/core/inference/llama_cpp.py index c320f03b2c..f768764c22 100644 --- a/studio/backend/core/inference/llama_cpp.py +++ b/studio/backend/core/inference/llama_cpp.py @@ -11,6 +11,7 @@ through its OpenAI-compatible /v1/chat/completions endpoint. import atexit import contextlib import json +import os import re import struct import structlog @@ -22,11 +23,12 @@ import sys import threading import time from pathlib import Path -from typing import Generator, Optional +from typing import Generator, List, Optional from urllib.parse import urlparse import httpx +from utils.native_path_leases import child_env_without_native_path_secret from utils.subprocess_compat import ( windows_hidden_subprocess_kwargs as _windows_hidden_subprocess_kwargs, ) @@ -53,6 +55,21 @@ _INTENT_SIGNAL = re.compile( r")" ) _MAX_REPROMPTS = 3 + +# Without max_tokens, llama-server defaults to n_predict = n_ctx (up to +# 262144 for Qwen3.5), producing many-minute zombie decodes when cancel +# fails. t_max_predict_ms is a wall-clock backstop applied unconditionally, +# but the llama.cpp README notes it ONLY fires after a newline has been +# generated -- a model stuck in a long unbroken non-newline sequence is +# unbounded by it. So we still want a token cap as the front-line limiter. +# +# The cap is the model's effective context length when we know it, +# falling back to a generous floor when metadata is unavailable. 4096 was +# too low: Qwen3 / gpt-oss reasoning traces routinely exceed it, and any +# OpenAI-API caller that omits max_tokens (langchain, llama-index, raw +# curl) sees responses silently truncated mid-sentence. +_DEFAULT_MAX_TOKENS_FLOOR = 32768 +_DEFAULT_T_MAX_PREDICT_MS = 600_000 # 10 min _REPROMPT_MAX_CHARS = 2000 # ── Pre-compiled patterns for GGUF shard detection ─────────── @@ -60,6 +77,238 @@ _SHARD_FULL_RE = re.compile(r"^(.*)-(\d{5})-of-(\d{5})\.gguf$") _SHARD_RE = re.compile(r"^(.*)-\d{5}-of-\d{5}\.gguf$") +# ── Sliding-window-pattern resolver ─────────────────────────── +# Resolves the per-layer SWA mask when a GGUF reports a sliding window +# but no `sliding_window_pattern` field. Tier order in +# `_resolve_swa_pattern`: GGUF metadata, on-disk cache, bootstrap dict +# below, transformers introspection, HF Hub config.json, legacy 1/4 +# fallback. Period N means layer i is SWA iff `(i + 1) % N != 0`, +# matching transformers. Skipped on purpose: phi3 (no key/val length +# in GGUF, window >= ctx anyway), qwen2 family (converter strips +# sliding_window when use_sliding_window=False), mistral v0.1/v0.2 +# (all-SWA can't be expressed as a period). +_BOOTSTRAP_SWA_DEFAULTS: dict[str, int] = { + "gemma2": 2, # Gemma2Config.sliding_window_pattern + "gemma3": 6, # Gemma3TextConfig.sliding_window_pattern + "gemma3n": 5, # text_config.layer_types: SWA*4 + FULL + "gpt_oss": 2, # text_config.layer_types: alternating + "cohere2": 4, # Cohere2Config.sliding_window_pattern +} + +# Process-wide cache backed by JSON on disk. Values are int period or +# list[bool] mask. Lazy-loaded. +_SWA_CACHE: Optional[dict] = None +_SWA_CACHE_LOCK = threading.Lock() + + +def _swa_cache_path() -> Path: + home = os.environ.get("UNSLOTH_STUDIO_HOME") or os.environ.get("STUDIO_HOME") + base = Path(home) if home else Path.home() / ".unsloth" / "studio" + return base / "swa_cache.json" + + +def _load_swa_cache() -> dict: + global _SWA_CACHE + with _SWA_CACHE_LOCK: + if _SWA_CACHE is not None: + return _SWA_CACHE + try: + with open(_swa_cache_path()) as f: + _SWA_CACHE = json.load(f) + if not isinstance(_SWA_CACHE, dict): + _SWA_CACHE = {} + except (FileNotFoundError, json.JSONDecodeError, OSError): + _SWA_CACHE = {} + return _SWA_CACHE + + +def _save_swa_cache(cache: dict) -> None: + try: + path = _swa_cache_path() + path.parent.mkdir(parents = True, exist_ok = True) + tmp = path.with_suffix(".json.tmp") + with open(tmp, "w") as f: + json.dump(cache, f, indent = 2, sort_keys = True) + tmp.replace(path) + except OSError: + pass + + +def _period_from_layer_types(layer_types: list) -> Optional[int]: + """Smallest period N where `(i+1) % N != 0` matches the SWA mask, + or None if no fixed period fits.""" + if not layer_types: + return None + is_swa = ["full" not in str(t).lower() for t in layer_types] + n = len(is_swa) + for N in range(1, n + 1): + if all(((i + 1) % N != 0) == is_swa[i] for i in range(n)): + return N + return None + + +def _fetch_swa_entry_from_hf(repo_id: str) -> Optional[object]: + try: + from huggingface_hub import hf_hub_download + + cfg_path = hf_hub_download(repo_id, "config.json", repo_type = "model") + with open(cfg_path) as f: + cfg = json.load(f) + except Exception: + return None + + src = cfg.get("text_config") if isinstance(cfg.get("text_config"), dict) else cfg + period = src.get("sliding_window_pattern") + if isinstance(period, int) and period > 0: + return period + lt = src.get("layer_types") + if isinstance(lt, list) and lt: + return _period_from_layer_types(lt) or [ + "full" not in str(t).lower() for t in lt + ] + return None + + +def _arch_aliases(arch: str) -> tuple: + # GGUF emits `falcon-h1`; HF model_type is `falcon_h1`. Normalise both ways. + seen = [] + for a in (arch, arch.replace("-", "_"), arch.replace("_", "-")): + if a and a not in seen: + seen.append(a) + return tuple(seen) + + +def _swa_entry_from_config_obj(cfg) -> Optional[object]: + src = getattr(cfg, "text_config", None) or cfg + period = getattr(src, "sliding_window_pattern", None) + if isinstance(period, int) and period > 0: + return period + lt = getattr(src, "layer_types", None) + if isinstance(lt, list) and lt: + return _period_from_layer_types(lt) or [ + "full" not in str(t).lower() for t in lt + ] + return None + + +_SWA_PATTERN_SOURCE_RE = re.compile( + r"sliding_window_pattern\s*(?::\s*[\w\[\], ]*)?\s*=\s*(\d+)" +) + + +def _resolve_swa_entry_from_transformers(arch: str) -> Optional[object]: + """Default-instantiate the matching Config; on failure, regex-parse + its source for `sliding_window_pattern = N`.""" + try: + from transformers.models.auto.configuration_auto import ( + CONFIG_MAPPING, + CONFIG_MAPPING_NAMES, + ) + except Exception: + return None + + cfg_class = None + for alias in _arch_aliases(arch): + if alias in CONFIG_MAPPING_NAMES: + try: + cfg_class = CONFIG_MAPPING[alias] + break + except Exception: + cfg_class = None + if cfg_class is None: + return None + + try: + if (entry := _swa_entry_from_config_obj(cfg_class())) is not None: + return entry + except Exception: + pass + + import inspect + + candidates = [cfg_class] + text_cfg_class = getattr(cfg_class, "sub_configs", {}).get("text_config") + if text_cfg_class is not None: + candidates.append(text_cfg_class) + for cls in candidates: + try: + src = inspect.getsource(cls) + except (OSError, TypeError): + continue + if m := _SWA_PATTERN_SOURCE_RE.search(src): + period = int(m.group(1)) + if period > 0: + return period + return None + + +def _resolve_swa_pattern( + arch: Optional[str], + n_layers: Optional[int], + source_repo_candidates: tuple = (), + *, + allow_network: Optional[bool] = None, +) -> Optional[list]: + if not arch or not n_layers: + return None + if allow_network is None: + allow_network = os.environ.get("UNSLOTH_STUDIO_OFFLINE", "0") not in ( + "1", + "true", + "True", + "yes", + ) + + cache = _load_swa_cache() + + def _entry_to_mask(entry): + if isinstance(entry, int) and entry > 0: + return [(i + 1) % entry != 0 for i in range(n_layers)] + if isinstance(entry, list) and entry: + return [bool(entry[i % len(entry)]) for i in range(n_layers)] + return None + + def _persist(entry): + with _SWA_CACHE_LOCK: + cache[arch] = entry + _save_swa_cache(cache) + + if (entry := cache.get(arch)) is not None: + if (mask := _entry_to_mask(entry)) is not None: + return mask + + if (entry := _BOOTSTRAP_SWA_DEFAULTS.get(arch)) is not None: + return _entry_to_mask(entry) + + entry = _resolve_swa_entry_from_transformers(arch) + if entry is not None: + _persist(entry) + return _entry_to_mask(entry) + + # Tier 3: live HF fetch (with persistent caching of the result) + if allow_network: + for repo_id in source_repo_candidates: + if not repo_id: + continue + entry = _fetch_swa_entry_from_hf(repo_id) + if entry is not None: + _persist(entry) + return _entry_to_mask(entry) + + return None + + +def _hf_repo_from_url(url: Optional[str]) -> Optional[str]: + """Strip `https://huggingface.co/owner/name(/...)` to `owner/name`.""" + if not url or "huggingface.co/" not in url: + return None + tail = url.split("huggingface.co/", 1)[1].rstrip("/") + parts = tail.split("/") + if len(parts) < 2: + return None + return f"{parts[0]}/{parts[1]}" + + # Model size extraction — lazy import to avoid pulling in transformers # at module level. See PR description for the full explanation. def _extract_model_size_b(model_id: str): @@ -199,17 +448,24 @@ class LlamaCppBackend: # KV-cache estimation fields (populated by _read_gguf_metadata) self._n_layers: Optional[int] = None self._n_kv_heads: Optional[int] = None + self._n_kv_heads_by_layer: Optional[list[int]] = None self._n_heads: Optional[int] = None self._embedding_length: Optional[int] = None - # Architecture-aware KV fields (8 new fields for 5-path estimation) + # Architecture-aware KV fields for 5-path estimation self._kv_key_length: Optional[int] = None self._kv_value_length: Optional[int] = None self._sliding_window: Optional[int] = None + self._sliding_window_pattern: Optional[list[bool]] = None self._full_attention_interval: Optional[int] = None self._kv_lora_rank: Optional[int] = None self._key_length_mla: Optional[int] = None + self._kv_key_length_swa: Optional[int] = None + self._kv_value_length_swa: Optional[int] = None self._ssm_inner_size: Optional[int] = None self._ssm_state_size: Optional[int] = None + # Last N layers reuse KV from earlier layers and don't allocate + # their own cache (Gemma 3n / Gemma 4: .attention.shared_kv_layers). + self._shared_kv_layers: Optional[int] = None self._lock = threading.Lock() self._stdout_lines: list[str] = [] self._stdout_thread: Optional[threading.Thread] = None @@ -549,14 +805,24 @@ class LlamaCppBackend: @staticmethod def _get_gpu_free_memory() -> list[tuple[int, int]]: - """Query free memory per GPU via nvidia-smi. + """Query free memory per GPU. - Returns list of (gpu_index, free_mib) sorted by index. - Respects CUDA_VISIBLE_DEVICES if set. - Returns empty list if nvidia-smi is not available. + Order: + 1. ``nvidia-smi`` (NVIDIA CUDA hosts) -- respects + ``CUDA_VISIBLE_DEVICES``. + 2. ``torch.cuda.mem_get_info`` -- universal fallback that + works on AMD ROCm too because the HIP runtime + reuses the entire ``torch.cuda.*`` namespace. Covers the + AMD case for issue #5106 (nvidia-smi-only probe silently + returned [] on AMD hosts) and also rescues NVIDIA hosts + where ``nvidia-smi`` is missing from PATH. + + Returns list of (gpu_index, free_mib) sorted by index. Empty + list if no supported GPU is reachable. """ import os + # ── NVIDIA via nvidia-smi ──────────────────────────────────── try: result = subprocess.run( [ @@ -567,32 +833,98 @@ class LlamaCppBackend: capture_output = True, text = True, timeout = 10, + env = child_env_without_native_path_secret(), **_windows_hidden_subprocess_kwargs(), ) - if result.returncode != 0: - return [] - - # Parse which GPUs are allowed by existing CUDA_VISIBLE_DEVICES - allowed = None - cvd = os.environ.get("CUDA_VISIBLE_DEVICES") - if cvd is not None and cvd.strip(): - try: - allowed = set(int(x.strip()) for x in cvd.split(",")) - except ValueError: - pass # Non-numeric (e.g., "GPU-uuid"), ignore filter - - gpus = [] - for line in result.stdout.strip().splitlines(): - parts = line.split(",") - if len(parts) == 2: - idx = int(parts[0].strip()) - free_mib = int(parts[1].strip()) - if allowed is not None and idx not in allowed: - continue - gpus.append((idx, free_mib)) - return gpus + if result.returncode == 0: + allowed: Optional[set[int]] = None + cvd = os.environ.get("CUDA_VISIBLE_DEVICES") + if cvd is not None: + try: + # `if x.strip()` filters trailing-comma masks like + # "0,1," which would otherwise raise ValueError on + # an empty token. An explicitly empty mask (CVD="") + # yields an empty `allowed` set so all GPUs are + # filtered out, matching the codebase convention. + allowed = set( + int(x.strip()) for x in cvd.split(",") if x.strip() + ) + except ValueError: + pass + gpus: list[tuple[int, int]] = [] + for line in result.stdout.strip().splitlines(): + parts = line.split(",") + if len(parts) == 2: + idx = int(parts[0].strip()) + free_mib = int(parts[1].strip()) + if allowed is not None and idx not in allowed: + continue + gpus.append((idx, free_mib)) + # Match the docstring's sort-by-id guarantee. nvidia-smi + # almost always returns sorted output, but driver order + # is not formally guaranteed. + gpus.sort(key = lambda g: g[0]) + if gpus: + return gpus except Exception as e: - logger.debug(f"Failed to query GPU free memory via nvidia-smi: {e}") + logger.debug(f"nvidia-smi probe failed: {e}") + + # ── Torch fallback (covers AMD ROCm and missing nvidia-smi) ── + try: + import torch + + if not hasattr(torch, "cuda") or not torch.cuda.is_available(): + return [] + if not hasattr(torch.cuda, "mem_get_info"): + return [] + # torch.cuda enumerates GPUs RELATIVE to the visibility mask. + # On NVIDIA builds the mask is CUDA_VISIBLE_DEVICES; on AMD + # ROCm builds it is HIP_VISIBLE_DEVICES (or ROCR_VISIBLE_DEVICES + # if HIP is unset). Downstream we feed these IDs back into the + # llama-server subprocess as CVD, so we must translate visible + # ordinals back to physical indices first; otherwise launching + # with ``CUDA_VISIBLE_DEVICES=2,3`` would get rewritten to + # ``CUDA_VISIBLE_DEVICES=0,1`` and target the wrong GPUs. + physical_ids: Optional[list[int]] = None + # Match the codebase convention in + # ``utils/hardware/hardware.py::_get_parent_visible_gpu_spec``: + # treat an explicitly empty mask (``HIP_VISIBLE_DEVICES=""``) + # as "set to no GPUs" rather than falling through to the next + # var. ``or`` would coerce empty string to falsy and silently + # promote the wrong source. + if getattr(torch.version, "hip", None) is not None: + hip_v = os.environ.get("HIP_VISIBLE_DEVICES") + rocr_v = os.environ.get("ROCR_VISIBLE_DEVICES") + cvd = ( + hip_v + if hip_v is not None + else rocr_v + if rocr_v is not None + else os.environ.get("CUDA_VISIBLE_DEVICES") + ) + else: + cvd = os.environ.get("CUDA_VISIBLE_DEVICES") + if cvd is not None: + try: + # Empty mask (CVD="") yields an empty list so the + # below loop produces no GPUs, consistent with the + # nvidia-smi path and utils/hardware/hardware.py. + physical_ids = [int(x.strip()) for x in cvd.split(",") if x.strip()] + except ValueError: + physical_ids = None + gpus = [] + for ordinal in range(torch.cuda.device_count()): + free_bytes, _total_bytes = torch.cuda.mem_get_info(ordinal) + idx = ( + physical_ids[ordinal] + if physical_ids is not None and ordinal < len(physical_ids) + else ordinal + ) + gpus.append((idx, free_bytes // (1024 * 1024))) + # Match the nvidia-smi path's docstring guarantee of sorted-by-id. + return sorted(gpus, key = lambda g: g[0]) + except Exception as e: + logger.debug(f"torch GPU probe failed: {e}") return [] @staticmethod @@ -652,13 +984,29 @@ class LlamaCppBackend: # New-style: need both explicit key AND value dimensions if self._kv_key_length is not None and self._kv_value_length is not None: return True - # Legacy: need embedding_length + head count + # Legacy: need embedding_length + a head count (scalar or per-layer). return self._embedding_length is not None and ( - self._n_kv_heads is not None or self._n_heads is not None + self._n_kv_heads is not None + or self._n_heads is not None + or self._n_kv_heads_by_layer is not None ) + def _kv_heads_for_layer(self, layer_idx: int, fallback: int) -> int: + if self._n_kv_heads_by_layer is not None and layer_idx < len( + self._n_kv_heads_by_layer + ): + return self._n_kv_heads_by_layer[layer_idx] + return fallback + def _estimate_kv_cache_bytes( - self, n_ctx: int, cache_type_kv: Optional[str] = None + self, + n_ctx: int, + cache_type_kv: Optional[str] = None, + *, + swa_full: bool = False, + n_parallel: int = 1, + kv_unified: bool = True, + ctx_checkpoints: int = 0, ) -> int: """Estimate KV cache VRAM for a given context length. @@ -669,12 +1017,34 @@ class LlamaCppBackend: 4. GQA -- standard full KV with explicit key/value dimensions 5. Legacy -- fallback using embed // n_heads + Server-flag knobs (mirror llama-server's CLI): + swa_full -- ``--swa-full``: force SWA layers to cache the + full ``n_ctx`` (collapses path 3 to path 4 + sizing for the SWA layers). + n_parallel -- ``--parallel``: number of server slots. + Verified empirically against llama-server: + non-SWA layers stay constant (cells split + across slots), SWA layers scale linearly + (per-slot window). + kv_unified -- ``--kv-unified`` (default on): retained for + API forward-compat. Currently a no-op for + memory math because the unified buffer total + matches per-slot buffers in measured cases. + ctx_checkpoints -- ``--ctx-checkpoints``: SWA snapshot count per + slot (PR #15293). Each snapshot stores one + sliding-window of state per SWA layer. + Returns 0 if metadata is insufficient for estimation. """ if not self._can_estimate_kv() or n_ctx <= 0: return 0 n_layers = self._n_layers # type: ignore[assignment] + # Gemma 3n / Gemma 4 reuse KV from earlier layers in the last + # ``shared_kv_layers`` blocks -- those don't allocate their own + # cache. Floor at 1 so a misconfigured GGUF can't zero out KV. + shared = self._shared_kv_layers or 0 + n_layers_kv = max(1, n_layers - shared) n_kv = self._n_kv_heads or self._n_heads or 1 # type: ignore[assignment] # Bytes per element depends on KV cache quantization @@ -690,6 +1060,8 @@ class LlamaCppBackend: "iq4_nl": 0.5625, }.get(cache_type_kv or "f16", 2.0) + slots = max(1, n_parallel) + # Path 1: MLA (DeepSeek-V2/V3, GLM-4.7, GLM-5, Kimi-K2.5) # MLA stores one compressed KV latent per token/layer (shared across heads). # V is reconstructed from the latent on the fly -- no separate V cache. @@ -700,7 +1072,7 @@ class LlamaCppBackend: n_kv_mla = self._n_kv_heads or 1 rope_dim = self._key_length_mla or 64 key_len = self._kv_key_length or (self._kv_lora_rank + rope_dim) - return int(n_layers * n_ctx * n_kv_mla * key_len * bpe) + return int(n_layers_kv * n_ctx * n_kv_mla * key_len * bpe) key_len = self._kv_key_length val_len = self._kv_value_length @@ -718,11 +1090,19 @@ class LlamaCppBackend: head_dim = self._embedding_length // self._n_heads if self._n_heads else 128 # type: ignore[operator] return int(n_attn * n_ctx * n_kv * 2 * head_dim * bpe) - # Path 3: Sliding Window (Gemma-3, gpt-oss) - # SWA layers only cache min(ctx, window) tokens; global layers cache full ctx. - # Most SWA architectures use few global layers (e.g., Gemma-3 uses 1 in 6). - # Without an explicit field, we conservatively assume 1/4 of layers are global - # which is still far more accurate than the legacy formula (which ignores SWA). + # Path 3: Sliding window (Gemma 2/3/3n/4, gpt-oss, Cohere2 ...). + # Pattern is filled in by the resolver at parse time; if absent, + # falls through to the legacy 1/4-global heuristic below. + # Per-layer-type ``--parallel N`` accounting (verified empirically + # against ``llama-server``): + # * non-SWA layers: total cells = n_ctx, partitioned across + # slots -> total memory CONSTANT in slots. + # * SWA layers: per-slot cells = 2 * sliding_window + # (capped at n_ctx and at per_slot_ctx + # when ctx is split among many slots) -> + # total memory grows LINEARLY in slots. + # ``--swa-full`` forces full n_ctx for SWA layers instead. + # ``--ctx-checkpoints N`` adds N snapshots per SWA layer per slot. if ( self._sliding_window is not None and self._sliding_window > 0 @@ -730,20 +1110,72 @@ class LlamaCppBackend: and val_len is not None ): swa = self._sliding_window - n_global = max(1, n_layers // 4) - n_swa = n_layers - n_global + per_slot_ctx = max(1, n_ctx // slots) + # ``--swa-full`` makes SWA layers cache the full context just + # like non-SWA: cells get partitioned across slots, so per-slot + # cells = per_slot_ctx and the slots*per-slot product collapses + # back to the constant ``n_ctx`` total. Otherwise SWA caches + # 2*sliding_window per slot, clamped at the per-slot ctx. + swa_cells_per_slot = ( + per_slot_ctx if swa_full else min(n_ctx, 2 * swa, per_slot_ctx) + ) + key_len_swa = self._kv_key_length_swa or key_len + val_len_swa = self._kv_value_length_swa or val_len + if self._sliding_window_pattern is not None: + global_bytes = 0.0 # constant across slots + swa_bytes_per_slot = 0.0 # multiplied by slots + checkpoint_extra_per_slot = 0.0 + # Iterate only over layers that allocate their own KV; + # the trailing ``shared`` layers reuse earlier caches. + for layer_idx in range(n_layers_kv): + layer_n_kv = self._kv_heads_for_layer(layer_idx, n_kv) + is_swa = ( + layer_idx < len(self._sliding_window_pattern) + and self._sliding_window_pattern[layer_idx] + ) + if is_swa: + swa_bytes_per_slot += ( + swa_cells_per_slot + * layer_n_kv + * (key_len_swa + val_len_swa) + * bpe + ) + if ctx_checkpoints > 0 and not swa_full: + checkpoint_extra_per_slot += ( + ctx_checkpoints + * swa + * layer_n_kv + * (key_len_swa + val_len_swa) + * bpe + ) + else: + global_bytes += n_ctx * layer_n_kv * (key_len + val_len) * bpe + return int( + global_bytes + + slots * (swa_bytes_per_slot + checkpoint_extra_per_slot) + ) + n_global = max(1, n_layers_kv // 4) + n_swa = n_layers_kv - n_global kv_per_token = n_kv * (key_len + val_len) * bpe + kv_per_token_swa = n_kv * (key_len_swa + val_len_swa) * bpe + global_bytes = n_global * n_ctx * kv_per_token + swa_bytes_per_slot = n_swa * swa_cells_per_slot * kv_per_token_swa + checkpoint_extra_per_slot = ( + ctx_checkpoints * n_swa * swa * kv_per_token_swa + if ctx_checkpoints > 0 and not swa_full + else 0.0 + ) return int( - n_global * n_ctx * kv_per_token + n_swa * min(n_ctx, swa) * kv_per_token + global_bytes + slots * (swa_bytes_per_slot + checkpoint_extra_per_slot) ) # Path 4: Standard GQA with explicit key/value dimensions if key_len is not None and val_len is not None: - return int(n_layers * n_ctx * n_kv * (key_len + val_len) * bpe) + return int(n_layers_kv * n_ctx * n_kv * (key_len + val_len) * bpe) # Path 5: Legacy fallback (old GGUFs without explicit dimensions) head_dim = self._embedding_length // self._n_heads if self._n_heads else 128 # type: ignore[operator] - return int(2 * n_kv * head_dim * n_layers * n_ctx * bpe) + return int(2 * n_kv * head_dim * n_layers_kv * n_ctx * bpe) def _fit_context_to_vram( self, @@ -752,6 +1184,12 @@ class LlamaCppBackend: model_size_bytes: int, cache_type_kv: Optional[str] = None, min_ctx: int = 4096, + *, + swa_full: bool = False, + n_parallel: int = 1, + kv_unified: bool = True, + ctx_checkpoints: int = 0, + kv_on_gpu: bool = True, ) -> int: """Return the largest context length that fits in GPU VRAM. @@ -759,6 +1197,11 @@ class LlamaCppBackend: threshold -- 10% reserved for compute buffers, CUDA context, scratch space, flash-attn workspace, etc.). If the model weights alone don't fit, returns min_ctx unchanged. + + ``kv_on_gpu`` mirrors ``--kv-offload`` (default on). When False + the KV cache lives in CPU RAM and doesn't compete with weights + for VRAM; the requested context is honored verbatim. The other + keyword args mirror ``_estimate_kv_cache_bytes``. """ if not self._can_estimate_kv(): logger.debug( @@ -768,11 +1211,22 @@ class LlamaCppBackend: ) return requested_ctx + # KV lives off-GPU: no VRAM accounting needed for the cache itself. + if not kv_on_gpu: + return requested_ctx + + kv_kwargs = dict( + swa_full = swa_full, + n_parallel = n_parallel, + kv_unified = kv_unified, + ctx_checkpoints = ctx_checkpoints, + ) + budget_bytes = available_mib * 1024 * 1024 * 0.90 model_footprint = model_size_bytes # Check if requested context already fits - kv = self._estimate_kv_cache_bytes(requested_ctx, cache_type_kv) + kv = self._estimate_kv_cache_bytes(requested_ctx, cache_type_kv, **kv_kwargs) if model_footprint + kv <= budget_bytes: return requested_ctx @@ -794,7 +1248,7 @@ class LlamaCppBackend: best = effective_min while lo <= hi: mid = (lo + hi) // 2 - kv = self._estimate_kv_cache_bytes(mid, cache_type_kv) + kv = self._estimate_kv_cache_bytes(mid, cache_type_kv, **kv_kwargs) if kv <= remaining: best = mid lo = mid + 1 @@ -927,6 +1381,19 @@ class LlamaCppBackend: for _ in range(alen): LlamaCppBackend._gguf_skip_value(f, atype) + @staticmethod + def _gguf_read_array_value(f, atype: int, alen: int) -> Optional[list]: + if atype == 4: # UINT32 + return [struct.unpack(" None: """Read context_length, architecture params, and chat_template from a GGUF header. @@ -945,23 +1412,44 @@ class LlamaCppBackend: self._supports_tools = False self._n_layers = None self._n_kv_heads = None + self._n_kv_heads_by_layer = None self._n_heads = None self._embedding_length = None self._kv_key_length = None self._kv_value_length = None self._sliding_window = None + self._sliding_window_pattern = None self._full_attention_interval = None self._kv_lora_rank = None self._key_length_mla = None + self._kv_key_length_swa = None + self._kv_value_length_swa = None self._ssm_inner_size = None self._ssm_state_size = None + self._shared_kv_layers = None try: - WANTED = {"general.architecture", "tokenizer.chat_template"} + WANTED = { + "general.architecture", + "tokenizer.chat_template", + # Source-repo hints for the SWA resolver's HF fallback. + "general.source.huggingface.repository", + "general.source.url", + "general.source.repo_url", + "general.base_model.0.repo_url", + "general.base_model.0.organization", + "general.base_model.0.name", + "general.basename", + "general.organization", + "general.size_label", + "general.finetune", + } # Additional arch-specific keys are added dynamically once # we know the architecture name. arch_keys: dict[str, str] = {} # gguf_key -> attribute name arch = None + sliding_window_pattern_period: Optional[int] = None + general: dict[str, str] = {} with open(gguf_path, "rb") as f: magic = struct.unpack(" bool: """ Start llama-server with a GGUF model. @@ -1411,8 +1992,11 @@ class LlamaCppBackend: pool_mib, model_size, cache_type_kv, + n_parallel = n_parallel, + ) + kv = self._estimate_kv_cache_bytes( + capped, cache_type_kv, n_parallel = n_parallel ) - kv = self._estimate_kv_cache_bytes(capped, cache_type_kv) total_mib = (model_size + kv) / (1024 * 1024) if total_mib <= pool_mib * 0.90: best_cap = max(best_cap, capped) @@ -1436,7 +2020,7 @@ class LlamaCppBackend: # have surfaced the "might be slower" warning before # the user submitted a ctx above the fit ceiling. requested_total = model_size + self._estimate_kv_cache_bytes( - effective_ctx, cache_type_kv + effective_ctx, cache_type_kv, n_parallel = n_parallel ) gpu_indices, use_fit = self._select_gpus(requested_total, gpus) # No silent shrink: effective_ctx stays == n_ctx. @@ -1451,8 +2035,11 @@ class LlamaCppBackend: pool_mib, model_size, cache_type_kv, + n_parallel = n_parallel, + ) + kv = self._estimate_kv_cache_bytes( + capped, cache_type_kv, n_parallel = n_parallel ) - kv = self._estimate_kv_cache_bytes(capped, cache_type_kv) total_mib = (model_size + kv) / (1024 * 1024) if total_mib <= pool_mib * 0.90: effective_ctx = capped @@ -1485,7 +2072,9 @@ class LlamaCppBackend: ) if effective_ctx < original_ctx: - kv_est = self._estimate_kv_cache_bytes(effective_ctx, cache_type_kv) + kv_est = self._estimate_kv_cache_bytes( + effective_ctx, cache_type_kv, n_parallel = n_parallel + ) logger.info( f"Context auto-reduced: {original_ctx} -> {effective_ctx} " f"(model: {model_size / (1024**3):.1f} GB, " @@ -1493,7 +2082,7 @@ class LlamaCppBackend: ) kv_cache_bytes = self._estimate_kv_cache_bytes( - effective_ctx, cache_type_kv + effective_ctx, cache_type_kv, n_parallel = n_parallel ) logger.info( f"GGUF size: {model_size / (1024**3):.1f} GB, " @@ -1528,8 +2117,10 @@ class LlamaCppBackend: # Model fits on selected GPU(s) -- offload all layers cmd.extend(["-ngl", "-1"]) - if n_threads is not None: - cmd.extend(["--threads", str(n_threads)]) + # -1 = llama.cpp auto-detect (physical cores). Pass explicitly so we + # do not inherit llama-server's internal default, which has historically + # varied (hardware concurrency incl. hyperthreads on some builds). + cmd.extend(["--threads", str(n_threads if n_threads is not None else -1)]) # Always enable Jinja chat template rendering for proper template support cmd.extend(["--jinja"]) @@ -1561,7 +2152,7 @@ class LlamaCppBackend: # existing text (code refactoring, summarization, reasoning). # For general chat with low repetition, overhead is ~5 ms. # - # Benchmarks from llama.cpp PRs #18471, #19164: + # Benchmarks from upstream llama.cpp speculative-decoding PRs: # Scenario | Without | With | Speedup # gpt-oss-120b code refactor | 181 t/s | 446 t/s | 2.5x # Qwen3-235B offloaded | 12 t/s | 21 t/s | 1.8x @@ -1574,11 +2165,21 @@ class LlamaCppBackend: # ref: https://github.com/ggml-org/llama.cpp/blob/master/docs/speculative.md # ref: https://github.com/ggml-org/llama.cpp/pull/19164 # ref: https://github.com/ggml-org/llama.cpp/pull/18471 + # ``"default"`` -> let llama-server pick a sensible spec + # config via ``--spec-default``. Explicit type names are + # passed through with the manual draft tuning we've shipped + # historically so power users keep their overrides. _valid_spec_types = {"ngram-simple", "ngram-mod"} - if speculative_type and speculative_type in _valid_spec_types: - if not is_vision: # spec decoding disabled for vision models - cmd.extend(["--spec-type", speculative_type]) - if speculative_type == "ngram-mod": + normalized_spec = ( + speculative_type.lower().strip() if speculative_type else None + ) + if normalized_spec and normalized_spec != "off" and not is_vision: + if normalized_spec == "default": + cmd.append("--spec-default") + self._speculative_type = "default" + elif normalized_spec in _valid_spec_types: + cmd.extend(["--spec-type", normalized_spec]) + if normalized_spec == "ngram-mod": cmd.extend( [ "--spec-ngram-size-n", @@ -1589,7 +2190,7 @@ class LlamaCppBackend: "64", ] ) - self._speculative_type = speculative_type + self._speculative_type = normalized_spec else: self._speculative_type = None else: @@ -1599,6 +2200,18 @@ class LlamaCppBackend: if chat_template_override: import tempfile + self._chat_template = chat_template_override + flags = detect_reasoning_flags( + self._chat_template, + self._model_identifier, + log_source = "GGUF chat template override", + ) + self._supports_reasoning = flags["supports_reasoning"] + self._reasoning_style = flags["reasoning_style"] + self._reasoning_always_on = flags["reasoning_always_on"] + self._supports_preserve_thinking = flags["supports_preserve_thinking"] + self._supports_tools = flags["supports_tools"] + self._chat_template_file = tempfile.NamedTemporaryFile( mode = "w", suffix = ".jinja", @@ -1651,6 +2264,17 @@ class LlamaCppBackend: else: self._api_key = None + # User-supplied pass-through args go last so llama.cpp's + # last-wins flag parsing lets the user override Studio's + # auto-set tier-2 flags (e.g. --cache-type-k, --spec-type). + # The route layer has already validated this list against + # the managed-flag denylist via validate_extra_args(). + if extra_args: + cmd.extend(str(a) for a in extra_args) + logger.info( + f"Appending user extra args to llama-server: {list(extra_args)}" + ) + _log_cmd = list(cmd) if "--api-key" in _log_cmd: _ki = _log_cmd.index("--api-key") + 1 @@ -1662,7 +2286,7 @@ class LlamaCppBackend: import os import sys - env = os.environ.copy() + env = child_env_without_native_path_secret() binary_dir = str(Path(binary).parent) if sys.platform == "win32": @@ -1746,9 +2370,29 @@ class LlamaCppBackend: f"{new_ld}:{existing_ld}" if existing_ld else new_ld ) - # Pin to selected GPU(s) via CUDA_VISIBLE_DEVICES + # Pin to selected GPU(s). On ROCm, llama-server (and any torch + # in the subprocess) honors HIP_VISIBLE_DEVICES / ROCR_VISIBLE_DEVICES; + # narrowing only CUDA_VISIBLE_DEVICES leaves an AMD child seeing + # the full HIP/ROCR set the parent inherited. if gpu_indices is not None: - env["CUDA_VISIBLE_DEVICES"] = ",".join(str(i) for i in gpu_indices) + pinned = ",".join(str(i) for i in gpu_indices) + env["CUDA_VISIBLE_DEVICES"] = pinned + try: + import torch as _torch + + if getattr(_torch.version, "hip", None) is not None: + env["HIP_VISIBLE_DEVICES"] = pinned + env["ROCR_VISIBLE_DEVICES"] = pinned + except Exception as e: + logger.debug( + "Failed to set ROCm visibility env vars for child: %s", e + ) + + # Defensive kill: if a concurrent load slipped past Phase 1 + # (because its `self._process` was None at the time) and + # already stored a Popen handle here, drop that orphan + # before we overwrite the reference. See issue #5161. + self._kill_process() self._stdout_lines = [] self._process = subprocess.Popen( @@ -1867,16 +2511,21 @@ class LlamaCppBackend: self._speculative_type = None self._n_layers = None self._n_kv_heads = None + self._n_kv_heads_by_layer = None self._n_heads = None self._embedding_length = None self._kv_key_length = None self._kv_value_length = None self._sliding_window = None + self._sliding_window_pattern = None self._full_attention_interval = None self._kv_lora_rank = None self._key_length_mla = None + self._kv_key_length_swa = None + self._kv_value_length_swa = None self._ssm_inner_size = None self._ssm_state_size = None + self._shared_kv_layers = None # Clean up temp chat template file if hasattr(self, "_chat_template_file") and self._chat_template_file: try: @@ -2032,6 +2681,7 @@ class LlamaCppBackend: capture_output = True, text = True, timeout = 5, + env = child_env_without_native_path_secret(), ) if result.returncode != 0: return @@ -2444,8 +3094,15 @@ class LlamaCppBackend: ) if _reasoning_kw is not None: payload["chat_template_kwargs"] = _reasoning_kw - if max_tokens is not None: - payload["max_tokens"] = max_tokens + # Default cap to the model's effective context length when known, + # otherwise the conservative floor. The wall-clock backstop below + # keeps a stuck model from running indefinitely either way. + payload["max_tokens"] = ( + max_tokens + if max_tokens is not None + else (self._effective_context_length or _DEFAULT_MAX_TOKENS_FLOOR) + ) + payload["t_max_predict_ms"] = _DEFAULT_T_MAX_PREDICT_MS if stop: payload["stop"] = stop payload["stream_options"] = {"include_usage": True} @@ -2465,7 +3122,9 @@ class LlamaCppBackend: _auth_headers = ( {"Authorization": f"Bearer {self._api_key}"} if self._api_key else None ) - with httpx.Client(timeout = stream_timeout) as client: + with httpx.Client( + timeout = stream_timeout, limits = httpx.Limits(max_keepalive_connections = 0) + ) as client: with self._stream_with_retry( client, url, @@ -2664,8 +3323,12 @@ class LlamaCppBackend: ) if _reasoning_kw is not None: payload["chat_template_kwargs"] = _reasoning_kw - if max_tokens is not None: - payload["max_tokens"] = max_tokens + payload["max_tokens"] = ( + max_tokens + if max_tokens is not None + else (self._effective_context_length or _DEFAULT_MAX_TOKENS_FLOOR) + ) + payload["t_max_predict_ms"] = _DEFAULT_T_MAX_PREDICT_MS if stop: payload["stop"] = stop @@ -2704,7 +3367,10 @@ class LlamaCppBackend: write = 10, pool = 10, ) - with httpx.Client(timeout = stream_timeout) as client: + with httpx.Client( + timeout = stream_timeout, + limits = httpx.Limits(max_keepalive_connections = 0), + ) as client: with self._stream_with_retry( client, url, @@ -3317,8 +3983,12 @@ class LlamaCppBackend: ) if _reasoning_kw is not None: stream_payload["chat_template_kwargs"] = _reasoning_kw - if max_tokens is not None: - stream_payload["max_tokens"] = max_tokens + stream_payload["max_tokens"] = ( + max_tokens + if max_tokens is not None + else (self._effective_context_length or _DEFAULT_MAX_TOKENS_FLOOR) + ) + stream_payload["t_max_predict_ms"] = _DEFAULT_T_MAX_PREDICT_MS if stop: stream_payload["stop"] = stop stream_payload["stream_options"] = {"include_usage": True} @@ -3337,7 +4007,9 @@ class LlamaCppBackend: _auth_headers = ( {"Authorization": f"Bearer {self._api_key}"} if self._api_key else None ) - with httpx.Client(timeout = stream_timeout) as client: + with httpx.Client( + timeout = stream_timeout, limits = httpx.Limits(max_keepalive_connections = 0) + ) as client: with self._stream_with_retry( client, url, diff --git a/studio/backend/core/inference/llama_server_args.py b/studio/backend/core/inference/llama_server_args.py new file mode 100644 index 0000000000..44c7d542c7 --- /dev/null +++ b/studio/backend/core/inference/llama_server_args.py @@ -0,0 +1,120 @@ +# SPDX-License-Identifier: AGPL-3.0-only +# Copyright 2026-present the Unsloth AI Inc. team. All rights reserved. See /studio/LICENSE.AGPL-3.0 + +"""Validator for user-supplied llama-server pass-through args. + +Studio runs llama-server as a managed subprocess and lets callers pass +extra flags directly (CLI: ``unsloth run ... --top-k 20``; HTTP: +``LoadRequest.llama_extra_args``). This module is the boundary that +rejects only flags Studio fundamentally cannot share with the user -- +model identity, the auth key, and the network endpoint Studio's HTTP +proxy targets. Anything else passes through. + +User-supplied args are appended to ``cmd`` after Studio's auto-set +flags, so llama.cpp's last-wins CLI parsing makes the user's value +override the auto-set one. That covers tunable knobs the user might +reasonably want to override -- ``-c``/``--ctx-size``, +``-np``/``--parallel``, ``-fa``/``--flash-attn``, +``-ngl``/``--gpu-layers``, ``-t``/``--threads``, ``-fit``/``--fit*``, +``--cache-type-k/v``, ``--chat-template-file/-kwargs``, +``--spec-*``, ``--jinja``/``--no-jinja``, +``--no-context-shift``/``--context-shift``, sampling params, etc. + +Reference: https://github.com/ggml-org/llama.cpp/blob/master/tools/server/README.md +""" + +from __future__ import annotations + +from typing import Iterable, Optional + +# Each group is the full set of aliases (short + long) for one +# hard-denied flag, taken from the llama-server README. If llama.cpp +# adds a new alias for an existing denied flag, extend the relevant +# group. +# +# Flags NOT in this list (e.g. -c, --parallel, --flash-attn, -ngl, +# -t/--threads, --jinja, --no-context-shift, --fit*, --cache-type-*, +# --chat-template-*, --spec-*) pass through and override Studio's +# auto-set version via llama.cpp's last-wins CLI parsing. +_DENYLIST_GROUPS: tuple[frozenset[str], ...] = ( + # Model identity -- Studio resolves the model from LoadRequest and + # passes -m / mmproj after downloading from HF if needed. A second + # -m would point at a different model than the one Studio thinks + # is loaded. + frozenset({"-m", "--model"}), + frozenset({"-mu", "--model-url"}), + frozenset({"-dr", "--docker-repo"}), + frozenset({"-hf", "-hfr", "--hf-repo"}), + frozenset({"-hff", "--hf-file"}), + frozenset({"-hfv", "-hfrv", "--hf-repo-v"}), + frozenset({"-hffv", "--hf-file-v"}), + frozenset({"-hft", "--hf-token"}), + frozenset({"-mm", "--mmproj"}), + frozenset({"-mmu", "--mmproj-url"}), + # Networking -- Studio binds llama-server's port and reverse-proxies + # HTTP traffic to it. Retargeting host/port/path/prefix would + # orphan Studio's proxy and the UI would lose the server. + frozenset({"--host"}), + frozenset({"--port"}), + frozenset({"--path"}), + frozenset({"--api-prefix"}), + frozenset({"--reuse-port"}), + # Auth / TLS -- Studio terminates auth at its own layer; an + # upstream --api-key would shadow Studio's UNSLOTH_DIRECT_STREAM + # key, and TLS on llama-server would break the local proxy hop. + frozenset({"--api-key"}), + frozenset({"--api-key-file"}), + frozenset({"--ssl-key-file"}), + frozenset({"--ssl-cert-file"}), + # Single-model server -- Studio runs one model per llama-server + # process and serves its own UI. Enabling multi-model loading or + # llama-server's built-in web UI changes the surface clients see. + frozenset({"--webui", "--no-webui"}), + frozenset({"--models-dir"}), + frozenset({"--models-preset"}), + frozenset({"--models-max"}), + frozenset({"--models-autoload", "--no-models-autoload"}), +) + +_DENYLIST: frozenset[str] = frozenset().union(*_DENYLIST_GROUPS) + + +def _flag_name(token: str) -> Optional[str]: + """Return the flag name for a token, or None if it isn't a flag. + + Peels ``--key=value`` to the bare ``--key``. Plain numeric values + like ``-1`` or ``-0.5`` (e.g. ``--seed -1``) are values, not flags; + llama-server short-form flags always start with a letter. + """ + if not token.startswith("-") or token in {"-", "--"}: + return None + if len(token) >= 2 and (token[1].isdigit() or token[1] == "."): + return None + return token.split("=", 1)[0] + + +def validate_extra_args(args: Optional[Iterable[str]]) -> list[str]: + """Validate user-supplied llama-server args. + + Returns the args as a flat list ready to extend the llama-server + command. Raises ``ValueError`` (with the offending flag in the + message) the moment a token resolves to a Studio-managed flag. + """ + if not args: + return [] + out: list[str] = [] + for raw in args: + token = str(raw) + flag = _flag_name(token) + if flag is not None and flag in _DENYLIST: + raise ValueError( + f"llama-server flag '{flag}' is managed by Unsloth Studio " + f"and cannot be passed as an extra arg" + ) + out.append(token) + return out + + +def is_managed_flag(flag: str) -> bool: + """True if ``flag`` is a Studio-managed llama-server flag.""" + return flag in _DENYLIST diff --git a/studio/backend/core/inference/orchestrator.py b/studio/backend/core/inference/orchestrator.py index cb5d9da34a..5562820f49 100644 --- a/studio/backend/core/inference/orchestrator.py +++ b/studio/backend/core/inference/orchestrator.py @@ -166,23 +166,30 @@ class InferenceOrchestrator: def _spawn_subprocess(self, config: dict) -> None: """Spawn a new inference subprocess.""" + from utils.native_path_leases import ( + native_path_secret_removed_for_child_start, + run_without_native_path_secret, + ) + from .worker import run_inference_process - self._cmd_queue = _CTX.Queue() - self._resp_queue = _CTX.Queue() - self._cancel_event = _CTX.Event() + with native_path_secret_removed_for_child_start(): + self._cmd_queue = _CTX.Queue() + self._resp_queue = _CTX.Queue() + self._cancel_event = _CTX.Event() - self._proc = _CTX.Process( - target = run_inference_process, - kwargs = { - "cmd_queue": self._cmd_queue, - "resp_queue": self._resp_queue, - "cancel_event": self._cancel_event, - "config": config, - }, - daemon = True, - ) - self._proc.start() + self._proc = _CTX.Process( + target = run_without_native_path_secret, + args = (run_inference_process,), + kwargs = { + "cmd_queue": self._cmd_queue, + "resp_queue": self._resp_queue, + "cancel_event": self._cancel_event, + "config": config, + }, + daemon = True, + ) + self._proc.start() logger.info("Inference subprocess started (pid=%s)", self._proc.pid) def _cancel_generation(self) -> None: @@ -708,6 +715,17 @@ class InferenceOrchestrator: def unload_model(self, model_name: str) -> bool: """Unload a model from the subprocess.""" + if model_name in self.loading_models: + logger.info( + "Cancelling in-flight load for model '%s' by terminating subprocess", + model_name, + ) + self._shutdown_subprocess(timeout = 0.5) + self.loading_models.discard(model_name) + self.active_model_name = None + self.models.clear() + return True + if not self._ensure_subprocess_alive(): # No subprocess — just clear local state self.models.pop(model_name, None) diff --git a/studio/backend/core/training/resume.py b/studio/backend/core/training/resume.py new file mode 100644 index 0000000000..165c1c2cf1 --- /dev/null +++ b/studio/backend/core/training/resume.py @@ -0,0 +1,75 @@ +# SPDX-License-Identifier: AGPL-3.0-only +# Copyright 2026-present the Unsloth AI Inc. team. All rights reserved. See /studio/LICENSE.AGPL-3.0 + +"""Helpers for validating resumable training outputs.""" + +from pathlib import Path +from typing import Optional + +from utils.paths import outputs_root, resolve_output_dir + + +def _is_under_outputs(path: Path) -> bool: + resolved = path.resolve(strict = False) + root = outputs_root().resolve(strict = False) + try: + resolved.relative_to(root) + return True + except ValueError: + return False + + +def has_resume_state(path_value: Optional[str]) -> bool: + if not path_value: + return False + return get_resume_checkpoint_path(path_value) is not None + + +def _checkpoint_step(path: Path) -> int: + try: + return int(path.name.removeprefix("checkpoint-")) + except ValueError: + return -1 + + +def get_resume_checkpoint_path(path_value: str) -> Optional[str]: + path = resolve_output_dir(path_value) + if not _is_under_outputs(path) or not path.is_dir(): + return None + if (path / "trainer_state.json").is_file(): + return str(path) + + checkpoints = [ + child + for child in path.glob("checkpoint-*") + if child.is_dir() and (child / "trainer_state.json").is_file() + ] + if not checkpoints: + return None + return str(max(checkpoints, key = _checkpoint_step)) + + +def normalize_resume_output_dir(path_value: str) -> str: + path = resolve_output_dir(path_value) + if not _is_under_outputs(path): + raise ValueError("Resume checkpoint must be inside Studio outputs.") + return str(path) + + +def can_resume_run(run: dict) -> bool: + if run.get("resumed_later"): + return False + + final_step = run.get("final_step") + total_steps = run.get("total_steps") + has_remaining_steps = ( + not isinstance(final_step, int) + or not isinstance(total_steps, int) + or total_steps <= 0 + or final_step < total_steps + ) + return ( + run.get("status") == "stopped" + and has_remaining_steps + and has_resume_state(run.get("output_dir")) + ) diff --git a/studio/backend/core/training/trainer.py b/studio/backend/core/training/trainer.py index c1e2ac4a85..fe8d277ac0 100644 --- a/studio/backend/core/training/trainer.py +++ b/studio/backend/core/training/trainer.py @@ -70,6 +70,7 @@ from utils.paths import ( ) from trl import SFTTrainer, SFTConfig +from utils.native_path_leases import child_env_without_native_path_secret from utils.subprocess_compat import ( windows_hidden_subprocess_kwargs as _windows_hidden_subprocess_kwargs, ) @@ -376,6 +377,7 @@ class UnslothTrainer: def _finalize_training(self, output_dir, label = ""): """Save model after training and update progress. Used by all training branches.""" if self.should_stop and self.save_on_stop: + self.trainer._save_checkpoint(self.trainer.model, trial = None) self.trainer.save_model() self.tokenizer.save_pretrained(output_dir) self._patch_adapter_config(output_dir) @@ -1770,6 +1772,7 @@ class UnslothTrainer: spark_code_dir, ], check = True, + env = child_env_without_native_path_secret(), **_windows_hidden_subprocess_kwargs(), ) @@ -2004,6 +2007,7 @@ class UnslothTrainer: outetts_code_dir, ], check = True, + env = child_env_without_native_path_secret(), **_windows_hidden_subprocess_kwargs(), ) for fpath in [ @@ -2828,7 +2832,9 @@ class UnslothTrainer: total_steps = total, status_message = "Starting CSM training..." ) logger.info(f"CSM training config: {config}\n") - self.trainer.train() + self.trainer.train( + resume_from_checkpoint = training_args.get("resume_from_checkpoint") + ) self._finalize_training(output_dir, "CSM") return @@ -2867,7 +2873,9 @@ class UnslothTrainer: total_steps = total, status_message = "Starting SNAC training..." ) logger.info(f"SNAC training config: {config}\n") - self.trainer.train() + self.trainer.train( + resume_from_checkpoint = training_args.get("resume_from_checkpoint") + ) self._finalize_training(output_dir, "SNAC") return @@ -2913,7 +2921,9 @@ class UnslothTrainer: total_steps = total, status_message = "Starting Whisper training..." ) logger.info(f"Whisper training config: {config}\n") - self.trainer.train() + self.trainer.train( + resume_from_checkpoint = training_args.get("resume_from_checkpoint") + ) self._finalize_training(output_dir, "Whisper") return @@ -3408,7 +3418,9 @@ class UnslothTrainer: # ========== START TRAINING ========== self._update_progress(status_message = "Starting training...") logger.info("Starting training...\n") - self.trainer.train() + self.trainer.train( + resume_from_checkpoint = training_args.get("resume_from_checkpoint") + ) # ========== SAVE MODEL ========== self._finalize_training(output_dir) diff --git a/studio/backend/core/training/training.py b/studio/backend/core/training/training.py index f35c7e8ad3..5642faa189 100644 --- a/studio/backend/core/training/training.py +++ b/studio/backend/core/training/training.py @@ -29,6 +29,10 @@ from typing import Optional, Tuple, Any import matplotlib.pyplot as plt from utils.hardware import prepare_gpu_selection +from utils.native_path_leases import ( + native_path_secret_removed_for_child_start, + run_without_native_path_secret, +) logger = get_logger(__name__) @@ -185,6 +189,7 @@ class TrainingBackend: "wandb_project": kwargs.get("wandb_project", "unsloth-training"), "enable_tensorboard": kwargs.get("enable_tensorboard", False), "tensorboard_dir": kwargs.get("tensorboard_dir", "runs"), + "resume_from_checkpoint": kwargs.get("resume_from_checkpoint"), "trust_remote_code": kwargs.get("trust_remote_code", False), "gpu_ids": kwargs.get("gpu_ids"), } @@ -212,20 +217,22 @@ class TrainingBackend: from .worker import run_training_process - event_queue = _CTX.Queue() - stop_queue = _CTX.Queue() - - proc = _CTX.Process( - target = run_training_process, - kwargs = { - "event_queue": event_queue, - "stop_queue": stop_queue, - "config": config, - }, - daemon = True, - ) try: - proc.start() + with native_path_secret_removed_for_child_start(): + event_queue = _CTX.Queue() + stop_queue = _CTX.Queue() + + proc = _CTX.Process( + target = run_without_native_path_secret, + args = (run_training_process,), + kwargs = { + "event_queue": event_queue, + "stop_queue": stop_queue, + "config": config, + }, + daemon = True, + ) + proc.start() except Exception: logger.error("Failed to start training subprocess", exc_info = True) return False diff --git a/studio/backend/core/training/worker.py b/studio/backend/core/training/worker.py index 8ab2b5b2be..60b9e994ab 100644 --- a/studio/backend/core/training/worker.py +++ b/studio/backend/core/training/worker.py @@ -35,6 +35,15 @@ from utils.wheel_utils import ( ) +def _output_dir_from_resume_checkpoint( + resume_from_checkpoint: str | None, +) -> str | None: + if not resume_from_checkpoint: + return None + path = Path(resume_from_checkpoint) + return str(path.parent if path.name.startswith("checkpoint-") else path) + + _CAUSAL_CONV1D_RELEASE_TAG = "v1.6.1.post4" _CAUSAL_CONV1D_PACKAGE_VERSION = "1.6.1" _MAMBA_SSM_RELEASE_TAG = "v2.3.1" @@ -50,6 +59,8 @@ def _model_wants_causal_conv1d(model_name: str) -> bool: for key in ( "qwen3.5", "qwen3_5", + "qwen3.6", + "qwen3_6", "qwen3-next", "qwen3_next", "nemotron_h", @@ -755,7 +766,10 @@ def run_training_process( return # Generate output dir - output_dir = config.get("output_dir") + resume_from_checkpoint = config.get("resume_from_checkpoint") + output_dir = config.get("output_dir") or _output_dir_from_resume_checkpoint( + resume_from_checkpoint + ) if not output_dir: output_dir = f"{model_name.replace('/', '_')}_{int(time.time())}" output_dir = str(resolve_output_dir(output_dir)) @@ -803,6 +817,7 @@ def run_training_process( max_seq_length = config.get("max_seq_length", 2048), optim = config.get("optim", "adamw_8bit"), lr_scheduler_type = config.get("lr_scheduler_type", "linear"), + resume_from_checkpoint = resume_from_checkpoint, ) _tqdm_stop.set() @@ -819,10 +834,13 @@ def run_training_process( } ) else: + saved_output_dir = ( + None if trainer.should_stop and not trainer.save_on_stop else output_dir + ) event_queue.put( { "type": "complete", - "output_dir": output_dir, + "output_dir": saved_output_dir, "status_message": progress.status_message or "Training completed", "ts": time.time(), } @@ -1107,11 +1125,15 @@ def _run_embedding_training(event_queue: Any, stop_queue: Any, config: dict) -> ) return - output_dir = config.get("output_dir") + resume_from_checkpoint = config.get("resume_from_checkpoint") + output_dir = config.get("output_dir") or _output_dir_from_resume_checkpoint( + resume_from_checkpoint + ) if not output_dir: output_dir = str( resolve_output_dir(f"{model_name.replace('/', '_')}_{int(time.time())}") ) + output_dir = str(resolve_output_dir(output_dir)) num_epochs = config.get("num_epochs", 2) batch_size = config.get("batch_size", 256) @@ -1219,7 +1241,7 @@ def _run_embedding_training(event_queue: Any, stop_queue: Any, config: dict) -> callbacks = [_EmbeddingProgressCallback()], ) - trainer.train() + trainer.train(resume_from_checkpoint = resume_from_checkpoint) except Exception as e: event_queue.put( { @@ -1245,6 +1267,8 @@ def _run_embedding_training(event_queue: Any, stop_queue: Any, config: dict) -> _send_status(event_queue, "Saving model...") try: + if _should_stop and _save_on_stop: + trainer._save_checkpoint(trainer.model, trial = None) model.save_pretrained(output_dir) model.tokenizer.save_pretrained(output_dir) logger.info("Embedding model saved to %s", output_dir) diff --git a/studio/backend/loggers/config.py b/studio/backend/loggers/config.py index 4c0d8ade28..4a27f13d38 100644 --- a/studio/backend/loggers/config.py +++ b/studio/backend/loggers/config.py @@ -22,6 +22,8 @@ from typing import Optional import structlog +from loggers.handlers import filter_sensitive_data + class LogConfig: """Structured logging configuration for the application. @@ -58,6 +60,8 @@ class LogConfig: structlog.processors.TimeStamper(fmt = "iso"), # timestamp first structlog.processors.add_log_level, # level second structlog.contextvars.merge_contextvars, + structlog.processors.format_exc_info, + filter_sensitive_data, # Custom processor to flatten the extra field lambda logger, method_name, event_dict: { "timestamp": event_dict.get("timestamp"), diff --git a/studio/backend/loggers/handlers.py b/studio/backend/loggers/handlers.py index 3add92ea1e..ddd404cdf3 100644 --- a/studio/backend/loggers/handlers.py +++ b/studio/backend/loggers/handlers.py @@ -15,6 +15,7 @@ Key Components: - get_logger: Factory function for structured loggers """ +import re import time from typing import Callable @@ -22,7 +23,12 @@ import structlog from fastapi import Request, Response from starlette.middleware.base import BaseHTTPMiddleware +from utils.native_path_leases import redact_native_paths + logger = structlog.get_logger(__name__) +_NATIVE_PATH_LEASE_RE = re.compile( + r"(?i)(\b(?:native_path_lease|nativePathLease)[\"']?\s*[:=]\s*[\"']?)[A-Za-z0-9_-]+\.[A-Za-z0-9_-]+" +) class LoggingMiddleware(BaseHTTPMiddleware): @@ -75,6 +81,12 @@ def filter_sensitive_data(logger, method_name, event_dict): """Structlog processor to filter out base64 data from logs.""" def filter_value(value): + if isinstance(value, str): + try: + value = redact_native_paths(value) + except Exception: + pass + value = _NATIVE_PATH_LEASE_RE.sub(r"\1", value) if ( isinstance(value, str) and len(value) > 100 @@ -83,12 +95,22 @@ def filter_sensitive_data(logger, method_name, event_dict): # Likely base64 data, truncate it return value[:20] + "..." elif isinstance(value, dict): - return {k: filter_value(v) for k, v in value.items()} + return { + k: "" + if str(k).replace("_", "").lower() == "nativepathlease" + else filter_value(v) + for k, v in value.items() + } elif isinstance(value, list): return [filter_value(item) for item in value] return value - return {k: filter_value(v) for k, v in event_dict.items()} + return { + k: "" + if str(k).replace("_", "").lower() == "nativepathlease" + else filter_value(v) + for k, v in event_dict.items() + } def get_logger(name: str) -> structlog.BoundLogger: diff --git a/studio/backend/main.py b/studio/backend/main.py index 05adcaa2ea..0958094ff0 100644 --- a/studio/backend/main.py +++ b/studio/backend/main.py @@ -62,6 +62,7 @@ from routes import ( datasets_router, export_router, inference_router, + inference_studio_router, models_router, training_history_router, training_router, @@ -77,6 +78,7 @@ from utils.hardware import ( import utils.hardware.hardware as _hw_module from utils.cache_cleanup import clear_unsloth_compiled_cache +from utils.native_path_leases import native_path_leases_supported def get_unsloth_version() -> str: @@ -186,6 +188,8 @@ if _api_only: "tauri://localhost", # Linux/macOS Tauri webview "http://tauri.localhost", # Windows Tauri webview "http://localhost", # dev fallback + "http://localhost:5173", # Tauri dev/Vite + "http://127.0.0.1:5173", # Tauri dev/Vite fallback ] _cors_origin_regex = None else: @@ -207,6 +211,9 @@ app.include_router(auth_router, prefix = "/api/auth", tags = ["auth"]) app.include_router(training_router, prefix = "/api/train", tags = ["training"]) app.include_router(models_router, prefix = "/api/models", tags = ["models"]) app.include_router(inference_router, prefix = "/api/inference", tags = ["inference"]) +# Studio-only inference endpoints (cancel, etc.) are intentionally NOT +# exposed on the /v1 OpenAI-compat prefix below. +app.include_router(inference_studio_router, prefix = "/api/inference", tags = ["inference"]) # OpenAI-compatible endpoints: mount the same inference router at /v1 # so external tools (Open WebUI, SillyTavern, etc.) can use the @@ -238,6 +245,7 @@ async def health_check(): "chat_only": _hw_module.CHAT_ONLY, "desktop_protocol_version": 1, "supports_desktop_auth": True, + "native_path_leases_supported": native_path_leases_supported(), } diff --git a/studio/backend/models/inference.py b/studio/backend/models/inference.py index e5b037755d..43087cc5bf 100644 --- a/studio/backend/models/inference.py +++ b/studio/backend/models/inference.py @@ -18,6 +18,9 @@ class LoadRequest(BaseModel): """Request to load a model for inference""" model_path: str = Field(..., description = "Model identifier or local path") + native_path_lease: Optional[str] = Field( + None, description = "Frontend-visible signed native path grant" + ) hf_token: Optional[str] = Field( None, description = "HuggingFace token for gated models" ) @@ -52,6 +55,16 @@ class LoadRequest(BaseModel): None, description = "Speculative decoding mode for GGUF models (e.g. 'ngram-simple', 'ngram-mod'). Ignored for non-GGUF and vision models.", ) + llama_extra_args: Optional[List[str]] = Field( + None, + description = ( + "Extra arguments forwarded verbatim to llama-server for GGUF models. " + "One token per list entry, e.g. ['--top-k', '20', '--seed', '42']. " + "Studio-managed flags (model identity, port, context length, GPU placement, " + "auth, --flash-attn, --no-context-shift, --jinja) are rejected. Ignored for " + "non-GGUF models." + ), + ) class UnloadRequest(BaseModel): @@ -69,6 +82,9 @@ class ValidateModelRequest(BaseModel): """ model_path: str = Field(..., description = "Model identifier or local path") + native_path_lease: Optional[str] = Field( + None, description = "Frontend-visible signed native path grant" + ) hf_token: Optional[str] = Field( None, description = "HuggingFace token for gated models" ) @@ -283,6 +299,10 @@ class InferenceStatusResponse(BaseModel): supports_tools: bool = Field( False, description = "Whether the active model supports tool calling" ) + chat_template: Optional[str] = Field( + None, + description = "Jinja2 chat template string for the active model", + ) context_length: Optional[int] = Field( None, description = "Context length of the active model" ) @@ -396,7 +416,10 @@ class ChatMessage(BaseModel): if self.name is not None and self.role != "tool": raise ValueError('"name" is only valid on role="tool" messages.') - # Per-role content requirements. + # Per-role content requirements. OpenAI-compatible clients may send + # ``content=""`` for image-only turns when the image travels in a + # companion field such as Studio's ``image_base64`` extension, so treat + # empty strings as present content for user/system messages. if self.role == "tool": if not self.tool_call_id: raise ValueError( @@ -411,10 +434,8 @@ class ChatMessage(BaseModel): 'role="assistant" messages require either "content" or "tool_calls".' ) else: # "user" | "system" - if not self.content: - raise ValueError( - f'role="{self.role}" messages require non-empty "content".' - ) + if self.content is None or self.content == []: + raise ValueError(f'role="{self.role}" messages require "content".') return self @@ -531,6 +552,10 @@ class ChatCompletionRequest(BaseModel): None, description = "[x-unsloth] Session/thread ID for scoping tool execution sandbox.", ) + cancel_id: Optional[str] = Field( + None, + description = "[x-unsloth] Per-request cancellation token. Frontend sends a fresh UUID per run so /inference/cancel matches one specific generation.", + ) # ── Streaming response chunks ──────────────────────────────────── @@ -992,6 +1017,7 @@ class AnthropicMessagesRequest(BaseModel): enable_tools: Optional[bool] = None enabled_tools: Optional[list[str]] = None session_id: Optional[str] = None + cancel_id: Optional[str] = None model_config = {"extra": "allow"} diff --git a/studio/backend/models/training.py b/studio/backend/models/training.py index 07a306ca39..a9f4caa1bb 100644 --- a/studio/backend/models/training.py +++ b/studio/backend/models/training.py @@ -127,6 +127,9 @@ class TrainingStartRequest(BaseModel): wandb_project: Optional[str] = Field(None, description = "W&B project name") enable_tensorboard: bool = Field(False, description = "Enable TensorBoard logging") tensorboard_dir: Optional[str] = Field(None, description = "TensorBoard directory") + resume_from_checkpoint: Optional[str] = Field( + None, description = "Saved training output directory to resume from" + ) # GPU selection gpu_ids: Optional[List[int]] = Field( @@ -220,6 +223,8 @@ class TrainingRunSummary(BaseModel): duration_seconds: Optional[float] = None error_message: Optional[str] = None loss_sparkline: Optional[List[float]] = None + can_resume: bool = False + resumed_later: bool = False class TrainingRunListResponse(BaseModel): diff --git a/studio/backend/plugins/data-designer-github-repo-seed/README.md b/studio/backend/plugins/data-designer-github-repo-seed/README.md new file mode 100644 index 0000000000..346d94b305 --- /dev/null +++ b/studio/backend/plugins/data-designer-github-repo-seed/README.md @@ -0,0 +1,73 @@ +# data-designer-github-repo-seed + +A Data Designer seed-reader plugin for **Unsloth Studio** that scrapes real +GitHub data (issues, pull requests, commits) from one or more repositories +and hands it to the recipe pipeline as a seed dataset. + +Designed to ship with Studio as a default seed source so any user with a +GitHub token can build training datasets straight from live repos. + +## What it does + +Given a list of `owner/name` repos, a GitHub token, and a per-resource +`limit`, the plugin uses GitHub's GraphQL API to fetch issues, pull +requests, and/or commits, with labels, state, authors, and the first N +comments of each item, and materialises a single JSONL with uniform +columns so the rest of the recipe (LLM text / LLM structured / processors) +can treat it like any other seed table. + +| Column | Description | +|---------------|------------------------------------------------| +| `item_type` | `issue` / `pull` / `commit` | +| `repo` | `owner/name` | +| `number` | Issue/PR number, or commit SHA | +| `title` | Title (or commit message headline) | +| `body` | Issue/PR body (or full commit message) | +| `state` | `OPEN` / `CLOSED` / `MERGED` (empty for commit)| +| `author` | GitHub login of the author | +| `created_at` | ISO8601 | +| `closed_at` | ISO8601 (empty for commits) | +| `url` | Permalink | +| `labels` | List of label names | +| `comments` | First N comments concatenated | + +## Usage in a recipe + +```json +{ + "seed_config": { + "source": { + "seed_type": "github_repo", + "repos": ["unslothai/unsloth", "unslothai/unsloth-zoo"], + "token": "", + "item_types": ["issues", "pulls"], + "limit": 100, + "include_comments": true, + "max_comments_per_item": 30 + }, + "sampling_strategy": "shuffle", + "selection_strategy": null + } +} +``` + +Leave `token` empty to fall back to the server's `GH_TOKEN` / `GITHUB_TOKEN` +environment variable, useful when the recipe is published and shouldn't +carry a secret. + +## Auth + +A GitHub personal access token with `public_repo` scope is enough for public +repositories; `repo` scope is required for private ones. GraphQL requests +are rate-limit aware: the client inspects `x-ratelimit-*` headers and +sleeps until reset when the budget drops below a safety threshold. + +## Install + +Shipped as a default Studio plugin. For development: + +```bash +pip install -e . +``` + +Registered automatically via the `data_designer.plugins` entry point. diff --git a/studio/backend/plugins/data-designer-github-repo-seed/pyproject.toml b/studio/backend/plugins/data-designer-github-repo-seed/pyproject.toml new file mode 100644 index 0000000000..e232adc60c --- /dev/null +++ b/studio/backend/plugins/data-designer-github-repo-seed/pyproject.toml @@ -0,0 +1,25 @@ +# SPDX-License-Identifier: AGPL-3.0-only +# Copyright 2026-present the Unsloth AI Inc. team. All rights reserved. See /studio/LICENSE.AGPL-3.0 + +[build-system] +requires = ["setuptools>=68", "wheel"] +build-backend = "setuptools.build_meta" + +[project] +name = "data-designer-github-repo-seed" +version = "0.1.0" +description = "Unsloth Studio seed plugin that scrapes GitHub issues, PRs, and commits." +requires-python = ">=3.11" +dependencies = [ + "data-designer-engine>=0.5.4,<0.6", + "requests>=2.31", +] + +[project.entry-points."data_designer.plugins"] +github_repo_seed = "data_designer_github_repo_seed.plugin:github_repo_seed_plugin" + +[tool.setuptools] +package-dir = {"" = "src"} + +[tool.setuptools.packages.find] +where = ["src"] diff --git a/studio/backend/plugins/data-designer-github-repo-seed/src/data_designer_github_repo_seed/__init__.py b/studio/backend/plugins/data-designer-github-repo-seed/src/data_designer_github_repo_seed/__init__.py new file mode 100644 index 0000000000..f57af4c6c3 --- /dev/null +++ b/studio/backend/plugins/data-designer-github-repo-seed/src/data_designer_github_repo_seed/__init__.py @@ -0,0 +1,7 @@ +# SPDX-License-Identifier: AGPL-3.0-only +# Copyright 2026-present the Unsloth AI Inc. team. All rights reserved. See /studio/LICENSE.AGPL-3.0 + +# Intentionally empty. Data-designer loads submodules lazily via qualified names +# (impl_qualified_name / config_qualified_name in plugin.py), so importing this +# package must NOT touch modules that depend on data_designer.engine.* during +# Studio's bootstrap (circular import). diff --git a/studio/backend/plugins/data-designer-github-repo-seed/src/data_designer_github_repo_seed/config.py b/studio/backend/plugins/data-designer-github-repo-seed/src/data_designer_github_repo_seed/config.py new file mode 100644 index 0000000000..6b347c4f83 --- /dev/null +++ b/studio/backend/plugins/data-designer-github-repo-seed/src/data_designer_github_repo_seed/config.py @@ -0,0 +1,64 @@ +# SPDX-License-Identifier: AGPL-3.0-only +# Copyright 2026-present the Unsloth AI Inc. team. All rights reserved. See /studio/LICENSE.AGPL-3.0 + +from __future__ import annotations + +from typing import Literal + +from pydantic import Field, field_validator, model_validator + +from data_designer.config.seed_source import SeedSource + + +class GitHubRepoSeedSource(SeedSource): + seed_type: Literal["github_repo"] = "github_repo" + + repos: list[str] = Field( + default_factory = list, + description = "List of GitHub repositories to scrape, each in `owner/name` form.", + ) + token: str = Field( + default = "", + description = "Personal access token. Leave blank to read GH_TOKEN / GITHUB_TOKEN from env at run time.", + ) + item_types: list[Literal["issues", "pulls", "commits"]] = Field( + default = ["issues", "pulls"], + description = "Which GitHub item types to fetch per repo.", + ) + limit: int = Field( + default = 100, + ge = 1, + le = 5000, + description = "Maximum items per repo per item type (e.g. limit=100 + ['issues','pulls'] => up to 200 items per repo).", + ) + include_comments: bool = Field( + default = True, + description = "Fetch the first N comments of each issue/PR and include them in the `comments` column.", + ) + max_comments_per_item: int = Field(default = 30, ge = 0, le = 200) + + @field_validator("repos") + @classmethod + def _validate_repos(cls, v: list[str]) -> list[str]: + out: list[str] = [] + for r in v or []: + r = r.strip() + if not r: + continue + if r.count("/") != 1 or not all(r.split("/")): + raise ValueError(f"Each repo must be `owner/name`; got {r!r}") + out.append(r) + return out + + @field_validator("item_types") + @classmethod + def _validate_item_types(cls, v: list[str]) -> list[str]: + if not v: + raise ValueError("item_types must not be empty") + return list(dict.fromkeys(v)) + + @model_validator(mode = "after") + def _ensure_repos(self) -> "GitHubRepoSeedSource": + if not self.repos: + raise ValueError("At least one repo is required") + return self diff --git a/studio/backend/plugins/data-designer-github-repo-seed/src/data_designer_github_repo_seed/impl.py b/studio/backend/plugins/data-designer-github-repo-seed/src/data_designer_github_repo_seed/impl.py new file mode 100644 index 0000000000..5a38e26d6b --- /dev/null +++ b/studio/backend/plugins/data-designer-github-repo-seed/src/data_designer_github_repo_seed/impl.py @@ -0,0 +1,83 @@ +# SPDX-License-Identifier: AGPL-3.0-only +# Copyright 2026-present the Unsloth AI Inc. team. All rights reserved. See /studio/LICENSE.AGPL-3.0 + +from __future__ import annotations + +import hashlib +import tempfile +import threading +from pathlib import Path +from typing import Optional + +import data_designer.lazy_heavy_imports as lazy +from data_designer.engine.resources.seed_reader import SeedReader + +from .config import GitHubRepoSeedSource +from .scraper import ScrapeConfig, materialize_to_jsonl + + +# In-process cache mapping a stable config signature to the JSONL materialization +# path. A single recipe job invokes the seed reader multiple times (validation, +# preview, per-column sampling), and the default flow re-scrapes the repo on +# every call: for a 2-repo preview that is ~15s of redundant GitHub GraphQL +# traffic before any generation fires. Memoize the materialization so the second +# and third passes reuse the file the first pass wrote. Cache key excludes the +# raw token and uses a short SHA-256 digest so token values never hit memory +# twice and token rotation invalidates cleanly. +_SCRAPE_CACHE: dict[tuple, str] = {} +_SCRAPE_CACHE_LOCK = threading.Lock() + + +def _scrape_cache_key(cfg: ScrapeConfig) -> tuple: + token_digest = hashlib.sha256( + (cfg.token or "").encode("utf-8"), + ).hexdigest()[:16] + return ( + tuple(cfg.repos), + tuple(cfg.item_types), + cfg.limit, + bool(cfg.include_comments), + cfg.max_comments_per_item, + token_digest, + ) + + +def _lookup_cached_scrape(key: tuple) -> Optional[str]: + with _SCRAPE_CACHE_LOCK: + path = _SCRAPE_CACHE.get(key) + if path and Path(path).exists(): + return path + # Stale entry (tmp cleanup, user restarted, ...); drop it so the caller + # materializes a fresh file rather than returning a dangling path. + if path: + with _SCRAPE_CACHE_LOCK: + _SCRAPE_CACHE.pop(key, None) + return None + + +def _store_cached_scrape(key: tuple, path: str) -> None: + with _SCRAPE_CACHE_LOCK: + _SCRAPE_CACHE[key] = path + + +class GitHubRepoSeedReader(SeedReader[GitHubRepoSeedSource]): + def create_duckdb_connection(self): + return lazy.duckdb.connect() + + def get_dataset_uri(self) -> str: + out_dir = Path(tempfile.gettempdir()) / "studio-github-repo-seed" + cfg = ScrapeConfig( + repos = list(self.source.repos), + token = self.source.token, + item_types = list(self.source.item_types), + limit = self.source.limit, + include_comments = self.source.include_comments, + max_comments_per_item = self.source.max_comments_per_item, + ) + cache_key = _scrape_cache_key(cfg) + cached_path = _lookup_cached_scrape(cache_key) + if cached_path is not None: + return cached_path + path = materialize_to_jsonl(cfg, out_dir) + _store_cached_scrape(cache_key, str(path)) + return str(path) diff --git a/studio/backend/plugins/data-designer-github-repo-seed/src/data_designer_github_repo_seed/plugin.py b/studio/backend/plugins/data-designer-github-repo-seed/src/data_designer_github_repo_seed/plugin.py new file mode 100644 index 0000000000..f87dbd0507 --- /dev/null +++ b/studio/backend/plugins/data-designer-github-repo-seed/src/data_designer_github_repo_seed/plugin.py @@ -0,0 +1,10 @@ +# SPDX-License-Identifier: AGPL-3.0-only +# Copyright 2026-present the Unsloth AI Inc. team. All rights reserved. See /studio/LICENSE.AGPL-3.0 + +from data_designer.plugins.plugin import Plugin, PluginType + +github_repo_seed_plugin = Plugin( + impl_qualified_name = "data_designer_github_repo_seed.impl.GitHubRepoSeedReader", + config_qualified_name = "data_designer_github_repo_seed.config.GitHubRepoSeedSource", + plugin_type = PluginType.SEED_READER, +) diff --git a/studio/backend/plugins/data-designer-github-repo-seed/src/data_designer_github_repo_seed/scraper.py b/studio/backend/plugins/data-designer-github-repo-seed/src/data_designer_github_repo_seed/scraper.py new file mode 100644 index 0000000000..d768fe37be --- /dev/null +++ b/studio/backend/plugins/data-designer-github-repo-seed/src/data_designer_github_repo_seed/scraper.py @@ -0,0 +1,236 @@ +# SPDX-License-Identifier: AGPL-3.0-only +# Copyright 2026-present the Unsloth AI Inc. team. All rights reserved. See /studio/LICENSE.AGPL-3.0 + +"""Multi-repo GitHub scraper for the Studio seed plugin. + +Drives the GraphQL-based scraper in `scraper_impl/` per repo. Each repo is +scraped with a trial_limits cap so we stop at `limit` items per resource. +After scraping, we read the per-resource JSONL shards and flatten them into +a single unified JSONL with stable columns (`item_type`, `repo`, `number`, +`title`, `body`, ...). +""" + +from __future__ import annotations + +import json +import os +import sys +import time +import uuid +from dataclasses import dataclass +from pathlib import Path + +# Defer scraper_impl imports until `scrape()` runs with a resolved token. +_IMPL_DIR = Path(__file__).parent / "scraper_impl" + + +def _ensure_impl_on_path() -> None: + if str(_IMPL_DIR) not in sys.path: + sys.path.insert(0, str(_IMPL_DIR)) + + +def _load_impl(): + _ensure_impl_on_path() + import importlib + + gh_client = importlib.import_module("gh_client") # type: ignore + scraper_mod = importlib.import_module("scraper") # type: ignore + return gh_client.GitHubClient, scraper_mod.RepoScraper + + +@dataclass +class ScrapeConfig: + repos: list[str] + token: str + item_types: list[str] + limit: int + include_comments: bool + max_comments_per_item: int + + +def _resolve_token(token: str) -> str: + tok = token or os.environ.get("GH_TOKEN", "") or os.environ.get("GITHUB_TOKEN", "") + if not tok: + raise ValueError( + "GitHub token is required. Set it in the recipe config or the GH_TOKEN / GITHUB_TOKEN env var." + ) + return tok + + +def _read_jsonl(path: Path, max_rows: int | None = None): + if not path.exists(): + return + with path.open(encoding = "utf-8") as f: + for i, line in enumerate(f): + if not line.strip(): + continue + if max_rows is not None and i >= max_rows: + return + try: + yield json.loads(line) + except json.JSONDecodeError: + continue + + +def _flatten_issue_row(r: dict, repo: str, include_comments: bool, max_c: int) -> dict: + labels = [ + l.get("name") + for l in (r.get("labels", {}) or {}).get("nodes", []) + if l.get("name") + ] + comments_nodes = (r.get("comments") or {}).get("nodes") or [] + comments_text = "" + if include_comments and comments_nodes: + kept = comments_nodes[:max_c] + comments_text = "\n\n".join( + f"[{(c.get('author') or {}).get('login', '?')}]: {c.get('body') or ''}" + for c in kept + ) + return { + "item_type": "issue", + "repo": repo, + "number": r.get("number"), + "title": r.get("title") or "", + "body": r.get("body") or "", + "state": r.get("state") or "", + "author": (r.get("author") or {}).get("login", ""), + "created_at": r.get("createdAt") or "", + "closed_at": r.get("closedAt") or "", + "url": r.get("url") or r.get("permalink") or "", + "labels": labels, + "comments": comments_text, + } + + +def _flatten_pr_row(r: dict, repo: str, include_comments: bool, max_c: int) -> dict: + labels = [ + l.get("name") + for l in (r.get("labels", {}) or {}).get("nodes", []) + if l.get("name") + ] + comments_nodes = (r.get("comments") or {}).get("nodes") or [] + comments_text = "" + if include_comments and comments_nodes: + kept = comments_nodes[:max_c] + comments_text = "\n\n".join( + f"[{(c.get('author') or {}).get('login', '?')}]: {c.get('body') or ''}" + for c in kept + ) + return { + "item_type": "pull", + "repo": repo, + "number": r.get("number"), + "title": r.get("title") or "", + "body": r.get("body") or "", + "state": r.get("state") or "", + "author": (r.get("author") or {}).get("login", ""), + "created_at": r.get("createdAt") or "", + "closed_at": r.get("closedAt") or "", + "url": r.get("url") or r.get("permalink") or "", + "labels": labels, + "comments": comments_text, + } + + +def _flatten_commit_row(r: dict, repo: str) -> dict: + msg = r.get("messageHeadline") or r.get("message") or "" + body = r.get("messageBody") or r.get("message") or msg + author = r.get("author") or {} + return { + "item_type": "commit", + "repo": repo, + "number": r.get("oid") or r.get("sha") or "", + "title": msg, + "body": body, + "state": "", + "author": (author.get("user") or {}).get("login") or author.get("name", ""), + "created_at": (author.get("date") or r.get("committedDate") or ""), + "closed_at": "", + "url": r.get("url") or "", + "labels": [], + "comments": "", + } + + +def scrape(cfg: ScrapeConfig, base_dir: Path): + token = _resolve_token(cfg.token) + GitHubClient, RepoScraper = _load_impl() + client = GitHubClient(token = token) + base_dir.mkdir(parents = True, exist_ok = True) + + # Per-resource trial limits. limit <= 0 means "all": use a very large cap. + effective_limit = cfg.limit if cfg.limit and cfg.limit > 0 else 1_000_000 + trial_limits: dict[str, int] = {} + if "issues" in cfg.item_types: + trial_limits["issues"] = effective_limit + if "pulls" in cfg.item_types: + trial_limits["pull_requests"] = effective_limit + if "commits" in cfg.item_types: + trial_limits["commits"] = effective_limit + + all_rows: list[dict] = [] + for repo in cfg.repos: + owner, name = repo.split("/", 1) + scraper = RepoScraper( + owner = owner, + name = name, + base_dir = base_dir, + client = client, + trial_limits = trial_limits, + light = True, + ) + try: + repo_meta = scraper.scrape_repo_meta() + if "issues" in cfg.item_types: + scraper.scrape_issues() + if "pulls" in cfg.item_types: + scraper.scrape_prs() + if "commits" in cfg.item_types: + default_ref = repo_meta.get("defaultBranchRef") or {} + default_branch = ( + default_ref.get("name") if isinstance(default_ref, dict) else None + ) + branch = ( + f"refs/heads/{default_branch}" + if default_branch + else "refs/heads/main" + ) + scraper.scrape_commits(branch = branch) + finally: + scraper.close() + + read_cap = cfg.limit if cfg.limit and cfg.limit > 0 else None + repo_dir = base_dir / f"{owner}__{name}" + if "issues" in cfg.item_types: + for row in _read_jsonl(repo_dir / "issues.jsonl", read_cap): + all_rows.append( + _flatten_issue_row( + row, repo, cfg.include_comments, cfg.max_comments_per_item + ) + ) + if "pulls" in cfg.item_types: + for row in _read_jsonl(repo_dir / "pull_requests.jsonl", read_cap): + all_rows.append( + _flatten_pr_row( + row, repo, cfg.include_comments, cfg.max_comments_per_item + ) + ) + if "commits" in cfg.item_types: + for row in _read_jsonl(repo_dir / "commits.jsonl", read_cap): + all_rows.append(_flatten_commit_row(row, repo)) + + return all_rows + + +def materialize_to_jsonl(cfg: ScrapeConfig, out_dir: Path) -> Path: + out_dir.mkdir(parents = True, exist_ok = True) + tag = "-".join(r.replace("/", "__") for r in cfg.repos)[:120] + kinds = "-".join(cfg.item_types) + run_id = f"{int(time.time())}-{uuid.uuid4().hex[:12]}" + fname = f"github_{tag}__{kinds}__{cfg.limit}_{run_id}.jsonl" + out = out_dir / fname + rows = scrape(cfg, out_dir / "raw-runs" / run_id) + with out.open("w", encoding = "utf-8") as f: + for r in rows: + f.write(json.dumps(r, ensure_ascii = False) + "\n") + return out diff --git a/studio/backend/plugins/data-designer-github-repo-seed/src/data_designer_github_repo_seed/scraper_impl/__init__.py b/studio/backend/plugins/data-designer-github-repo-seed/src/data_designer_github_repo_seed/scraper_impl/__init__.py new file mode 100644 index 0000000000..32014236c6 --- /dev/null +++ b/studio/backend/plugins/data-designer-github-repo-seed/src/data_designer_github_repo_seed/scraper_impl/__init__.py @@ -0,0 +1,2 @@ +# SPDX-License-Identifier: AGPL-3.0-only +# Copyright 2026-present the Unsloth AI Inc. team. All rights reserved. See /studio/LICENSE.AGPL-3.0 diff --git a/studio/backend/plugins/data-designer-github-repo-seed/src/data_designer_github_repo_seed/scraper_impl/gh_client.py b/studio/backend/plugins/data-designer-github-repo-seed/src/data_designer_github_repo_seed/scraper_impl/gh_client.py new file mode 100644 index 0000000000..dd2de2f5ce --- /dev/null +++ b/studio/backend/plugins/data-designer-github-repo-seed/src/data_designer_github_repo_seed/scraper_impl/gh_client.py @@ -0,0 +1,248 @@ +# SPDX-License-Identifier: AGPL-3.0-only +# Copyright 2026-present the Unsloth AI Inc. team. All rights reserved. See /studio/LICENSE.AGPL-3.0 + +"""GitHub API client with rate-limit awareness, retry, and dual REST/GraphQL support.""" + +from __future__ import annotations + +import json +import os +import time +import logging +from typing import Any, Dict, Iterable, Iterator, List, Optional + +import requests + +log = logging.getLogger("gh_client") + +GRAPHQL_URL = "https://api.github.com/graphql" +REST_BASE = "https://api.github.com" + +BASE_HEADERS = { + "Accept": "application/vnd.github+json", + "X-GitHub-Api-Version": "2022-11-28", + "User-Agent": "github-data-gatherer/1.0", +} + + +class RateLimitError(Exception): + pass + + +class GitHubClient: + def __init__( + self, + min_remaining_graphql: int = 100, + min_remaining_rest: int = 100, + token: str | None = None, + ): + token = token or os.environ.get("GH_TOKEN") or os.environ.get("GITHUB_TOKEN") + if not token: + raise RuntimeError("GH_TOKEN not set in environment") + self.session = requests.Session() + self.session.headers.update( + {**BASE_HEADERS, "Authorization": f"Bearer {token}"} + ) + self.min_remaining_graphql = min_remaining_graphql + self.min_remaining_rest = min_remaining_rest + self.graphql_remaining: Optional[int] = None + self.graphql_reset: Optional[int] = None + self.rest_remaining: Optional[int] = None + self.rest_reset: Optional[int] = None + self.calls_graphql = 0 + self.calls_rest = 0 + self.retry_count = 0 + + def _sleep_until(self, reset_ts: int, buffer_s: int = 10) -> None: + now = int(time.time()) + wait = max(0, reset_ts - now) + buffer_s + log.warning("Rate limit hit. Sleeping %ds until reset.", wait) + time.sleep(wait) + + def _check_rate_and_wait(self, kind: str) -> None: + if kind == "graphql": + remaining = self.graphql_remaining + reset = self.graphql_reset + min_remaining = self.min_remaining_graphql + else: + remaining = self.rest_remaining + reset = self.rest_reset + min_remaining = self.min_remaining_rest + if remaining is not None and remaining < min_remaining: + if reset: + self._sleep_until(reset) + # Reset remaining so we don't spin + if kind == "graphql": + self.graphql_remaining = None + else: + self.rest_remaining = None + + def graphql( + self, + query: str, + variables: Optional[Dict[str, Any]] = None, + max_retries: int = 20, + ) -> Dict[str, Any]: + self._check_rate_and_wait("graphql") + backoff = 2 + last_err = None + for attempt in range(max_retries): + try: + r = self.session.post( + GRAPHQL_URL, + json = {"query": query, "variables": variables or {}}, + timeout = 120, + ) + self.calls_graphql += 1 + # Update rate info from response headers + rem = r.headers.get("X-RateLimit-Remaining") + rst = r.headers.get("X-RateLimit-Reset") + if rem is not None: + try: + self.graphql_remaining = int(rem) + except ValueError: + pass + if rst is not None: + try: + self.graphql_reset = int(rst) + except ValueError: + pass + if r.status_code in (502, 503, 504): + log.warning("GraphQL %s transient, retrying", r.status_code) + time.sleep(backoff) + backoff = min(backoff * 2, 60) + continue + if r.status_code == 403 or r.status_code == 429: + # Check for secondary/abuse + retry_after = r.headers.get("Retry-After") + if retry_after: + t = int(retry_after) + log.warning("Secondary rate limit. Sleep %ds.", t) + time.sleep(t + 2) + continue + if self.graphql_reset: + self._sleep_until(self.graphql_reset) + continue + time.sleep(60) + continue + r.raise_for_status() + data = r.json() + if "errors" in data and data["errors"]: + # Surface errors but allow partial data + errs = data["errors"] + # Retry on RATE_LIMITED + for e in errs: + if e.get("type") == "RATE_LIMITED": + self._sleep_until( + (self.graphql_reset or int(time.time()) + 60) + ) + break + else: + # No rate-limit error, log and return partial + log.warning("GraphQL errors: %s", json.dumps(errs)[:400]) + return data + continue + return data + except requests.RequestException as e: + last_err = e + log.warning("GraphQL network error: %s. Retry.", e) + time.sleep(backoff) + backoff = min(backoff * 2, 60) + raise RuntimeError(f"GraphQL failed after {max_retries} retries: {last_err}") + + def rest( + self, + method: str, + path: str, + params: Optional[Dict[str, Any]] = None, + json_body: Optional[Dict[str, Any]] = None, + max_retries: int = 6, + ) -> requests.Response: + self._check_rate_and_wait("rest") + if path.startswith("http"): + url = path + else: + url = REST_BASE + path + backoff = 2 + last_err = None + for attempt in range(max_retries): + try: + r = self.session.request( + method, url, params = params, json = json_body, timeout = 120 + ) + self.calls_rest += 1 + rem = r.headers.get("X-RateLimit-Remaining") + rst = r.headers.get("X-RateLimit-Reset") + if rem is not None: + try: + self.rest_remaining = int(rem) + except ValueError: + pass + if rst is not None: + try: + self.rest_reset = int(rst) + except ValueError: + pass + if r.status_code in (502, 503, 504): + log.warning("REST %s transient, retrying", r.status_code) + time.sleep(backoff) + backoff = min(backoff * 2, 60) + continue + if r.status_code in (403, 429): + retry_after = r.headers.get("Retry-After") + if retry_after: + t = int(retry_after) + log.warning("Secondary rate limit on REST. Sleep %ds.", t) + time.sleep(t + 2) + continue + # Check if primary rate + if self.rest_remaining == 0 and self.rest_reset: + self._sleep_until(self.rest_reset) + continue + log.warning("REST 403/429, sleep 60") + time.sleep(60) + continue + return r + except requests.RequestException as e: + last_err = e + log.warning("REST network error: %s. Retry.", e) + time.sleep(backoff) + backoff = min(backoff * 2, 60) + raise RuntimeError(f"REST failed after {max_retries} retries: {last_err}") + + def rest_paginate( + self, path: str, params: Optional[Dict[str, Any]] = None, per_page: int = 100 + ) -> Iterator[dict]: + params = dict(params or {}) + params.setdefault("per_page", per_page) + url = path + while True: + r = self.rest("GET", url, params = params if url == path else None) + if r.status_code != 200: + log.error( + "REST paginate got %s at %s: %s", r.status_code, url, r.text[:200] + ) + return + items = r.json() + if isinstance(items, dict): + # Some endpoints return dict with list field + items = items.get("items", []) + for it in items: + yield it + # Follow link header + link = r.headers.get("Link", "") + nxt = None + for part in link.split(","): + if 'rel="next"' in part: + nxt = part.split(";")[0].strip().strip("<>") + break + if not nxt: + return + url = nxt + params = None + + def rate_snapshot(self) -> Dict[str, Any]: + r = self.rest("GET", "/rate_limit") + if r.status_code == 200: + return r.json() + return {} diff --git a/studio/backend/plugins/data-designer-github-repo-seed/src/data_designer_github_repo_seed/scraper_impl/queries.py b/studio/backend/plugins/data-designer-github-repo-seed/src/data_designer_github_repo_seed/scraper_impl/queries.py new file mode 100644 index 0000000000..9dc7613db5 --- /dev/null +++ b/studio/backend/plugins/data-designer-github-repo-seed/src/data_designer_github_repo_seed/scraper_impl/queries.py @@ -0,0 +1,685 @@ +# SPDX-License-Identifier: AGPL-3.0-only +# Copyright 2026-present the Unsloth AI Inc. team. All rights reserved. See /studio/LICENSE.AGPL-3.0 + +"""GraphQL queries for GitHub data scraping. + +GitHub's GraphQL rejects queries that define unused fragments, so each query +only includes the fragments it actually references. +""" + +# ---- Fragments (kept as raw strings, composed per query) ---- +F_ACTOR = """ +fragment ActorFields on Actor { + __typename + login + url + avatarUrl + ... on User { id databaseId name } + ... on Bot { id databaseId } + ... on Organization { id databaseId name } +} +""" + +F_LABEL = """ +fragment LabelFields on Label { + id + name + color + description + createdAt +} +""" + +F_TIMELINE = """ +fragment TimelineItem on IssueTimelineItems { + __typename + ... on Node { id } + ... on AddedToProjectEvent { createdAt actor { ...ActorFields } } + ... on AssignedEvent { createdAt actor { ...ActorFields } assignee { __typename ... on User { login } ... on Bot { login } } } + ... on ClosedEvent { createdAt actor { ...ActorFields } stateReason closer { __typename ... on Commit { oid url } ... on PullRequest { number url } } } + ... on CommentDeletedEvent { createdAt actor { ...ActorFields } } + ... on ConnectedEvent { createdAt actor { ...ActorFields } source { __typename ... on Issue { number url repository { nameWithOwner } } ... on PullRequest { number url repository { nameWithOwner } } } subject { __typename ... on Issue { number url } ... on PullRequest { number url } } } + ... on ConvertedNoteToIssueEvent { createdAt actor { ...ActorFields } } + ... on CrossReferencedEvent { createdAt actor { ...ActorFields } isCrossRepository willCloseTarget source { __typename ... on Issue { number url repository { nameWithOwner } title } ... on PullRequest { number url repository { nameWithOwner } title } } } + ... on DemilestonedEvent { createdAt actor { ...ActorFields } milestoneTitle } + ... on DisconnectedEvent { createdAt actor { ...ActorFields } subject { __typename ... on Issue { number url } ... on PullRequest { number url } } source { __typename ... on Issue { number url } ... on PullRequest { number url } } } + ... on IssueComment { id databaseId createdAt updatedAt author { ...ActorFields } body url reactionGroups { content reactors { totalCount } } } + ... on LabeledEvent { createdAt actor { ...ActorFields } label { name color } } + ... on LockedEvent { createdAt actor { ...ActorFields } lockReason } + ... on MarkedAsDuplicateEvent { createdAt actor { ...ActorFields } canonical { __typename ... on Issue { number url } ... on PullRequest { number url } } } + ... on MentionedEvent { createdAt actor { ...ActorFields } } + ... on MilestonedEvent { createdAt actor { ...ActorFields } milestoneTitle } + ... on MovedColumnsInProjectEvent { createdAt actor { ...ActorFields } } + ... on PinnedEvent { createdAt actor { ...ActorFields } } + ... on ReferencedEvent { createdAt actor { ...ActorFields } commit { oid url } commitRepository { nameWithOwner } } + ... on RemovedFromProjectEvent { createdAt actor { ...ActorFields } } + ... on RenamedTitleEvent { createdAt actor { ...ActorFields } previousTitle currentTitle } + ... on ReopenedEvent { createdAt actor { ...ActorFields } } + ... on SubscribedEvent { createdAt actor { ...ActorFields } } + ... on TransferredEvent { createdAt actor { ...ActorFields } fromRepository { nameWithOwner } } + ... on UnassignedEvent { createdAt actor { ...ActorFields } assignee { __typename ... on User { login } ... on Bot { login } } } + ... on UnlabeledEvent { createdAt actor { ...ActorFields } label { name color } } + ... on UnlockedEvent { createdAt actor { ...ActorFields } } + ... on UnmarkedAsDuplicateEvent { createdAt actor { ...ActorFields } } + ... on UnpinnedEvent { createdAt actor { ...ActorFields } } + ... on UnsubscribedEvent { createdAt actor { ...ActorFields } } + ... on UserBlockedEvent { createdAt actor { ...ActorFields } blockDuration } +} +""" + +F_PR_TIMELINE = """ +fragment PRTimelineItem on PullRequestTimelineItems { + __typename + ... on Node { id } + ... on AssignedEvent { createdAt actor { ...ActorFields } assignee { __typename ... on User { login } ... on Bot { login } } } + ... on AutoMergeDisabledEvent { createdAt actor { ...ActorFields } reason } + ... on AutoMergeEnabledEvent { createdAt actor { ...ActorFields } } + ... on AutoRebaseEnabledEvent { createdAt actor { ...ActorFields } } + ... on AutoSquashEnabledEvent { createdAt actor { ...ActorFields } } + ... on AutomaticBaseChangeFailedEvent { createdAt actor { ...ActorFields } oldBase newBase } + ... on AutomaticBaseChangeSucceededEvent { createdAt actor { ...ActorFields } oldBase newBase } + ... on BaseRefChangedEvent { createdAt actor { ...ActorFields } previousRefName currentRefName } + ... on BaseRefDeletedEvent { createdAt actor { ...ActorFields } baseRefName } + ... on BaseRefForcePushedEvent { createdAt actor { ...ActorFields } beforeCommit { oid } afterCommit { oid } ref { name } } + ... on ClosedEvent { createdAt actor { ...ActorFields } stateReason } + ... on CommentDeletedEvent { createdAt actor { ...ActorFields } } + ... on ConnectedEvent { createdAt actor { ...ActorFields } source { __typename ... on Issue { number url } ... on PullRequest { number url } } subject { __typename ... on Issue { number url } ... on PullRequest { number url } } } + ... on ConvertToDraftEvent { createdAt actor { ...ActorFields } } + ... on CrossReferencedEvent { createdAt actor { ...ActorFields } isCrossRepository willCloseTarget source { __typename ... on Issue { number url repository { nameWithOwner } title } ... on PullRequest { number url repository { nameWithOwner } title } } } + ... on DemilestonedEvent { createdAt actor { ...ActorFields } milestoneTitle } + ... on DeployedEvent { createdAt actor { ...ActorFields } } + ... on DeploymentEnvironmentChangedEvent { createdAt actor { ...ActorFields } } + ... on DisconnectedEvent { createdAt actor { ...ActorFields } subject { __typename ... on Issue { number url } ... on PullRequest { number url } } source { __typename ... on Issue { number url } ... on PullRequest { number url } } } + ... on HeadRefDeletedEvent { createdAt actor { ...ActorFields } headRefName } + ... on HeadRefForcePushedEvent { createdAt actor { ...ActorFields } beforeCommit { oid } afterCommit { oid } ref { name } } + ... on HeadRefRestoredEvent { createdAt actor { ...ActorFields } } + ... on IssueComment { id databaseId createdAt updatedAt author { ...ActorFields } body url reactionGroups { content reactors { totalCount } } } + ... on LabeledEvent { createdAt actor { ...ActorFields } label { name color } } + ... on LockedEvent { createdAt actor { ...ActorFields } lockReason } + ... on MarkedAsDuplicateEvent { createdAt actor { ...ActorFields } canonical { __typename ... on Issue { number url } ... on PullRequest { number url } } } + ... on MentionedEvent { createdAt actor { ...ActorFields } } + ... on MergedEvent { createdAt actor { ...ActorFields } commit { oid url } mergeRefName } + ... on MilestonedEvent { createdAt actor { ...ActorFields } milestoneTitle } + ... on MovedColumnsInProjectEvent { createdAt actor { ...ActorFields } } + ... on PinnedEvent { createdAt actor { ...ActorFields } } + ... on PullRequestCommit { commit { oid url message author { user { login } date } committedDate } } + ... on PullRequestCommitCommentThread { commit { oid } } + ... on PullRequestReview { id databaseId createdAt submittedAt author { ...ActorFields } body state url reactionGroups { content reactors { totalCount } } } + ... on PullRequestReviewThread { id isResolved isOutdated path line diffSide } + ... on PullRequestRevisionMarker { createdAt lastSeenCommit { oid } } + ... on ReadyForReviewEvent { createdAt actor { ...ActorFields } } + ... on ReferencedEvent { createdAt actor { ...ActorFields } commit { oid url } commitRepository { nameWithOwner } } + ... on RenamedTitleEvent { createdAt actor { ...ActorFields } previousTitle currentTitle } + ... on ReopenedEvent { createdAt actor { ...ActorFields } } + ... on ReviewDismissedEvent { createdAt actor { ...ActorFields } dismissalMessage previousReviewState } + ... on ReviewRequestRemovedEvent { createdAt actor { ...ActorFields } requestedReviewer { __typename ... on User { login } ... on Team { name } } } + ... on ReviewRequestedEvent { createdAt actor { ...ActorFields } requestedReviewer { __typename ... on User { login } ... on Team { name } } } + ... on SubscribedEvent { createdAt actor { ...ActorFields } } + ... on TransferredEvent { createdAt actor { ...ActorFields } fromRepository { nameWithOwner } } + ... on UnassignedEvent { createdAt actor { ...ActorFields } assignee { __typename ... on User { login } ... on Bot { login } } } + ... on UnlabeledEvent { createdAt actor { ...ActorFields } label { name color } } + ... on UnlockedEvent { createdAt actor { ...ActorFields } } + ... on UnmarkedAsDuplicateEvent { createdAt actor { ...ActorFields } } + ... on UnpinnedEvent { createdAt actor { ...ActorFields } } + ... on UnsubscribedEvent { createdAt actor { ...ActorFields } } + ... on UserBlockedEvent { createdAt actor { ...ActorFields } blockDuration } +} +""" + + +def _q(parts: list[str], body: str) -> str: + return "\n".join(parts + [body]) + + +ISSUES_PAGE_QUERY = _q( + [F_ACTOR, F_LABEL, F_TIMELINE], + """ +query IssuesPage($owner: String!, $name: String!, $first: Int!, $after: String) { + repository(owner: $owner, name: $name) { + issues(first: $first, after: $after, orderBy: {field: CREATED_AT, direction: ASC}) { + pageInfo { hasNextPage endCursor } + totalCount + nodes { + id databaseId number title body state stateReason + createdAt updatedAt closedAt + url + author { ...ActorFields } + editor { ...ActorFields } + labels(first: 50) { nodes { ...LabelFields } } + assignees(first: 20) { nodes { login id } } + milestone { title number state dueOn } + reactionGroups { content reactors { totalCount } } + comments(first: 100) { + totalCount + pageInfo { hasNextPage endCursor } + nodes { + id databaseId createdAt updatedAt url body + author { ...ActorFields } + editor { ...ActorFields } + reactionGroups { content reactors { totalCount } } + } + } + timelineItems(first: 100) { + totalCount + pageInfo { hasNextPage endCursor } + nodes { ...TimelineItem } + } + trackedInIssues(first: 20) { totalCount nodes { number url repository { nameWithOwner } } } + trackedIssues(first: 20) { totalCount nodes { number url repository { nameWithOwner } } } + } + } + } + rateLimit { cost remaining resetAt } +} +""", +) + +PRS_PAGE_QUERY = _q( + [F_ACTOR, F_LABEL, F_PR_TIMELINE], + """ +query PRsPage($owner: String!, $name: String!, $first: Int!, $after: String) { + repository(owner: $owner, name: $name) { + pullRequests(first: $first, after: $after, orderBy: {field: CREATED_AT, direction: ASC}) { + pageInfo { hasNextPage endCursor } + totalCount + nodes { + id databaseId number title body state isDraft + createdAt updatedAt closedAt mergedAt + url + headRefName headRefOid + baseRefName baseRefOid + additions deletions changedFiles + mergeable merged mergeStateStatus + author { ...ActorFields } + editor { ...ActorFields } + mergedBy { ...ActorFields } + labels(first: 50) { nodes { ...LabelFields } } + assignees(first: 20) { nodes { login id } } + milestone { title number state dueOn } + reactionGroups { content reactors { totalCount } } + closingIssuesReferences(first: 20) { totalCount nodes { number url repository { nameWithOwner } title } } + comments(first: 100) { + totalCount + pageInfo { hasNextPage endCursor } + nodes { + id databaseId createdAt updatedAt url body + author { ...ActorFields } + editor { ...ActorFields } + reactionGroups { content reactors { totalCount } } + } + } + reviewThreads(first: 50) { + totalCount + pageInfo { hasNextPage endCursor } + nodes { + id isResolved isOutdated path line diffSide + comments(first: 50) { + totalCount + pageInfo { hasNextPage endCursor } + nodes { + id databaseId createdAt updatedAt url body path diffHunk + author { ...ActorFields } + editor { ...ActorFields } + position originalPosition line originalLine + commit { oid } + reactionGroups { content reactors { totalCount } } + } + } + } + } + reviews(first: 50) { + totalCount + pageInfo { hasNextPage endCursor } + nodes { + id databaseId state createdAt submittedAt body url + author { ...ActorFields } + reactionGroups { content reactors { totalCount } } + } + } + commits(first: 100) { + totalCount + pageInfo { hasNextPage endCursor } + nodes { + commit { + oid + message + messageHeadline + committedDate + authoredDate + author { name email user { login } date } + committer { name email user { login } date } + additions deletions changedFilesIfAvailable + parents(first: 3) { nodes { oid } } + } + } + } + files(first: 100) { + totalCount + pageInfo { hasNextPage endCursor } + nodes { + path additions deletions changeType + } + } + timelineItems(first: 100) { + totalCount + pageInfo { hasNextPage endCursor } + nodes { ...PRTimelineItem } + } + } + } + } + rateLimit { cost remaining resetAt } +} +""", +) + +PRS_PAGE_QUERY_LIGHT = _q( + [F_ACTOR, F_LABEL], + """ +query PRsPageLight($owner: String!, $name: String!, $first: Int!, $after: String) { + repository(owner: $owner, name: $name) { + pullRequests(first: $first, after: $after, orderBy: {field: CREATED_AT, direction: ASC}) { + pageInfo { hasNextPage endCursor } + totalCount + nodes { + id databaseId number title body state isDraft + createdAt updatedAt closedAt mergedAt + url + author { ...ActorFields } + labels(first: 50) { nodes { ...LabelFields } } + comments(first: 30) { + totalCount + pageInfo { hasNextPage endCursor } + nodes { + id databaseId createdAt updatedAt url body + author { ...ActorFields } + } + } + } + } + } + rateLimit { cost remaining resetAt } +} +""", +) + +ISSUES_PAGE_QUERY_LIGHT = _q( + [F_ACTOR, F_LABEL], + """ +query IssuesPageLight($owner: String!, $name: String!, $first: Int!, $after: String) { + repository(owner: $owner, name: $name) { + issues(first: $first, after: $after, orderBy: {field: CREATED_AT, direction: ASC}) { + pageInfo { hasNextPage endCursor } + totalCount + nodes { + id databaseId number title body state + createdAt updatedAt closedAt + url + author { ...ActorFields } + labels(first: 50) { nodes { ...LabelFields } } + comments(first: 30) { + totalCount + pageInfo { hasNextPage endCursor } + nodes { + id databaseId createdAt updatedAt url body + author { ...ActorFields } + } + } + } + } + } + rateLimit { cost remaining resetAt } +} +""", +) + +ISSUE_COMMENTS_QUERY = _q( + [F_ACTOR], + """ +query IssueComments($owner: String!, $name: String!, $number: Int!, $after: String) { + repository(owner: $owner, name: $name) { + issueOrPullRequest(number: $number) { + __typename + ... on Issue { + comments(first: 100, after: $after) { + pageInfo { hasNextPage endCursor } + nodes { + id databaseId createdAt updatedAt url body + author { ...ActorFields } + editor { ...ActorFields } + reactionGroups { content reactors { totalCount } } + } + } + } + ... on PullRequest { + comments(first: 100, after: $after) { + pageInfo { hasNextPage endCursor } + nodes { + id databaseId createdAt updatedAt url body + author { ...ActorFields } + editor { ...ActorFields } + reactionGroups { content reactors { totalCount } } + } + } + } + } + } + rateLimit { cost remaining resetAt } +} +""", +) + +ISSUE_TIMELINE_QUERY = _q( + [F_ACTOR, F_TIMELINE], + """ +query IssueTimeline($owner: String!, $name: String!, $number: Int!, $after: String) { + repository(owner: $owner, name: $name) { + issue(number: $number) { + timelineItems(first: 100, after: $after) { + pageInfo { hasNextPage endCursor } + nodes { ...TimelineItem } + } + } + } + rateLimit { cost remaining resetAt } +} +""", +) + +PR_TIMELINE_QUERY = _q( + [F_ACTOR, F_PR_TIMELINE], + """ +query PRTimeline($owner: String!, $name: String!, $number: Int!, $after: String) { + repository(owner: $owner, name: $name) { + pullRequest(number: $number) { + timelineItems(first: 100, after: $after) { + pageInfo { hasNextPage endCursor } + nodes { ...PRTimelineItem } + } + } + } + rateLimit { cost remaining resetAt } +} +""", +) + +PR_COMMITS_QUERY = """ +query PRCommits($owner: String!, $name: String!, $number: Int!, $after: String) { + repository(owner: $owner, name: $name) { + pullRequest(number: $number) { + commits(first: 100, after: $after) { + pageInfo { hasNextPage endCursor } + nodes { + commit { + oid message messageHeadline committedDate authoredDate + author { name email user { login } date } + committer { name email user { login } date } + additions deletions changedFilesIfAvailable + parents(first: 3) { nodes { oid } } + } + } + } + } + } + rateLimit { cost remaining resetAt } +} +""" + +PR_FILES_QUERY = """ +query PRFiles($owner: String!, $name: String!, $number: Int!, $after: String) { + repository(owner: $owner, name: $name) { + pullRequest(number: $number) { + files(first: 100, after: $after) { + pageInfo { hasNextPage endCursor } + nodes { path additions deletions changeType } + } + } + } + rateLimit { cost remaining resetAt } +} +""" + +PR_REVIEW_THREADS_QUERY = _q( + [F_ACTOR], + """ +query PRReviewThreads($owner: String!, $name: String!, $number: Int!, $after: String) { + repository(owner: $owner, name: $name) { + pullRequest(number: $number) { + reviewThreads(first: 50, after: $after) { + pageInfo { hasNextPage endCursor } + nodes { + id isResolved isOutdated path line diffSide + comments(first: 50) { + totalCount + nodes { + id databaseId createdAt updatedAt url body path diffHunk + author { ...ActorFields } + editor { ...ActorFields } + position originalPosition line originalLine + commit { oid } + reactionGroups { content reactors { totalCount } } + } + } + } + } + } + } + rateLimit { cost remaining resetAt } +} +""", +) + +DISCUSSIONS_PAGE_QUERY = _q( + [F_ACTOR, F_LABEL], + """ +query DiscussionsPage($owner: String!, $name: String!, $first: Int!, $after: String) { + repository(owner: $owner, name: $name) { + discussions(first: $first, after: $after, orderBy: {field: CREATED_AT, direction: ASC}) { + pageInfo { hasNextPage endCursor } + totalCount + nodes { + id databaseId number title body + createdAt updatedAt url + author { ...ActorFields } + editor { ...ActorFields } + locked + answerChosenAt + closed closedAt + category { id name emoji description isAnswerable } + labels(first: 30) { nodes { ...LabelFields } } + upvoteCount + answer { id databaseId body author { ...ActorFields } createdAt url } + reactionGroups { content reactors { totalCount } } + comments(first: 50) { + totalCount + pageInfo { hasNextPage endCursor } + nodes { + id databaseId body createdAt updatedAt url + author { ...ActorFields } + editor { ...ActorFields } + upvoteCount + isAnswer + reactionGroups { content reactors { totalCount } } + replies(first: 50) { + totalCount + pageInfo { hasNextPage endCursor } + nodes { + id databaseId body createdAt updatedAt url + author { ...ActorFields } + editor { ...ActorFields } + reactionGroups { content reactors { totalCount } } + } + } + } + } + } + } + } + rateLimit { cost remaining resetAt } +} +""", +) + +DISCUSSION_COMMENTS_QUERY = _q( + [F_ACTOR], + """ +query DiscussionComments($owner: String!, $name: String!, $number: Int!, $after: String) { + repository(owner: $owner, name: $name) { + discussion(number: $number) { + comments(first: 50, after: $after) { + pageInfo { hasNextPage endCursor } + nodes { + id databaseId body createdAt updatedAt url + author { ...ActorFields } + editor { ...ActorFields } + upvoteCount + isAnswer + reactionGroups { content reactors { totalCount } } + replies(first: 50) { + totalCount + nodes { + id databaseId body createdAt updatedAt url + author { ...ActorFields } + editor { ...ActorFields } + reactionGroups { content reactors { totalCount } } + } + } + } + } + } + } + rateLimit { cost remaining resetAt } +} +""", +) + +DISCUSSION_REPLIES_QUERY = _q( + [F_ACTOR], + """ +query DiscussionReplies($commentId: ID!, $after: String) { + node(id: $commentId) { + ... on DiscussionComment { + replies(first: 50, after: $after) { + pageInfo { hasNextPage endCursor } + nodes { + id databaseId body createdAt updatedAt url + author { ...ActorFields } + editor { ...ActorFields } + reactionGroups { content reactors { totalCount } } + } + } + } + } + rateLimit { cost remaining resetAt } +} +""", +) + +COMMITS_PAGE_QUERY = """ +query CommitsPage($owner: String!, $name: String!, $first: Int!, $after: String, $branch: String!) { + repository(owner: $owner, name: $name) { + ref(qualifiedName: $branch) { + target { + ... on Commit { + history(first: $first, after: $after) { + pageInfo { hasNextPage endCursor } + totalCount + nodes { + oid + message + messageHeadline + committedDate + authoredDate + url + additions deletions changedFilesIfAvailable + author { name email date user { login id } } + committer { name email date user { login id } } + parents(first: 3) { nodes { oid } } + associatedPullRequests(first: 5) { nodes { number url state } } + } + } + } + } + } + } + rateLimit { cost remaining resetAt } +} +""" + +RELEASES_QUERY = _q( + [F_ACTOR], + """ +query Releases($owner: String!, $name: String!, $first: Int!, $after: String) { + repository(owner: $owner, name: $name) { + releases(first: $first, after: $after, orderBy: {field: CREATED_AT, direction: ASC}) { + pageInfo { hasNextPage endCursor } + nodes { + id databaseId name tagName description + createdAt publishedAt updatedAt + isDraft isPrerelease isLatest + url + author { ...ActorFields } + tagCommit { oid url } + reactionGroups { content reactors { totalCount } } + releaseAssets(first: 50) { + nodes { name contentType size downloadUrl createdAt updatedAt } + } + } + } + } + rateLimit { cost remaining resetAt } +} +""", +) + +LABELS_QUERY = _q( + [F_LABEL], + """ +query LabelsList($owner: String!, $name: String!, $first: Int!, $after: String) { + repository(owner: $owner, name: $name) { + labels(first: $first, after: $after) { + pageInfo { hasNextPage endCursor } + nodes { ...LabelFields } + } + } + rateLimit { cost remaining resetAt } +} +""", +) + +MILESTONES_QUERY = """ +query Milestones($owner: String!, $name: String!, $first: Int!, $after: String) { + repository(owner: $owner, name: $name) { + milestones(first: $first, after: $after) { + pageInfo { hasNextPage endCursor } + nodes { + id number title description state + createdAt updatedAt closedAt dueOn + creator { login } + } + } + } + rateLimit { cost remaining resetAt } +} +""" + +REPO_META_QUERY = """ +query RepoMeta($owner: String!, $name: String!) { + repository(owner: $owner, name: $name) { + id databaseId name nameWithOwner description url + createdAt updatedAt pushedAt + isArchived isDisabled isFork isPrivate + primaryLanguage { name } + languages(first: 20, orderBy: {field: SIZE, direction: DESC}) { + edges { size node { name } } + totalSize + } + stargazerCount forkCount watchers { totalCount } + diskUsage + licenseInfo { key name } + homepageUrl + defaultBranchRef { name } + } + rateLimit { cost remaining resetAt } +} +""" diff --git a/studio/backend/plugins/data-designer-github-repo-seed/src/data_designer_github_repo_seed/scraper_impl/scraper.py b/studio/backend/plugins/data-designer-github-repo-seed/src/data_designer_github_repo_seed/scraper_impl/scraper.py new file mode 100644 index 0000000000..127129e18b --- /dev/null +++ b/studio/backend/plugins/data-designer-github-repo-seed/src/data_designer_github_repo_seed/scraper_impl/scraper.py @@ -0,0 +1,756 @@ +# SPDX-License-Identifier: AGPL-3.0-only +# Copyright 2026-present the Unsloth AI Inc. team. All rights reserved. See /studio/LICENSE.AGPL-3.0 + +"""Main scraper orchestration. Collects issues, PRs, discussions, commits, releases, etc. + +Resumable via state file. Writes JSONL shards under data/{repo}/{resource}.jsonl. +""" + +from __future__ import annotations + +import argparse +import json +import logging +import os +import subprocess +import sys +import time +from pathlib import Path +from typing import Any, Dict, Iterable, List, Optional, Tuple + +# Allow running as a module or script +THIS_DIR = Path(__file__).resolve().parent +if str(THIS_DIR) not in sys.path: + sys.path.insert(0, str(THIS_DIR)) + +from gh_client import GitHubClient +from state_store import JsonlWriter, StateStore +import queries as Q + +log = logging.getLogger("scraper") + + +def ts() -> str: + return time.strftime("%Y-%m-%d %H:%M:%S") + + +class RepoScraper: + def __init__( + self, + owner: str, + name: str, + base_dir: Path, + client: GitHubClient, + trial_limits: Optional[Dict[str, int]] = None, + light: bool = False, + ): + self.owner = owner + self.name = name + self.base_dir = base_dir + self.client = client + self.trial_limits = trial_limits or {} + # When light=True, use trimmed GraphQL queries (no reviewThreads, + # reviews, commits, timelineItems, files) so PR pages can be much + # larger without blowing GitHub's node-count ceiling. + self.light = light + self.repo_dir = base_dir / f"{owner}__{name}" + self.repo_dir.mkdir(parents = True, exist_ok = True) + self.state = StateStore(base_dir / "state" / f"{owner}__{name}.json") + + # Writers + self.writers: Dict[str, JsonlWriter] = {} + for key in ( + "issues", + "pull_requests", + "discussions", + "commits", + "releases", + "labels", + "milestones", + "pr_extra_comments", + "pr_extra_timeline", + "pr_extra_reviews", + "issue_extra_comments", + "issue_extra_timeline", + "discussion_extra_comments", + "discussion_extra_replies", + "repo_meta", + ): + self.writers[key] = JsonlWriter(self.repo_dir / f"{key}.jsonl") + + # ----- helpers ----- + def _trial_stop(self, key: str, counter: int) -> bool: + lim = self.trial_limits.get(key) + if lim is None: + return False + return counter >= lim + + def _log_rate(self, where: str, data: Dict[str, Any]) -> None: + rl = ( + data.get("data", {}).get("rateLimit") + if isinstance(data.get("data"), dict) + else None + ) + if rl: + log.debug( + "[%s] rate cost=%s remaining=%s resetAt=%s", + where, + rl.get("cost"), + rl.get("remaining"), + rl.get("resetAt"), + ) + + # ----- repo meta ----- + def scrape_repo_meta(self) -> Dict[str, Any]: + data = self.client.graphql( + Q.REPO_META_QUERY, {"owner": self.owner, "name": self.name} + ) + self._log_rate("repo_meta", data) + repo = data.get("data", {}).get("repository") or {} + repo["_fetchedAt"] = ts() + self.writers["repo_meta"].write(repo) + return repo + + # ----- issues ----- + def scrape_issues(self) -> int: + key = "issues" + cursor = self.state.get(f"{key}_cursor") + done = self.state.get(f"{key}_done", False) + if done: + log.info("%s/%s issues already complete", self.owner, self.name) + return 0 + total_new = 0 + page = 0 + # Light query skips heavy nested fields; safe at 50 per page. + # Clamp by trial_limit so e.g. limit=1 asks GitHub for first:1 + # instead of fetching a full 50-item page and discarding 49. + page_cap = 50 if self.light else 15 + trial_cap = self.trial_limits.get(key) + per_page = min(page_cap, trial_cap) if trial_cap and trial_cap > 0 else page_cap + while True: + page += 1 + vars_ = { + "owner": self.owner, + "name": self.name, + "first": per_page, + "after": cursor, + } + query = Q.ISSUES_PAGE_QUERY_LIGHT if self.light else Q.ISSUES_PAGE_QUERY + data = self.client.graphql(query, vars_) + self._log_rate("issues", data) + repo = (data.get("data") or {}).get("repository") or {} + issues = repo.get("issues") or {} + nodes = issues.get("nodes") or [] + for it in nodes: + it["_owner"] = self.owner + it["_repo"] = self.name + it["_fetchedAt"] = ts() + if not self.light: + if it.get("comments", {}).get("pageInfo", {}).get("hasNextPage"): + self._paginate_issue_comments( + it["number"], it["comments"]["pageInfo"]["endCursor"] + ) + if ( + it.get("timelineItems", {}) + .get("pageInfo", {}) + .get("hasNextPage") + ): + self._paginate_issue_timeline( + it["number"], + it["timelineItems"]["pageInfo"]["endCursor"], + ) + if self.writers[key].write(it): + total_new += 1 + info = issues.get("pageInfo") or {} + cursor = info.get("endCursor") + self.state.set(f"{key}_cursor", cursor) + log.info( + "[%s/%s] issues page %d (+%d) cursor=%s remaining=%s", + self.owner, + self.name, + page, + len(nodes), + str(cursor)[:20], + self.client.graphql_remaining, + ) + if self._trial_stop(key, total_new): + log.info("Trial limit reached for issues (%d)", total_new) + return total_new + if not info.get("hasNextPage"): + self.state.set(f"{key}_done", True) + break + return total_new + + def _paginate_issue_comments(self, number: int, after: str) -> None: + cur = after + while cur: + vars_ = { + "owner": self.owner, + "name": self.name, + "number": number, + "after": cur, + } + data = self.client.graphql(Q.ISSUE_COMMENTS_QUERY, vars_) + item = ((data.get("data") or {}).get("repository") or {}).get( + "issueOrPullRequest" + ) or {} + comments = item.get("comments") or {} + for c in comments.get("nodes") or []: + c["_owner"] = self.owner + c["_repo"] = self.name + c["_issueNumber"] = number + self.writers["issue_extra_comments"].write(c) + info = comments.get("pageInfo") or {} + cur = info.get("endCursor") if info.get("hasNextPage") else None + + def _paginate_issue_timeline(self, number: int, after: str) -> None: + cur = after + while cur: + vars_ = { + "owner": self.owner, + "name": self.name, + "number": number, + "after": cur, + } + data = self.client.graphql(Q.ISSUE_TIMELINE_QUERY, vars_) + item = ((data.get("data") or {}).get("repository") or {}).get("issue") or {} + tl = item.get("timelineItems") or {} + for ev in tl.get("nodes") or []: + ev["_owner"] = self.owner + ev["_repo"] = self.name + ev["_issueNumber"] = number + self.writers["issue_extra_timeline"].write(ev) + info = tl.get("pageInfo") or {} + cur = info.get("endCursor") if info.get("hasNextPage") else None + + # ----- PRs ----- + def scrape_prs(self) -> int: + key = "pull_requests" + cursor = self.state.get(f"{key}_cursor") + done = self.state.get(f"{key}_done", False) + if done: + log.info("%s/%s PRs already complete", self.owner, self.name) + return 0 + total_new = 0 + page = 0 + # Heavy nested PR query is capped at 3 per page (GitHub node-count + # ceiling); light query skips reviewThreads/reviews/commits/etc and + # can safely go to 25 per page. Clamp by trial_limit for small + # previews so limit=1 does not fetch a whole 25-item page. + page_cap = 25 if self.light else 3 + trial_cap = self.trial_limits.get(key) + per_page = min(page_cap, trial_cap) if trial_cap and trial_cap > 0 else page_cap + while True: + page += 1 + vars_ = { + "owner": self.owner, + "name": self.name, + "first": per_page, + "after": cursor, + } + query = Q.PRS_PAGE_QUERY_LIGHT if self.light else Q.PRS_PAGE_QUERY + data = self.client.graphql(query, vars_) + self._log_rate("prs", data) + repo = (data.get("data") or {}).get("repository") or {} + prs = repo.get("pullRequests") or {} + nodes = prs.get("nodes") or [] + for pr in nodes: + pr["_owner"] = self.owner + pr["_repo"] = self.name + pr["_fetchedAt"] = ts() + num = pr["number"] + if not self.light: + if pr.get("comments", {}).get("pageInfo", {}).get("hasNextPage"): + self._paginate_pr_comments( + num, pr["comments"]["pageInfo"]["endCursor"] + ) + if ( + pr.get("timelineItems", {}) + .get("pageInfo", {}) + .get("hasNextPage") + ): + self._paginate_pr_timeline( + num, pr["timelineItems"]["pageInfo"]["endCursor"] + ) + if pr.get("commits", {}).get("pageInfo", {}).get("hasNextPage"): + self._paginate_pr_commits( + num, pr["commits"]["pageInfo"]["endCursor"] + ) + if pr.get("files", {}).get("pageInfo", {}).get("hasNextPage"): + self._paginate_pr_files( + num, pr["files"]["pageInfo"]["endCursor"] + ) + if ( + pr.get("reviewThreads", {}) + .get("pageInfo", {}) + .get("hasNextPage") + ): + self._paginate_pr_review_threads( + num, pr["reviewThreads"]["pageInfo"]["endCursor"] + ) + if self.writers[key].write(pr): + total_new += 1 + info = prs.get("pageInfo") or {} + cursor = info.get("endCursor") + self.state.set(f"{key}_cursor", cursor) + log.info( + "[%s/%s] PRs page %d (+%d) cursor=%s remaining=%s", + self.owner, + self.name, + page, + len(nodes), + str(cursor)[:20], + self.client.graphql_remaining, + ) + if self._trial_stop(key, total_new): + log.info("Trial limit reached for PRs (%d)", total_new) + return total_new + if not info.get("hasNextPage"): + self.state.set(f"{key}_done", True) + break + return total_new + + def _paginate_pr_comments(self, number: int, after: str) -> None: + cur = after + while cur: + vars_ = { + "owner": self.owner, + "name": self.name, + "number": number, + "after": cur, + } + data = self.client.graphql(Q.ISSUE_COMMENTS_QUERY, vars_) + item = ((data.get("data") or {}).get("repository") or {}).get( + "issueOrPullRequest" + ) or {} + comments = item.get("comments") or {} + for c in comments.get("nodes") or []: + c["_owner"] = self.owner + c["_repo"] = self.name + c["_prNumber"] = number + self.writers["pr_extra_comments"].write(c) + info = comments.get("pageInfo") or {} + cur = info.get("endCursor") if info.get("hasNextPage") else None + + def _paginate_pr_timeline(self, number: int, after: str) -> None: + cur = after + while cur: + vars_ = { + "owner": self.owner, + "name": self.name, + "number": number, + "after": cur, + } + data = self.client.graphql(Q.PR_TIMELINE_QUERY, vars_) + item = ((data.get("data") or {}).get("repository") or {}).get( + "pullRequest" + ) or {} + tl = item.get("timelineItems") or {} + for ev in tl.get("nodes") or []: + ev["_owner"] = self.owner + ev["_repo"] = self.name + ev["_prNumber"] = number + self.writers["pr_extra_timeline"].write(ev) + info = tl.get("pageInfo") or {} + cur = info.get("endCursor") if info.get("hasNextPage") else None + + def _paginate_pr_commits(self, number: int, after: str) -> None: + cur = after + out_key = "pr_extra_commits" + if out_key not in self.writers: + self.writers[out_key] = JsonlWriter(self.repo_dir / f"{out_key}.jsonl") + while cur: + vars_ = { + "owner": self.owner, + "name": self.name, + "number": number, + "after": cur, + } + data = self.client.graphql(Q.PR_COMMITS_QUERY, vars_) + item = ((data.get("data") or {}).get("repository") or {}).get( + "pullRequest" + ) or {} + cc = item.get("commits") or {} + for c in cc.get("nodes") or []: + c["_owner"] = self.owner + c["_repo"] = self.name + c["_prNumber"] = number + self.writers[out_key].write(c) + info = cc.get("pageInfo") or {} + cur = info.get("endCursor") if info.get("hasNextPage") else None + + def _paginate_pr_files(self, number: int, after: str) -> None: + cur = after + out_key = "pr_extra_files" + if out_key not in self.writers: + self.writers[out_key] = JsonlWriter(self.repo_dir / f"{out_key}.jsonl") + while cur: + vars_ = { + "owner": self.owner, + "name": self.name, + "number": number, + "after": cur, + } + data = self.client.graphql(Q.PR_FILES_QUERY, vars_) + item = ((data.get("data") or {}).get("repository") or {}).get( + "pullRequest" + ) or {} + ff = item.get("files") or {} + for f in ff.get("nodes") or []: + f["_owner"] = self.owner + f["_repo"] = self.name + f["_prNumber"] = number + # files don't have id, synthesize one + f["_syntheticId"] = f"{self.owner}/{self.name}#{number}:{f.get('path')}" + self.writers[out_key].write(f) + info = ff.get("pageInfo") or {} + cur = info.get("endCursor") if info.get("hasNextPage") else None + + def _paginate_pr_review_threads(self, number: int, after: str) -> None: + cur = after + out_key = "pr_extra_review_threads" + if out_key not in self.writers: + self.writers[out_key] = JsonlWriter(self.repo_dir / f"{out_key}.jsonl") + while cur: + vars_ = { + "owner": self.owner, + "name": self.name, + "number": number, + "after": cur, + } + data = self.client.graphql(Q.PR_REVIEW_THREADS_QUERY, vars_) + item = ((data.get("data") or {}).get("repository") or {}).get( + "pullRequest" + ) or {} + rt = item.get("reviewThreads") or {} + for th in rt.get("nodes") or []: + th["_owner"] = self.owner + th["_repo"] = self.name + th["_prNumber"] = number + self.writers[out_key].write(th) + info = rt.get("pageInfo") or {} + cur = info.get("endCursor") if info.get("hasNextPage") else None + + # ----- Discussions ----- + def scrape_discussions(self) -> int: + key = "discussions" + cursor = self.state.get(f"{key}_cursor") + done = self.state.get(f"{key}_done", False) + if done: + log.info("%s/%s discussions already complete", self.owner, self.name) + return 0 + total_new = 0 + page = 0 + per_page = 15 + while True: + page += 1 + vars_ = { + "owner": self.owner, + "name": self.name, + "first": per_page, + "after": cursor, + } + data = self.client.graphql(Q.DISCUSSIONS_PAGE_QUERY, vars_) + self._log_rate("discussions", data) + repo = (data.get("data") or {}).get("repository") or {} + dd = repo.get("discussions") or {} + nodes = dd.get("nodes") or [] + for d in nodes: + d["_owner"] = self.owner + d["_repo"] = self.name + d["_fetchedAt"] = ts() + num = d["number"] + if d.get("comments", {}).get("pageInfo", {}).get("hasNextPage"): + self._paginate_discussion_comments( + num, d["comments"]["pageInfo"]["endCursor"] + ) + # paginate replies per comment if needed + for c in d.get("comments", {}).get("nodes", []) or []: + if c.get("replies", {}).get("pageInfo", {}).get("hasNextPage"): + self._paginate_discussion_replies( + c["id"], c["replies"]["pageInfo"]["endCursor"], num + ) + if self.writers[key].write(d): + total_new += 1 + info = dd.get("pageInfo") or {} + cursor = info.get("endCursor") + self.state.set(f"{key}_cursor", cursor) + log.info( + "[%s/%s] discussions page %d (+%d) cursor=%s remaining=%s", + self.owner, + self.name, + page, + len(nodes), + str(cursor)[:20], + self.client.graphql_remaining, + ) + if self._trial_stop(key, total_new): + return total_new + if not info.get("hasNextPage"): + self.state.set(f"{key}_done", True) + break + return total_new + + def _paginate_discussion_comments(self, number: int, after: str) -> None: + cur = after + while cur: + vars_ = { + "owner": self.owner, + "name": self.name, + "number": number, + "after": cur, + } + data = self.client.graphql(Q.DISCUSSION_COMMENTS_QUERY, vars_) + disc = ((data.get("data") or {}).get("repository") or {}).get( + "discussion" + ) or {} + cc = disc.get("comments") or {} + for c in cc.get("nodes") or []: + c["_owner"] = self.owner + c["_repo"] = self.name + c["_discussionNumber"] = number + self.writers["discussion_extra_comments"].write(c) + info = cc.get("pageInfo") or {} + cur = info.get("endCursor") if info.get("hasNextPage") else None + + def _paginate_discussion_replies( + self, comment_id: str, after: str, disc_number: int + ) -> None: + cur = after + while cur: + vars_ = { + "owner": self.owner, + "name": self.name, + "commentId": comment_id, + "after": cur, + } + data = self.client.graphql(Q.DISCUSSION_REPLIES_QUERY, vars_) + node = (data.get("data") or {}).get("node") or {} + replies = node.get("replies") or {} + for r in replies.get("nodes") or []: + r["_owner"] = self.owner + r["_repo"] = self.name + r["_discussionNumber"] = disc_number + r["_commentId"] = comment_id + self.writers["discussion_extra_replies"].write(r) + info = replies.get("pageInfo") or {} + cur = info.get("endCursor") if info.get("hasNextPage") else None + + # ----- Commits ----- + def scrape_commits(self, branch: str = "refs/heads/main") -> int: + key = "commits" + cursor = self.state.get(f"{key}_cursor") + done = self.state.get(f"{key}_done", False) + if done: + return 0 + total_new = 0 + page = 0 + page_cap = 100 + trial_cap = self.trial_limits.get(key) + per_page = min(page_cap, trial_cap) if trial_cap and trial_cap > 0 else page_cap + while True: + page += 1 + vars_ = { + "owner": self.owner, + "name": self.name, + "first": per_page, + "after": cursor, + "branch": branch, + } + data = self.client.graphql(Q.COMMITS_PAGE_QUERY, vars_) + self._log_rate("commits", data) + ref = ((data.get("data") or {}).get("repository") or {}).get("ref") or {} + tgt = ref.get("target") or {} + hist = tgt.get("history") or {} + nodes = hist.get("nodes") or [] + for c in nodes: + c["_owner"] = self.owner + c["_repo"] = self.name + c["_fetchedAt"] = ts() + if self.writers[key].write(c): + total_new += 1 + info = hist.get("pageInfo") or {} + cursor = info.get("endCursor") + self.state.set(f"{key}_cursor", cursor) + log.info( + "[%s/%s] commits page %d (+%d) remaining=%s", + self.owner, + self.name, + page, + len(nodes), + self.client.graphql_remaining, + ) + if self._trial_stop(key, total_new): + return total_new + if not info.get("hasNextPage"): + self.state.set(f"{key}_done", True) + break + return total_new + + # ----- Releases/Labels/Milestones ----- + def scrape_releases(self) -> int: + return self._scrape_simple("releases", Q.RELEASES_QUERY, "releases") + + def scrape_labels(self) -> int: + return self._scrape_simple("labels", Q.LABELS_QUERY, "labels") + + def scrape_milestones(self) -> int: + return self._scrape_simple("milestones", Q.MILESTONES_QUERY, "milestones") + + def _scrape_simple(self, key: str, query: str, field: str) -> int: + cursor = self.state.get(f"{key}_cursor") + done = self.state.get(f"{key}_done", False) + if done: + return 0 + total_new = 0 + while True: + vars_ = { + "owner": self.owner, + "name": self.name, + "first": 50, + "after": cursor, + } + data = self.client.graphql(query, vars_) + repo = (data.get("data") or {}).get("repository") or {} + col = repo.get(field) or {} + for it in col.get("nodes") or []: + it["_owner"] = self.owner + it["_repo"] = self.name + it["_fetchedAt"] = ts() + if self.writers[key].write(it): + total_new += 1 + info = col.get("pageInfo") or {} + cursor = info.get("endCursor") + self.state.set(f"{key}_cursor", cursor) + if self._trial_stop(key, total_new): + return total_new + if not info.get("hasNextPage"): + self.state.set(f"{key}_done", True) + break + log.info("[%s/%s] %s done +%d", self.owner, self.name, key, total_new) + return total_new + + def close(self) -> None: + for w in self.writers.values(): + try: + w.close() + except Exception: + pass + + +def setup_logging(log_file: Path) -> None: + log_file.parent.mkdir(parents = True, exist_ok = True) + fmt = "%(asctime)s %(levelname)s [%(name)s] %(message)s" + handlers = [ + logging.StreamHandler(sys.stdout), + logging.FileHandler(log_file, mode = "a", encoding = "utf-8"), + ] + logging.basicConfig(level = logging.INFO, format = fmt, handlers = handlers, force = True) + + +def main(): + ap = argparse.ArgumentParser() + ap.add_argument( + "--base-dir", default = "/mnt/disks/unslothai/ubuntu/workspace_34/github_scraper" + ) + ap.add_argument( + "--repos", nargs = "+", default = ["unslothai/unsloth", "unslothai/unsloth-zoo"] + ) + ap.add_argument("--trial", action = "store_true", help = "Small trial run") + ap.add_argument( + "--only", + nargs = "+", + default = None, + help = "Only run these resource keys: issues,pulls,discussions,commits,releases,labels,milestones,meta", + ) + ap.add_argument( + "--hf-upload-interval", + type = int, + default = 900, + help = "Seconds between HF uploads (0 to disable)", + ) + args = ap.parse_args() + + base = Path(args.base_dir) + data_dir = base / "data" + data_dir.mkdir(parents = True, exist_ok = True) + setup_logging(base / "logs" / f"scraper_{time.strftime('%Y%m%d_%H%M%S')}.log") + log.info("Scraper starting: repos=%s trial=%s", args.repos, args.trial) + + client = GitHubClient(min_remaining_graphql = 80, min_remaining_rest = 80) + rl = client.rate_snapshot() + log.info( + "Rate limit snapshot: %s", + json.dumps(rl.get("resources", {}), default = str)[:400], + ) + + # Start HF uploader in background if requested + uploader = None + if args.hf_upload_interval > 0: + from hf_uploader import HFUploader + + uploader = HFUploader(data_dir, interval_s = args.hf_upload_interval) + uploader.start() + + trial_limits = None + if args.trial: + trial_limits = { + "issues": 5, + "pull_requests": 5, + "discussions": 3, + "commits": 20, + "releases": 3, + "labels": 20, + "milestones": 20, + } + + only = set(args.only or []) + + try: + for repo_spec in args.repos: + owner, name = repo_spec.split("/") + scraper = RepoScraper(owner, name, data_dir, client, trial_limits) + try: + repo_meta: Dict[str, Any] = {} + if not only or "meta" in only or "commits" in only: + repo_meta = scraper.scrape_repo_meta() + if not only or "labels" in only: + scraper.scrape_labels() + if not only or "milestones" in only: + scraper.scrape_milestones() + if not only or "releases" in only: + scraper.scrape_releases() + if not only or "discussions" in only: + scraper.scrape_discussions() + if not only or "issues" in only: + scraper.scrape_issues() + if not only or "pulls" in only: + scraper.scrape_prs() + if not only or "commits" in only: + default_ref = repo_meta.get("defaultBranchRef") or {} + default_branch = ( + default_ref.get("name") + if isinstance(default_ref, dict) + else None + ) + branch = ( + f"refs/heads/{default_branch}" + if default_branch + else "refs/heads/main" + ) + scraper.scrape_commits(branch = branch) + finally: + scraper.close() + finally: + if uploader: + log.info("Stopping uploader and final sync...") + uploader.stop(final_upload = True) + log.info( + "Scraper complete. GraphQL calls=%d REST calls=%d", + client.calls_graphql, + client.calls_rest, + ) + + +if __name__ == "__main__": + main() diff --git a/studio/backend/plugins/data-designer-github-repo-seed/src/data_designer_github_repo_seed/scraper_impl/state_store.py b/studio/backend/plugins/data-designer-github-repo-seed/src/data_designer_github_repo_seed/scraper_impl/state_store.py new file mode 100644 index 0000000000..efa663db2f --- /dev/null +++ b/studio/backend/plugins/data-designer-github-repo-seed/src/data_designer_github_repo_seed/scraper_impl/state_store.py @@ -0,0 +1,105 @@ +# SPDX-License-Identifier: AGPL-3.0-only +# Copyright 2026-present the Unsloth AI Inc. team. All rights reserved. See /studio/LICENSE.AGPL-3.0 + +"""Checkpoint state management for resumable scraping.""" + +from __future__ import annotations + +import json +import os +import threading +from pathlib import Path +from typing import Any, Dict + + +class StateStore: + def __init__(self, path: str | Path): + self.path = Path(path) + self.path.parent.mkdir(parents = True, exist_ok = True) + self._lock = threading.Lock() + self._data: Dict[str, Any] = {} + if self.path.exists(): + try: + with self.path.open() as f: + self._data = json.load(f) + except Exception: + self._data = {} + + def get(self, key: str, default: Any = None) -> Any: + with self._lock: + return self._data.get(key, default) + + def set(self, key: str, value: Any) -> None: + with self._lock: + self._data[key] = value + self._flush() + + def update(self, key: str, **kwargs) -> None: + with self._lock: + sub = dict(self._data.get(key, {})) + sub.update(kwargs) + self._data[key] = sub + self._flush() + + def all(self) -> Dict[str, Any]: + with self._lock: + return dict(self._data) + + def _flush(self) -> None: + tmp = self.path.with_suffix(self.path.suffix + ".tmp") + with tmp.open("w") as f: + json.dump(self._data, f, indent = 2, default = str) + os.replace(tmp, self.path) + + +class JsonlWriter: + """Append-only JSONL writer, thread-safe, with line buffering.""" + + def __init__(self, path: str | Path): + self.path = Path(path) + self.path.parent.mkdir(parents = True, exist_ok = True) + self._lock = threading.Lock() + self._fh = self.path.open("a", buffering = 1) + self._count_seen_keys: set[str] = set() + # Preload seen keys if file exists (for dedup across resumes) + if self.path.exists() and self.path.stat().st_size > 0: + try: + with self.path.open() as f: + for line in f: + try: + obj = json.loads(line) + k = self._key(obj) + if k is not None: + self._count_seen_keys.add(k) + except Exception: + pass + except Exception: + pass + + def _key(self, obj: dict) -> str | None: + for k in ("id", "node_id", "number", "sha", "url"): + if k in obj: + return f"{k}:{obj[k]}" + return None + + def has(self, key: str) -> bool: + return key in self._count_seen_keys + + def write(self, obj: dict) -> bool: + """Return True if newly written, False if already present.""" + k = self._key(obj) + with self._lock: + if k is not None and k in self._count_seen_keys: + return False + if k is not None: + self._count_seen_keys.add(k) + self._fh.write(json.dumps(obj, default = str, ensure_ascii = False)) + self._fh.write("\n") + self._fh.flush() + return True + + def close(self) -> None: + try: + self._fh.close() + except Exception: + pass diff --git a/studio/backend/requirements/single-env/data-designer-deps.txt b/studio/backend/requirements/single-env/data-designer-deps.txt index fc63230922..f63c076621 100644 --- a/studio/backend/requirements/single-env/data-designer-deps.txt +++ b/studio/backend/requirements/single-env/data-designer-deps.txt @@ -19,7 +19,8 @@ ruff<1,>=0.14.10 scipy<2,>=1.11.0 sqlfluff<4,>=3.2.0 tiktoken<1,>=0.8.0 -# Unstructured-seed plugin deps (plugin installed with --no-deps) +# Local seed plugin deps (plugins installed with --no-deps) +requests>=2.31 pymupdf>=1.24.0 pymupdf4llm>=0.0.17 mammoth>=1.8.0 diff --git a/studio/backend/routes/__init__.py b/studio/backend/routes/__init__.py index e79f6553f9..cf4586281b 100644 --- a/studio/backend/routes/__init__.py +++ b/studio/backend/routes/__init__.py @@ -8,6 +8,7 @@ API Routes from routes.training import router as training_router from routes.models import router as models_router from routes.inference import router as inference_router +from routes.inference import studio_router as inference_studio_router from routes.datasets import router as datasets_router from routes.auth import router as auth_router from routes.data_recipe import router as data_recipe_router @@ -18,6 +19,7 @@ __all__ = [ "training_router", "models_router", "inference_router", + "inference_studio_router", "datasets_router", "auth_router", "data_recipe_router", diff --git a/studio/backend/routes/data_recipe/jobs.py b/studio/backend/routes/data_recipe/jobs.py index 606ef1832c..da6416e324 100644 --- a/studio/backend/routes/data_recipe/jobs.py +++ b/studio/backend/routes/data_recipe/jobs.py @@ -5,8 +5,9 @@ from __future__ import annotations -from datetime import timedelta -from typing import Any +import copy +from datetime import datetime, timedelta, timezone +from typing import Any, Optional from urllib.parse import urlparse from fastapi import APIRouter, HTTPException, Query, Request @@ -94,14 +95,111 @@ def _used_llm_model_aliases(recipe: dict[str, Any]) -> set[str]: return aliases -def _inject_local_providers(recipe: dict[str, Any], request: Request) -> None: +def _inject_local_structured_response_format( + recipe: dict[str, Any], local_provider_names: set[str] +) -> None: + """For each llm-structured column that targets a local-provider model_config, + clone the model_config and inject an OpenAI ``response_format`` with the + column's ``output_format`` JSON schema. The column is rewritten to point at + the clone so llm-text / llm-judge columns that share the same alias keep + free-form sampling. + + Without this, data_designer only injects a prompt-level "return JSON in a + ```json fence" instruction. Small GGUF models frequently break format, + wasting the full ``max_tokens`` budget per row and then failing to parse. + Forwarding ``response_format`` lets llama-server apply grammar-constrained + sampling from the JSON schema, which guarantees a parseable response and + terminates early. + """ + columns = recipe.get("columns") + model_configs = recipe.get("model_configs") + if not isinstance(columns, list) or not isinstance(model_configs, list): + return + + # alias -> model_config (only configs referencing a local provider qualify). + alias_to_local_mc: dict[str, dict[str, Any]] = {} + for mc in model_configs: + if not isinstance(mc, dict): + continue + if mc.get("provider") in local_provider_names and isinstance( + mc.get("alias"), str + ): + alias_to_local_mc[mc["alias"]] = mc + + if not alias_to_local_mc: + return + + # Clone per (alias, column) so each llm-structured column gets its own + # schema without leaking response_format onto other columns that share the + # same base alias. + seen_clone_aliases: set[str] = { + mc.get("alias") for mc in model_configs if isinstance(mc.get("alias"), str) + } + new_configs: list[dict[str, Any]] = [] + for column in columns: + if not isinstance(column, dict): + continue + if column.get("column_type") != "llm-structured": + continue + alias = column.get("model_alias") + if not isinstance(alias, str) or alias not in alias_to_local_mc: + continue + output_format = column.get("output_format") + if not isinstance(output_format, dict) or not output_format: + continue + base_mc = alias_to_local_mc[alias] + column_name = column.get("name") or "structured" + clone_alias_base = f"{alias}__{column_name}_structured" + clone_alias = clone_alias_base + counter = 1 + while clone_alias in seen_clone_aliases: + counter += 1 + clone_alias = f"{clone_alias_base}_{counter}" + seen_clone_aliases.add(clone_alias) + + clone = copy.deepcopy(base_mc) + clone["alias"] = clone_alias + params = clone.get("inference_parameters") + if not isinstance(params, dict): + params = {} + clone["inference_parameters"] = params + # data_designer's BaseInferenceParams is a pydantic model with + # extra="forbid", so response_format cannot sit at the top level of + # inference_parameters. It does expose an `extra_body: dict` pass- + # through that the OpenAI client spreads into the request body at the + # top level, which is where llama-server reads response_format from. + # llama.cpp server shape (tools/server/README.md): the schema sits + # directly under response_format, not nested in a json_schema object + # the way OpenAI's Chat Completions API expects. llama-server converts + # the schema to a GBNF grammar and applies it during sampling. + extra_body = params.get("extra_body") + if not isinstance(extra_body, dict): + extra_body = {} + extra_body["response_format"] = { + "type": "json_schema", + "schema": output_format, + } + params["extra_body"] = extra_body + new_configs.append(clone) + column["model_alias"] = clone_alias + + if new_configs: + model_configs.extend(new_configs) + + +def _inject_local_providers(recipe: dict[str, Any], request: Request) -> Optional[int]: """ Mutate recipe dict in-place: for any provider with is_local=True, - generate a JWT and fill in the endpoint pointing at this server. + fill in the endpoint pointing at this server and inject a short-lived + internal sk-unsloth-* API key for workflow auth. + + Returns the row id of the minted internal key (so the caller can + revoke it on job completion) or ``None`` when no local provider is + actually reachable from an LLM column. """ providers = recipe.get("model_providers") if not providers: - return + return None # Collect local providers and pop is_local from ALL dicts unconditionally. # Strict `is True` guard so malformed payloads (is_local: 1, @@ -115,7 +213,7 @@ def _inject_local_providers(recipe: dict[str, Any], request: Request) -> None: local_indices.append(i) if not local_indices: - return + return None endpoint = _resolve_local_v1_endpoint(request) @@ -138,6 +236,7 @@ def _inject_local_providers(recipe: dict[str, Any], request: Request) -> None: } token = "" + internal_key_id: Optional[int] = None if local_names & referenced_providers: # Verify a model is loaded. # NOTE: This is a point-in-time check (TOCTOU). The model could be unloaded @@ -158,18 +257,21 @@ def _inject_local_providers(recipe: dict[str, Any], request: Request) -> None: "No model loaded in Chat. Load a model first, then run the recipe." ) - from auth.authentication import ( - create_access_token, - ) # deferred: avoids circular import + from auth import storage # deferred: avoids circular import - # Uses the "unsloth" admin subject. If the user changes their password, - # the JWT secret rotates and this token becomes invalid mid-run. - # Acceptable for v1 - recipes typically finish well within one session. - token = create_access_token( - subject = "unsloth", - expires_delta = timedelta(hours = 24), - desktop = _request_has_desktop_access_token(request), + # Mint an internal sk-unsloth-* key scoped to this workflow run. + # Uses the unified API-key issuance path (one mint/revoke/verify + # surface instead of a second JWT code path). The key is marked + # internal so it is hidden from the user's API-key list, and the + # caller revokes it when the job terminates. + expires_at = (datetime.now(timezone.utc) + timedelta(hours = 24)).isoformat() + token, row = storage.create_api_key( + username = "unsloth", + name = "data-recipe workflow", + expires_at = expires_at, + internal = True, ) + internal_key_id = int(row["id"]) # Defensively strip any stale "external"-only fields the frontend may # have left on the dict (extra_headers/extra_body/api_key_env). The UI @@ -196,6 +298,37 @@ def _inject_local_providers(recipe: dict[str, Any], request: Request) -> None: continue if mc.get("provider") in local_names: mc["skip_health_check"] = True + # Disable thinking for data-recipe inference on local providers. + # Reasoning models emit a ... preamble before the + # answer, which roughly doubles generated token count per row and + # pushes the visible answer past data_designer's json-fence + # regex. Forward chat_template_kwargs={enable_thinking: False} + # through the OpenAI SDK's extra_body passthrough so llama-server + # renders the template without the reasoning preamble. Free-form + # llm-text columns benefit from the latency cut, and structured + # columns also stop leaking think tags into the grammar- + # constrained JSON (llama-server's GBNF path still enforces the + # schema either way). + params = mc.get("inference_parameters") + if not isinstance(params, dict): + params = {} + mc["inference_parameters"] = params + extra_body = params.get("extra_body") + if not isinstance(extra_body, dict): + extra_body = {} + tpl_kwargs = extra_body.get("chat_template_kwargs") + if not isinstance(tpl_kwargs, dict): + tpl_kwargs = {} + tpl_kwargs.setdefault("enable_thinking", False) + extra_body["chat_template_kwargs"] = tpl_kwargs + params["extra_body"] = extra_body + + # Forward each llm-structured column's output_format as an OpenAI + # response_format so llama-server uses grammar-constrained sampling and + # small GGUFs stop wasting the full max_tokens budget on broken JSON. + _inject_local_structured_response_format(recipe, local_names) + + return internal_key_id def _normalize_run_name(value: Any) -> str | None: @@ -240,21 +373,49 @@ def create_job(payload: RecipePayload, request: Request): ) from exc try: - _inject_local_providers(recipe, request) + internal_api_key_id = _inject_local_providers(recipe, request) except ValueError as exc: raise HTTPException(status_code = 400, detail = str(exc)) from exc - mgr = get_job_manager() + # Single try block covers get_job_manager() AND mgr.start() so a workflow + # key minted above never outlives the request even when an unexpected + # exception type (TypeError from a stale kwarg, OSError from a queue + # write, etc.) bubbles up. Without the bare except, such exceptions let + # the sk-unsloth-* key live until its 24h TTL. try: - job_id = mgr.start(recipe = recipe, run = run) + mgr = get_job_manager() + job_id = mgr.start( + recipe = recipe, + run = run, + internal_api_key_id = internal_api_key_id, + ) except RuntimeError as exc: + if internal_api_key_id is not None: + _revoke_internal_api_key_safe(internal_api_key_id) raise HTTPException(status_code = 409, detail = str(exc)) from exc except ValueError as exc: + if internal_api_key_id is not None: + _revoke_internal_api_key_safe(internal_api_key_id) raise HTTPException(status_code = 400, detail = str(exc)) from exc + except Exception: + if internal_api_key_id is not None: + _revoke_internal_api_key_safe(internal_api_key_id) + raise return {"job_id": job_id} +def _revoke_internal_api_key_safe(key_id: int) -> None: + """Best-effort revoke of a workflow-minted key; swallow any error so + that revocation failures never mask the caller's own error path.""" + try: + from auth import storage # deferred: avoids circular import + + storage.revoke_internal_api_key(key_id) + except Exception: + pass + + @router.get("/jobs/{job_id}/status") def job_status(job_id: str): mgr = get_job_manager() diff --git a/studio/backend/routes/data_recipe/seed.py b/studio/backend/routes/data_recipe/seed.py index e9cf828610..91cf718e6e 100644 --- a/studio/backend/routes/data_recipe/seed.py +++ b/studio/backend/routes/data_recipe/seed.py @@ -8,6 +8,7 @@ from __future__ import annotations import base64 import binascii import json +import os import re from itertools import islice from pathlib import Path @@ -627,3 +628,14 @@ def inspect_seed_upload(payload: SeedInspectUploadRequest) -> SeedInspectRespons split = None, subset = None, ) + + +@router.get("/seed/github/env-token") +def get_github_env_token_status() -> dict: + """Report whether the server has a GH_TOKEN / GITHUB_TOKEN env var. + + The value is never returned; the UI uses this to tell the user they + can leave the token field blank. + """ + has_token = bool(os.environ.get("GH_TOKEN") or os.environ.get("GITHUB_TOKEN")) + return {"has_token": has_token} diff --git a/studio/backend/routes/data_recipe/validate.py b/studio/backend/routes/data_recipe/validate.py index 555e3eaa06..e794d68e54 100644 --- a/studio/backend/routes/data_recipe/validate.py +++ b/studio/backend/routes/data_recipe/validate.py @@ -14,10 +14,63 @@ from core.data_recipe.service import ( create_data_designer, validate_recipe, ) +from loggers import get_logger from models.data_recipe import RecipePayload, ValidateError, ValidateResponse +logger = get_logger(__name__) router = APIRouter() +_GITHUB_VALIDATE_NOTE = "Recipe shape is valid. GitHub access and rate limits are checked when the run starts." +_GITHUB_ITEM_TYPES = {"issues", "pulls", "commits"} + + +def _github_seed_source(recipe: dict[str, Any]) -> dict[str, Any] | None: + seed_config = recipe.get("seed_config") + if not isinstance(seed_config, dict): + return None + source = seed_config.get("source") + if not isinstance(source, dict) or source.get("seed_type") != "github_repo": + return None + return source + + +def _validate_github_seed_static(source: dict[str, Any]) -> list[ValidateError]: + errors: list[ValidateError] = [] + + repos = source.get("repos") + if not isinstance(repos, list) or not repos: + errors.append(ValidateError(message = "GitHub seed requires at least one repo.")) + else: + for repo in repos: + if not isinstance(repo, str) or not repo.strip() or "/" not in repo: + errors.append( + ValidateError(message = "GitHub repos must be owner/name strings.") + ) + break + + item_types = source.get("item_types") + if not isinstance(item_types, list) or not item_types: + errors.append( + ValidateError(message = "GitHub seed requires at least one item type.") + ) + else: + invalid_items = [item for item in item_types if item not in _GITHUB_ITEM_TYPES] + if invalid_items: + errors.append( + ValidateError( + message = "GitHub item types must be issues, pulls, or commits." + ) + ) + + try: + limit = int(source.get("limit")) + except (TypeError, ValueError): + limit = 0 + if limit < 1 or limit > 5000: + errors.append(ValidateError(message = "GitHub limit must be from 1 to 5000.")) + + return errors + def _collect_validation_errors(recipe: dict[str, Any]) -> list[ValidateError]: try: @@ -93,6 +146,38 @@ def validate(payload: RecipePayload) -> ValidateResponse: _patch_local_providers(recipe) + github_source = _github_seed_source(recipe) + if github_source is not None: + static_errors = _validate_github_seed_static(github_source) + if static_errors: + return ValidateResponse(valid = False, errors = static_errors) + try: + build_config_builder(recipe) + except ModuleNotFoundError as exc: + # data_designer is an optional runtime dep. Static validation + # already passed; live access + full config validation are + # deferred to run start (per _GITHUB_VALIDATE_NOTE), so a missing + # optional import at validate time should not block the recipe. + # Restrict the bypass to the data_designer module specifically so + # other ImportErrors (e.g. broken internal imports or missing + # transitive deps after a package upgrade) still surface as + # validation failures instead of being silently swallowed. + if not (exc.name or "").startswith("data_designer"): + raise + logger.debug( + "data_designer not installed; deferring full config " + "validation to run start", + missing_module = exc.name, + ) + except Exception as exc: + detail = str(exc).strip() or "Validation failed." + return ValidateResponse( + valid = False, + errors = [ValidateError(message = detail)], + raw_detail = detail, + ) + return ValidateResponse(valid = True, raw_detail = _GITHUB_VALIDATE_NOTE) + try: validate_recipe(recipe) except RuntimeError as exc: diff --git a/studio/backend/routes/inference.py b/studio/backend/routes/inference.py index ed331a5660..a6b00360af 100644 --- a/studio/backend/routes/inference.py +++ b/studio/backend/routes/inference.py @@ -113,19 +113,45 @@ if str(backend_path) not in sys.path: # Import backend functions try: from core.inference import get_inference_backend - from core.inference.llama_cpp import LlamaCppBackend, detect_reasoning_flags + from core.inference.llama_cpp import ( + LlamaCppBackend, + _DEFAULT_MAX_TOKENS_FLOOR, + _DEFAULT_T_MAX_PREDICT_MS, + detect_reasoning_flags, + ) + from core.inference.llama_server_args import validate_extra_args from utils.models import ModelConfig from utils.inference import load_inference_config from utils.models.model_config import load_model_defaults + from utils.native_path_leases import ( + NativePathLeaseError, + display_label_for_native_path, + is_registered_native_path_label, + redact_native_paths, + verify_native_path_lease, + ) except ImportError: parent_backend = backend_path.parent / "backend" if str(parent_backend) not in sys.path: sys.path.insert(0, str(parent_backend)) from core.inference import get_inference_backend - from core.inference.llama_cpp import LlamaCppBackend, detect_reasoning_flags + from core.inference.llama_cpp import ( + LlamaCppBackend, + _DEFAULT_MAX_TOKENS_FLOOR, + _DEFAULT_T_MAX_PREDICT_MS, + detect_reasoning_flags, + ) + from core.inference.llama_server_args import validate_extra_args from utils.models import ModelConfig from utils.inference import load_inference_config from utils.models.model_config import load_model_defaults + from utils.native_path_leases import ( + NativePathLeaseError, + display_label_for_native_path, + is_registered_native_path_label, + redact_native_paths, + verify_native_path_lease, + ) from models.inference import ( LoadRequest, @@ -185,6 +211,138 @@ import numpy as np from datetime import date as _date router = APIRouter() +# Studio-only router (not mounted on /v1 OpenAI-compat). +studio_router = APIRouter() + + +def _effective_enable_tools(payload) -> Optional[bool]: + """Resolve `payload.enable_tools` against the process-level tool policy. + + Returns the policy value when set (CLI hard-override from `unsloth run`), + otherwise the per-request value. + """ + from state.tool_policy import get_tool_policy + + policy = get_tool_policy() + return policy if policy is not None else payload.enable_tools + + +# Cancel registry. Proxies (e.g. Colab) can swallow client fetch aborts +# so is_disconnected() never fires. POST /inference/cancel looks up +# in-flight cancel_events here by cancel_id (per-run) or session_id / +# completion_id (fallbacks). +_CANCEL_REGISTRY: dict[str, set[threading.Event]] = {} +_CANCEL_LOCK = threading.Lock() + +# Cancel POSTs that arrive before registration are stashed; the next +# matching __enter__ replays set() within the TTL. +_PENDING_CANCELS: dict[str, float] = {} +_PENDING_CANCEL_TTL_S = 30.0 + + +def _prune_pending(now: float) -> None: + for k in [ + k for k, ts in _PENDING_CANCELS.items() if now - ts > _PENDING_CANCEL_TTL_S + ]: + _PENDING_CANCELS.pop(k, None) + + +class _TrackedCancel: + """Register cancel_event in _CANCEL_REGISTRY for the block's duration.""" + + def __init__(self, event: threading.Event, *keys): + self.event = event + self.keys = tuple(k for k in keys if k) + + def __enter__(self): + # Register + consume-pending must be one critical section to close + # the TOCTOU race against a concurrent cancel POST. + should_cancel = False + with _CANCEL_LOCK: + for k in self.keys: + _CANCEL_REGISTRY.setdefault(k, set()).add(self.event) + now = time.monotonic() + _prune_pending(now) + for k in self.keys: + if k and _PENDING_CANCELS.pop(k, None) is not None: + should_cancel = True + if should_cancel: + self.event.set() + return self.event + + def __exit__(self, *exc): + with _CANCEL_LOCK: + for k in self.keys: + bucket = _CANCEL_REGISTRY.get(k) + if bucket is None: + continue + bucket.discard(self.event) + if not bucket: + _CANCEL_REGISTRY.pop(k, None) + return False + + +def _cancel_by_keys(keys) -> int: + """Set cancel_event for matching registry entries; no stash. + session_id/completion_id are shared across runs on the same thread, + so stashing them would ghost-cancel the user's next request. Only + cancel_id is per-run unique (see _cancel_by_cancel_id_or_stash).""" + if not keys: + return 0 + events: set[threading.Event] = set() + with _CANCEL_LOCK: + _prune_pending(time.monotonic()) + for k in keys: + bucket = _CANCEL_REGISTRY.get(k) + if bucket: + events.update(bucket) + for ev in events: + ev.set() + return len(events) + + +def _cancel_by_cancel_id_or_stash(cancel_id: str) -> int: + """Atomic lookup-or-stash; pairs with _TrackedCancel.__enter__ to + close the TOCTOU race.""" + now = time.monotonic() + events: set[threading.Event] = set() + with _CANCEL_LOCK: + _prune_pending(now) + bucket = _CANCEL_REGISTRY.get(cancel_id) + if bucket: + events.update(bucket) + else: + _PENDING_CANCELS[cancel_id] = now + for ev in events: + ev.set() + return len(events) + + +async def _await_cancel_then_close(cancel_event, resp) -> None: + """Watch a threading.Event from asyncio and close ``resp`` when it fires. + + Used by the passthrough streamers so a /cancel POST can interrupt + while the async iterator is blocked waiting for llama-server prefill. + Without this watcher the in-loop ``cancel_event.is_set()`` check is + unreachable until the first SSE chunk arrives, which is exactly the + proxy/Colab scenario the cancel POST exists to handle. + + Polls a threading.Event because the cancel registry is keyed by + threading.Event so the synchronous /cancel handler can call .set(). + 50ms cadence adds at most that much latency to a prefill cancel; the + common-case streaming cancel path still observes the event in the + iterator's first iteration after the next chunk. + """ + try: + while not cancel_event.is_set(): + await asyncio.sleep(0.05) + try: + await resp.aclose() + except Exception: + pass + except asyncio.CancelledError: + return + # Appended to tool-use nudge to discourage plan-without-action _TOOL_ACTION_NUDGE = ( @@ -202,6 +360,65 @@ _TOOL_XML_RE = _re.compile( logger = get_logger(__name__) +def _validate_native_mmproj_companion( + mmproj_path: str | None, gguf_path: str | None +) -> None: + if not mmproj_path or not gguf_path: + return + import stat as _stat_module + + mm = Path(mmproj_path) + gguf = Path(gguf_path) + try: + mm_lstat = os.lstat(mm) + except OSError as exc: + raise HTTPException( + status_code = 400, + detail = "Native vision companion is no longer accessible.", + ) from exc + if _stat_module.S_ISLNK(mm_lstat.st_mode) or not _stat_module.S_ISREG( + mm_lstat.st_mode + ): + raise HTTPException( + status_code = 400, + detail = "Native vision companion must be a regular file.", + ) + try: + if mm.resolve(strict = True).parent != gguf.resolve(strict = True).parent: + raise HTTPException( + status_code = 400, + detail = "Native vision companion must live next to the selected GGUF.", + ) + except OSError as exc: + raise HTTPException( + status_code = 400, + detail = "Native vision companion is no longer accessible.", + ) from exc + + +def _resolve_model_identifier_for_request( + request: LoadRequest | ValidateModelRequest, + *, + operation: str, +) -> tuple[str, str, bool]: + if not request.native_path_lease: + return request.model_path, request.model_path, False + try: + grant = verify_native_path_lease( + request.native_path_lease, + operation = operation, + expected_kind = "model", + expected_path_type = "file", + allowed_suffixes = (".gguf",), + ) + except NativePathLeaseError as exc: + raise HTTPException(status_code = 400, detail = str(exc)) from exc + display_label = ( + grant.display_label or Path(request.model_path).name or "Native model" + ) + return str(grant.canonical_path), display_label, True + + # GGUF inference backend (llama-server) _llama_cpp_backend = LlamaCppBackend() @@ -225,7 +442,19 @@ async def load_model( GGUF models are loaded via llama-server (llama.cpp) instead of Unsloth. """ + native_grant_backed = False + model_log_label = request.model_path try: + # Validate user-supplied llama-server pass-through args up front + # so a managed-flag collision returns 400 before any model work. + try: + extra_llama_args = validate_extra_args(request.llama_extra_args) + except ValueError as exc: + raise HTTPException(status_code = 400, detail = str(exc)) + + model_identifier, model_log_label, native_grant_backed = ( + _resolve_model_identifier_for_request(request, operation = "load-model") + ) # Version switching is handled automatically by the subprocess-based # inference backend — no need for ensure_transformers_version() here. @@ -239,10 +468,10 @@ async def load_model( and llama_backend.hf_variant and llama_backend.hf_variant.lower() == request.gguf_variant.lower() and llama_backend.model_identifier - and llama_backend.model_identifier.lower() == request.model_path.lower() + and llama_backend.model_identifier.lower() == model_identifier.lower() ): logger.info( - f"Model already loaded (GGUF): {request.model_path} variant={request.gguf_variant}, skipping reload" + f"Model already loaded (GGUF): {model_log_label} variant={request.gguf_variant}, skipping reload" ) inference_config = load_inference_config(llama_backend.model_identifier) from utils.models import is_audio_input_type @@ -255,8 +484,12 @@ async def load_model( _gguf_is_audio = getattr(llama_backend, "_is_audio", False) return LoadResponse( status = "already_loaded", - model = llama_backend.model_identifier, - display_name = llama_backend.model_identifier, + model = model_log_label + if native_grant_backed + else llama_backend.model_identifier, + display_name = model_log_label + if native_grant_backed + else llama_backend.model_identifier, is_vision = llama_backend._is_vision, is_lora = False, is_gguf = True, @@ -282,10 +515,10 @@ async def load_model( else: if ( backend.active_model_name - and backend.active_model_name.lower() == request.model_path.lower() + and backend.active_model_name.lower() == model_identifier.lower() ): logger.info( - f"Model already loaded (Unsloth): {request.model_path}, skipping reload" + f"Model already loaded (Unsloth): {model_log_label}, skipping reload" ) inference_config = load_inference_config(backend.active_model_name) _model_info = backend.models.get(backend.active_model_name, {}) @@ -314,8 +547,12 @@ async def load_model( pass return LoadResponse( status = "already_loaded", - model = backend.active_model_name, - display_name = backend.active_model_name, + model = model_log_label + if native_grant_backed + else backend.active_model_name, + display_name = model_log_label + if native_grant_backed + else backend.active_model_name, is_vision = _model_info.get("is_vision", False), is_lora = _model_info.get("is_lora", False), is_gguf = False, @@ -337,7 +574,7 @@ async def load_model( # Create config using clean factory method # is_lora is auto-detected from adapter_config.json on disk/HF config = ModelConfig.from_identifier( - model_id = request.model_path, + model_id = model_identifier, hf_token = request.hf_token, gguf_variant = request.gguf_variant, ) @@ -345,7 +582,7 @@ async def load_model( if not config: raise HTTPException( status_code = 400, - detail = f"Invalid model identifier: {request.model_path}", + detail = f"Invalid model identifier: {model_log_label}", ) # Normalize gpu_ids: empty list means auto-selection, same as None @@ -389,9 +626,14 @@ async def load_model( cache_type_kv = request.cache_type_kv, speculative_type = request.speculative_type, n_parallel = _n_parallel, + extra_args = extra_llama_args, ) else: # Local mode: llama-server loads via -m + if native_grant_backed and config.gguf_mmproj_file: + _validate_native_mmproj_companion( + config.gguf_mmproj_file, config.gguf_file + ) success = await asyncio.to_thread( llama_backend.load_model, gguf_path = config.gguf_file, @@ -403,15 +645,18 @@ async def load_model( cache_type_kv = request.cache_type_kv, speculative_type = request.speculative_type, n_parallel = _n_parallel, + extra_args = extra_llama_args, ) if not success: raise HTTPException( status_code = 500, - detail = f"Failed to load GGUF model: {config.display_name}", + detail = f"Failed to load GGUF model: {model_log_label if native_grant_backed else config.display_name}", ) - logger.info(f"Loaded GGUF model via llama-server: {config.identifier}") + logger.info( + f"Loaded GGUF model via llama-server: {model_log_label if native_grant_backed else config.identifier}" + ) # Detect TTS audio by probing the loaded model's vocabulary from utils.models import is_audio_input_type @@ -420,6 +665,10 @@ async def load_model( _gguf_is_audio = _gguf_audio in ("snac", "bicodec", "dac") llama_backend._is_audio = _gguf_is_audio llama_backend._audio_type = _gguf_audio + llama_backend._native_display_label = ( + model_log_label if native_grant_backed else None + ) + llama_backend._native_grant_backed = bool(native_grant_backed) if _gguf_is_audio: logger.info(f"GGUF model detected as audio: audio_type={_gguf_audio}") await asyncio.to_thread(llama_backend.init_audio_codec, _gguf_audio) @@ -428,8 +677,10 @@ async def load_model( return LoadResponse( status = "loaded", - model = config.identifier, - display_name = config.display_name, + model = model_log_label if native_grant_backed else config.identifier, + display_name = model_log_label + if native_grant_backed + else config.display_name, is_vision = config.is_vision, is_lora = False, is_gguf = True, @@ -552,10 +803,13 @@ async def load_model( ), ) raise HTTPException( - status_code = 500, detail = f"Failed to load model: {config.display_name}" + status_code = 500, + detail = f"Failed to load model: {model_log_label if native_grant_backed else config.display_name}", ) - logger.info(f"Loaded model: {config.identifier}") + logger.info( + f"Loaded model: {model_log_label if native_grant_backed else config.identifier}" + ) # Load inference configuration parameters inference_config = load_inference_config(config.identifier) @@ -585,8 +839,10 @@ async def load_model( return LoadResponse( status = "loaded", - model = config.identifier, - display_name = config.display_name, + model = model_log_label if native_grant_backed else config.identifier, + display_name = model_log_label + if native_grant_backed + else config.display_name, is_vision = config.is_vision, is_lora = config.is_lora, is_gguf = False, @@ -608,11 +864,17 @@ async def load_model( except HTTPException: raise except ValueError as e: + if native_grant_backed: + redacted_msg = redact_native_paths(str(e)) + logger.warning( + "Rejected inference selection for native model %s: %s", + model_log_label, + redacted_msg, + ) + raise HTTPException(status_code = 400, detail = redacted_msg) logger.warning("Rejected inference GPU selection: %s", e) raise HTTPException(status_code = 400, detail = str(e)) except Exception as e: - logger.error(f"Error loading model: {e}", exc_info = True) - msg = str(e) # Surface a friendlier message for models that Unsloth cannot load not_supported_hints = [ "No config file found", @@ -620,6 +882,22 @@ async def load_model( "is not supported", "does not support", ] + if native_grant_backed: + redacted_msg = redact_native_paths(str(e)) + logger.error( + "Error loading native model %s: %s", + model_log_label, + redacted_msg, + ) + msg = redacted_msg + if any(h.lower() in msg.lower() for h in not_supported_hints): + msg = f"This model is not supported yet. Try a different model. (Original error: {msg})" + raise HTTPException( + status_code = 500, + detail = f"Failed to load native model {model_log_label}: {msg}", + ) + logger.error(f"Error loading model: {e}", exc_info = True) + msg = str(e) if any(h.lower() in msg.lower() for h in not_supported_hints): msg = f"This model is not supported yet. Try a different model. (Original error: {msg})" raise HTTPException(status_code = 500, detail = f"Failed to load model: {msg}") @@ -636,9 +914,14 @@ async def validate_model( This checks that ModelConfig.from_identifier() can resolve the given model_path, but it does NOT actually load model weights into GPU memory. """ + native_grant_backed = False + model_log_label = request.model_path try: + model_identifier, model_log_label, native_grant_backed = ( + _resolve_model_identifier_for_request(request, operation = "validate-model") + ) config = ModelConfig.from_identifier( - model_id = request.model_path, + model_id = model_identifier, hf_token = request.hf_token, gguf_variant = request.gguf_variant, ) @@ -646,14 +929,16 @@ async def validate_model( if not config: raise HTTPException( status_code = 400, - detail = f"Invalid model identifier: {request.model_path}", + detail = f"Invalid model identifier: {model_log_label}", ) return ValidateModelResponse( valid = True, message = "Model identifier is valid.", - identifier = config.identifier, - display_name = getattr(config, "display_name", config.identifier), + identifier = model_log_label if native_grant_backed else config.identifier, + display_name = model_log_label + if native_grant_backed + else getattr(config, "display_name", config.identifier), is_gguf = getattr(config, "is_gguf", False), is_lora = getattr(config, "is_lora", False), is_vision = getattr(config, "is_vision", False), @@ -665,6 +950,26 @@ async def validate_model( except HTTPException: raise except Exception as e: + not_supported_hints = [ + "No config file found", + "not yet supported", + "is not supported", + "does not support", + ] + if native_grant_backed: + redacted_msg = redact_native_paths(str(e)) + logger.error( + "Error validating native model %s: %s", + model_log_label, + redacted_msg, + ) + msg = redacted_msg + if any(h.lower() in msg.lower() for h in not_supported_hints): + msg = f"This model is not supported yet. Try a different model. (Original error: {msg})" + raise HTTPException( + status_code = 400, + detail = f"Invalid native model {model_log_label}: {msg}", + ) logger.error( f"Error validating model identifier '{request.model_path}': {e}", exc_info = True, @@ -689,6 +994,9 @@ async def unload_model( llama_backend = get_llama_cpp_backend() if llama_backend.is_active and ( llama_backend.model_identifier == request.model_path + or is_registered_native_path_label( + llama_backend.model_identifier, request.model_path + ) or not llama_backend.is_loaded ): llama_backend.unload_model() @@ -706,6 +1014,48 @@ async def unload_model( raise HTTPException(status_code = 500, detail = f"Failed to unload model: {str(e)}") +@studio_router.post("/cancel") +async def cancel_inference( + request: Request, + current_subject: str = Depends(get_current_subject), +): + """Cancel in-flight inference requests. + + Body (JSON, at least one key required): + cancel_id - preferred: per-run UUID, matched exclusively. + session_id - fallback when cancel_id is absent. + completion_id - fallback when cancel_id is absent. + + A cancel_id arriving before its stream registers is stashed briefly + and replayed on registration. Returns {"cancelled": N}. + """ + try: + body = await request.json() + if not isinstance(body, dict): + body = {} + except Exception as e: + logger.debug("Failed to parse cancel request body: %s", e) + body = {} + + cancel_id = body.get("cancel_id") + if isinstance(cancel_id, str) and cancel_id: + return {"cancelled": _cancel_by_cancel_id_or_stash(cancel_id)} + + keys = [] + # `message_id` is the Anthropic passthrough's per-run identifier -- + # included so /v1/messages clients can cancel by their native id. + for k in ("completion_id", "session_id", "message_id"): + v = body.get(k) + if isinstance(v, str) and v: + keys.append(v) + + if not keys: + return {"cancelled": 0} + + n = _cancel_by_keys(keys) + return {"cancelled": n} + + @router.post("/generate/stream") async def generate_stream( request: GenerateRequest, @@ -794,16 +1144,27 @@ async def get_status( # If a GGUF model is loaded via llama-server, report that if llama_backend.is_loaded: _model_id = llama_backend.model_identifier + _native_grant_backed = getattr(llama_backend, "_native_grant_backed", False) + _display_model_id = getattr( + llama_backend, "_native_display_label", None + ) or display_label_for_native_path(_model_id) + if ( + _native_grant_backed + and _model_id + and _display_model_id == _model_id + and os.path.isabs(_model_id) + ): + _display_model_id = os.path.basename(_model_id) _inference_cfg = load_inference_config(_model_id) if _model_id else None return InferenceStatusResponse( - active_model = _model_id, + active_model = _display_model_id, is_vision = llama_backend.is_vision, is_gguf = True, gguf_variant = llama_backend.hf_variant, is_audio = getattr(llama_backend, "_is_audio", False), audio_type = getattr(llama_backend, "_audio_type", None), loading = [], - loaded = [_model_id], + loaded = [_display_model_id] if _display_model_id else [], inference = _inference_cfg, requires_trust_remote_code = bool( (_inference_cfg or {}).get("trust_remote_code", False) @@ -813,6 +1174,7 @@ async def get_status( reasoning_always_on = llama_backend.reasoning_always_on, supports_preserve_thinking = llama_backend.supports_preserve_thinking, supports_tools = llama_backend.supports_tools, + chat_template = llama_backend.chat_template, context_length = llama_backend.context_length, max_context_length = llama_backend.max_context_length, native_context_length = llama_backend.native_context_length, @@ -826,12 +1188,19 @@ async def get_status( is_audio = False audio_type = None has_audio_input = False + model_info = {} if backend.active_model_name: model_info = backend.models.get(backend.active_model_name, {}) is_vision = model_info.get("is_vision", False) is_audio = model_info.get("is_audio", False) audio_type = model_info.get("audio_type") has_audio_input = model_info.get("has_audio_input", False) + chat_template_info = model_info.get("chat_template_info", {}) + chat_template = ( + chat_template_info.get("template") + if isinstance(chat_template_info, dict) + else None + ) # Non-GGUF: only gpt-oss Harmony is wired through the transformers # generation path. Other template-level reasoning / tool kwargs @@ -869,6 +1238,7 @@ async def get_status( reasoning_always_on = False, supports_preserve_thinking = False, supports_tools = False, + chat_template = chat_template, ) except Exception as e: @@ -1114,6 +1484,20 @@ async def openai_chat_completions( llama_backend = get_llama_cpp_backend() using_gguf = llama_backend.is_loaded + # OpenAI-SDK clients send ``chat_template_kwargs`` via ``extra_body``, + # which the SDK spreads into the request body at the top level. Studio's + # ChatCompletionRequest has ``extra="allow"`` so pydantic stashes them in + # ``model_extra``, but the typed ``payload.enable_thinking`` path is what + # downstream generators actually consume. Lift ``enable_thinking`` from + # the extra-body chat_template_kwargs onto the typed field so clients + # that only know the OpenAI shape (data_designer recipe runs, etc.) + # can still control the reasoning preamble. + _extra = getattr(payload, "model_extra", None) + if payload.enable_thinking is None and isinstance(_extra, dict): + _tpl_kw = _extra.get("chat_template_kwargs") + if isinstance(_tpl_kw, dict) and "enable_thinking" in _tpl_kw: + payload.enable_thinking = bool(_tpl_kw["enable_thinking"]) + # ── Determine which backend is active ───────────────────── if using_gguf: model_name = llama_backend.model_identifier or payload.model @@ -1169,6 +1553,9 @@ async def openai_chat_completions( ) if payload.stream: + _cancel_keys = (payload.cancel_id, payload.session_id, completion_id) + _tracker = _TrackedCancel(cancel_event, *_cancel_keys) + _tracker.__enter__() async def audio_input_stream(): try: @@ -1185,10 +1572,17 @@ async def openai_chat_completions( ) yield f"data: {first_chunk.model_dump_json(exclude_none = True)}\n\n" - for chunk_text in audio_input_generate(): + gen = audio_input_generate() + _DONE = object() + while True: + if cancel_event.is_set(): + break if await request.is_disconnected(): cancel_event.set() return + chunk_text = await asyncio.to_thread(next, gen, _DONE) + if chunk_text is _DONE: + break if chunk_text: chunk = ChatCompletionChunk( id = completion_id, @@ -1221,6 +1615,8 @@ async def openai_chat_completions( f"Error during audio input streaming: {e}", exc_info = True ) yield f"data: {json.dumps({'error': {'message': _friendly_error(e), 'type': 'server_error'}})}\n\n" + finally: + _tracker.__exit__(None, None, None) return StreamingResponse( audio_input_stream(), @@ -1256,11 +1652,22 @@ async def openai_chat_completions( # carry `tool_calls` (content=None) — both of which are valid in # multi-turn client-side tool loops. _has_tool_messages = any(m.role == "tool" or m.tool_calls for m in payload.messages) + # Route guided-decoding requests through the verbatim passthrough so + # ``response_format`` (JSON schema) actually reaches llama-server and + # the model's GBNF-constrained output comes back unmodified. The + # non-passthrough GGUF path below calls ``generate_chat_completion`` + # which has no response_format kwarg, so the schema gets silently + # dropped and data_designer falls back to free-form sampling. Guided + # decoding does not require ``supports_tools`` - the grammar machinery + # is independent of tool-call parsing. + _has_response_format = _extract_response_format(payload) is not None + _tools_passthrough = llama_backend.supports_tools and ( + (payload.tools and len(payload.tools) > 0) or _has_tool_messages + ) if ( using_gguf - and llama_backend.supports_tools - and not payload.enable_tools - and ((payload.tools and len(payload.tools) > 0) or _has_tool_messages) + and not _effective_enable_tools(payload) + and (_tools_passthrough or _has_response_format) ): # Preserve the vision guard that would otherwise run in the # non-passthrough path below: text-only tool-capable GGUFs @@ -1350,8 +1757,13 @@ async def openai_chat_completions( created = int(time.time()) # ── Tool-calling path (agentic loop) ────────────────── + # `_effective_enable_tools` lets `unsloth run --enable-tools/--disable-tools` + # hard-override the per-request value. Without a CLI override, falls + # back to `payload.enable_tools` (existing behavior). use_tools = ( - payload.enable_tools and llama_backend.supports_tools and not image_b64 + _effective_enable_tools(payload) + and llama_backend.supports_tools + and not image_b64 ) if use_tools: @@ -1466,6 +1878,10 @@ async def openai_chat_completions( _tool_sentinel = object() + _cancel_keys = (payload.cancel_id, payload.session_id, completion_id) + _tracker = _TrackedCancel(cancel_event, *_cancel_keys) + _tracker.__enter__() + async def gguf_tool_stream(): try: first_chunk = ChatCompletionChunk( @@ -1488,6 +1904,8 @@ async def openai_chat_completions( _stream_usage = None _stream_timings = None while True: + if cancel_event.is_set(): + break if await request.is_disconnected(): cancel_event.set() return @@ -1595,6 +2013,8 @@ async def openai_chat_completions( }, } yield f"data: {json.dumps(error_chunk)}\n\n" + finally: + _tracker.__exit__(None, None, None) return StreamingResponse( gguf_tool_stream(), @@ -1628,6 +2048,9 @@ async def openai_chat_completions( _gguf_sentinel = object() if payload.stream: + _cancel_keys = (payload.cancel_id, payload.session_id, completion_id) + _tracker = _TrackedCancel(cancel_event, *_cancel_keys) + _tracker.__enter__() async def gguf_stream_chunks(): try: @@ -1652,6 +2075,8 @@ async def openai_chat_completions( _stream_usage = None _stream_timings = None while True: + if cancel_event.is_set(): + break if await request.is_disconnected(): cancel_event.set() return @@ -1735,6 +2160,8 @@ async def openai_chat_completions( }, } yield f"data: {json.dumps(error_chunk)}\n\n" + finally: + _tracker.__exit__(None, None, None) return StreamingResponse( gguf_stream_chunks(), @@ -1834,6 +2261,9 @@ async def openai_chat_completions( # ── Streaming response ──────────────────────────────────────── if payload.stream: + _cancel_keys = (payload.cancel_id, payload.session_id, completion_id) + _tracker = _TrackedCancel(cancel_event, *_cancel_keys) + _tracker.__enter__() async def stream_chunks(): try: @@ -1861,6 +2291,9 @@ async def openai_chat_completions( loop = asyncio.get_event_loop() gen = generate() while True: + if cancel_event.is_set(): + backend.reset_generation_state() + break # next(gen, _DONE) returns _DONE instead of raising # StopIteration — StopIteration cannot propagate # through asyncio futures (Python limitation). @@ -1916,6 +2349,8 @@ async def openai_chat_completions( }, } yield f"data: {json.dumps(error_chunk)}\n\n" + finally: + _tracker.__exit__(None, None, None) return StreamingResponse( stream_chunks(), @@ -2596,7 +3031,9 @@ async def _responses_stream( ), ) - body = _build_openai_passthrough_body(chat_req) + body = _build_openai_passthrough_body( + chat_req, backend_ctx = llama_backend.context_length + ) target_url = f"{llama_backend.base_url}/v1/chat/completions" async def event_generator(): @@ -3050,7 +3487,9 @@ async def anthropic_messages( # Server-side agentic loop doesn't support multimodal input — matches # the `not image_b64` gate in /v1/chat/completions. server_tools = ( - payload.enable_tools and llama_backend.supports_tools and not _has_image + _effective_enable_tools(payload) + and llama_backend.supports_tools + and not _has_image ) client_tools = ( not server_tools @@ -3081,6 +3520,8 @@ async def anthropic_messages( repetition_penalty = repetition_penalty, presence_penalty = presence_penalty, tool_choice = openai_tool_choice, + session_id = payload.session_id, + cancel_id = payload.cancel_id, ) return await _anthropic_passthrough_non_streaming( llama_backend, @@ -3441,6 +3882,9 @@ def _build_passthrough_payload( repetition_penalty = None, presence_penalty = None, tool_choice = "auto", + response_format = None, + chat_template_kwargs = None, + backend_ctx = None, ): body = { "messages": openai_messages, @@ -3453,8 +3897,12 @@ def _build_passthrough_payload( } if stream: body["stream_options"] = {"include_usage": True} - if max_tokens is not None: - body["max_tokens"] = max_tokens + body["max_tokens"] = ( + max_tokens + if max_tokens is not None + else (backend_ctx or _DEFAULT_MAX_TOKENS_FLOOR) + ) + body["t_max_predict_ms"] = _DEFAULT_T_MAX_PREDICT_MS if stop: body["stop"] = stop if min_p is not None: @@ -3464,6 +3912,17 @@ def _build_passthrough_payload( body["repeat_penalty"] = repetition_penalty if presence_penalty is not None: body["presence_penalty"] = presence_penalty + if response_format is not None: + # llama-server applies a GBNF grammar derived from the JSON schema + # when response_format is present. Field is documented flat at the + # request root (tools/server/README.md), which is also what the + # OpenAI SDK produces by spreading extra_body into the body top. + body["response_format"] = response_format + if chat_template_kwargs is not None: + # Propagate reasoning / template overrides (e.g. enable_thinking) + # so llama-server renders the Jinja template in the mode the caller + # asked for instead of whatever default the model was loaded with. + body["chat_template_kwargs"] = chat_template_kwargs return body @@ -3484,6 +3943,8 @@ async def _anthropic_passthrough_stream( repetition_penalty = None, presence_penalty = None, tool_choice = "auto", + session_id = None, + cancel_id = None, ): """Streaming client-side pass-through: forward tools to llama-server and translate its streaming response to Anthropic SSE without executing anything.""" @@ -3501,8 +3962,14 @@ async def _anthropic_passthrough_stream( repetition_penalty = repetition_penalty, presence_penalty = presence_penalty, tool_choice = tool_choice, + backend_ctx = llama_backend.context_length, ) + # cancel_id mirrors the OpenAI passthrough so a per-run cancel POST + # works without the caller having to know the local message_id. + _tracker = _TrackedCancel(cancel_event, cancel_id, session_id, message_id) + _tracker.__enter__() + async def _stream(): emitter = AnthropicPassthroughEmitter() for line in emitter.start(message_id, model_name): @@ -3535,15 +4002,28 @@ async def _anthropic_passthrough_stream( # has anything orphaned to finalize. Each aclose is wrapped in # `try: ... except Exception: pass` so anyio cleanup noise from # nested aclose paths can't bubble out. - client = httpx.AsyncClient(timeout = 600) + client = httpx.AsyncClient( + timeout = 600, + limits = httpx.Limits(max_keepalive_connections = 0), + ) resp = None lines_iter = None + cancel_watcher = None try: req = client.build_request("POST", target_url, json = body) resp = await client.send(req, stream = True) + # See _openai_passthrough_stream for rationale: aiter_lines() + # blocks during llama-server prefill, so the in-loop cancel + # check is unreachable until the first SSE chunk arrives. + # The watcher closes `resp` on cancel, raising in aiter_lines. + cancel_watcher = asyncio.create_task( + _await_cancel_then_close(cancel_event, resp) + ) lines_iter = resp.aiter_lines() async for raw_line in lines_iter: + if cancel_event.is_set(): + break if await request.is_disconnected(): cancel_event.set() break @@ -3558,9 +4038,18 @@ async def _anthropic_passthrough_stream( continue for line in emitter.feed_chunk(chunk): yield line + except (httpx.RemoteProtocolError, httpx.ReadError, httpx.CloseError): + if not cancel_event.is_set(): + raise except Exception as e: logger.error("anthropic_messages passthrough stream error: %s", e) finally: + if cancel_watcher is not None: + cancel_watcher.cancel() + try: + await cancel_watcher + except (asyncio.CancelledError, Exception): + pass if lines_iter is not None: try: await lines_iter.aclose() @@ -3575,6 +4064,7 @@ async def _anthropic_passthrough_stream( await client.aclose() except Exception: pass + _tracker.__exit__(None, None, None) for line in emitter.finish(): yield line @@ -3621,6 +4111,7 @@ async def _anthropic_passthrough_non_streaming( repetition_penalty = repetition_penalty, presence_penalty = presence_penalty, tool_choice = tool_choice, + backend_ctx = llama_backend.context_length, ) async with httpx.AsyncClient() as client: @@ -3742,7 +4233,21 @@ def _openai_messages_for_passthrough(payload) -> list[dict]: return messages -def _build_openai_passthrough_body(payload) -> dict: +def _extract_response_format(payload): + """Return the ``response_format`` field on an incoming ChatCompletionRequest + (or None). The model is declared with ``extra="allow"`` so pydantic stashes + unknown top-level fields in ``model_extra``; OpenAI-SDK clients spread + ``extra_body`` into the request body top level, which is where guided- + decoding recipes park their JSON-schema response_format. + """ + extra = getattr(payload, "model_extra", None) + if not isinstance(extra, dict): + return None + rf = extra.get("response_format") + return rf if isinstance(rf, dict) else None + + +def _build_openai_passthrough_body(payload, backend_ctx = None) -> dict: """Assemble the llama-server request body from a ChatCompletionRequest. Only explicitly-known OpenAI / llama-server fields are forwarded so that @@ -3751,6 +4256,12 @@ def _build_openai_passthrough_body(payload) -> dict: """ messages = _openai_messages_for_passthrough(payload) tool_choice = payload.tool_choice if payload.tool_choice is not None else "auto" + # When the caller asked for a specific reasoning mode, forward it to + # llama-server via chat_template_kwargs so the Jinja template renders + # with (or without) the reasoning preamble. + tpl_kwargs = None + if payload.enable_thinking is not None: + tpl_kwargs = {"enable_thinking": bool(payload.enable_thinking)} return _build_passthrough_payload( messages, payload.tools, @@ -3764,6 +4275,9 @@ def _build_openai_passthrough_body(payload) -> dict: repetition_penalty = payload.repetition_penalty, presence_penalty = payload.presence_penalty, tool_choice = tool_choice, + response_format = _extract_response_format(payload), + chat_template_kwargs = tpl_kwargs, + backend_ctx = backend_ctx, ) @@ -3784,103 +4298,56 @@ async def _openai_passthrough_stream( observes a standard OpenAI response. """ target_url = f"{llama_backend.base_url}/v1/chat/completions" - body = _build_openai_passthrough_body(payload) + body = _build_openai_passthrough_body( + payload, backend_ctx = llama_backend.context_length + ) - # Dispatch the upstream request BEFORE returning StreamingResponse so - # transport errors and non-200 upstream statuses surface as real HTTP - # errors to the client. OpenAI SDKs rely on status codes to raise - # ``APIError``/``BadRequestError``/...; burying the failure inside a - # 200 SSE ``error`` frame silently breaks their error handling. - client = httpx.AsyncClient(timeout = 600) - resp = None + _cancel_keys = (payload.cancel_id, payload.session_id, completion_id) + _tracker = _TrackedCancel(cancel_event, *_cancel_keys) + _tracker.__enter__() + + # Outer guard: asyncio.CancelledError at `await client.send(...)` is + # a BaseException that bypasses `except httpx.RequestError`; without + # this the tracker leaks. The generator's finally only runs once + # iteration starts. try: - req = client.build_request("POST", target_url, json = body) - resp = await client.send(req, stream = True) - except httpx.RequestError as e: - # llama-server subprocess crashed / still starting / unreachable. - logger.error("openai passthrough stream: upstream unreachable: %s", e) - if resp is not None: - try: - await resp.aclose() - except Exception: - pass - try: - await client.aclose() - except Exception: - pass - raise HTTPException( - status_code = 502, - detail = _friendly_error(e), + # Dispatch BEFORE returning StreamingResponse so transport errors + # and non-200 upstream statuses surface as real HTTP errors -- + # OpenAI SDKs rely on status codes to raise APIError/BadRequestError. + client = httpx.AsyncClient( + timeout = 600, + limits = httpx.Limits(max_keepalive_connections = 0), ) - - if resp.status_code != 200: - err_bytes = await resp.aread() - err_text = err_bytes.decode("utf-8", errors = "replace") - logger.error( - "openai passthrough upstream error: status=%s body=%s", - resp.status_code, - err_text[:500], - ) - upstream_status = resp.status_code + resp = None try: - await resp.aclose() - except Exception: - pass - try: - await client.aclose() - except Exception: - pass - raise HTTPException( - status_code = upstream_status, - detail = f"llama-server error: {err_text[:500]}", - ) - - async def _stream(): - # Same httpx lifecycle pattern as _anthropic_passthrough_stream: - # avoid `async with` on the client/response AND explicitly save - # resp.aiter_lines() so we can close it ourselves in the finally - # block. See the long comment there for the full rationale on - # why the anonymous `async for raw_line in resp.aiter_lines():` - # pattern leaks an unclosed async generator that Python's - # asyncgen GC hook then finalizes in a different asyncio task, - # producing "Exception ignored in:" / "async generator ignored - # GeneratorExit" / anyio cancel-scope traces on Python 3.13 + - # httpcore 1.0.x. - lines_iter = None - try: - lines_iter = resp.aiter_lines() - async for raw_line in lines_iter: - if await request.is_disconnected(): - cancel_event.set() - break - if not raw_line: - continue - if not raw_line.startswith("data: "): - continue - # Relay the llama-server SSE chunk verbatim so the client - # sees its native `id`, `finish_reason`, `delta.tool_calls`, - # and final `usage` unchanged. - yield raw_line + "\n\n" - if raw_line[6:].strip() == "[DONE]": - break - except Exception as e: - # Mid-stream failures still have to be reported inside the SSE - # body because the 200 response headers have already been - # committed by the time the first chunk flushes. - logger.error("openai passthrough stream error: %s", e) - err = { - "error": { - "message": _friendly_error(e), - "type": "server_error", - }, - } - yield f"data: {json.dumps(err)}\n\n" - finally: - if lines_iter is not None: + req = client.build_request("POST", target_url, json = body) + resp = await client.send(req, stream = True) + except httpx.RequestError as e: + # llama-server subprocess crashed / still starting / unreachable. + logger.error("openai passthrough stream: upstream unreachable: %s", e) + if resp is not None: try: - await lines_iter.aclose() + await resp.aclose() except Exception: pass + try: + await client.aclose() + except Exception: + pass + raise HTTPException( + status_code = 502, + detail = _friendly_error(e), + ) + + if resp.status_code != 200: + err_bytes = await resp.aread() + err_text = err_bytes.decode("utf-8", errors = "replace") + logger.error( + "openai passthrough upstream error: status=%s body=%s", + resp.status_code, + err_text[:500], + ) + upstream_status = resp.status_code try: await resp.aclose() except Exception: @@ -3889,16 +4356,91 @@ async def _openai_passthrough_stream( await client.aclose() except Exception: pass + raise HTTPException( + status_code = upstream_status, + detail = f"llama-server error: {err_text[:500]}", + ) - return StreamingResponse( - _stream(), - media_type = "text/event-stream", - headers = { - "Cache-Control": "no-cache", - "Connection": "keep-alive", - "X-Accel-Buffering": "no", - }, - ) + async def _stream(): + # Same httpx lifecycle pattern as _anthropic_passthrough_stream: + # save resp.aiter_lines() so the finally block can aclose() it + # on our task. See that function for full rationale. + lines_iter = None + # During llama-server prefill, `aiter_lines()` blocks until the + # first SSE chunk arrives. The in-loop `cancel_event` check + # cannot fire until then, which is the exact proxy/Colab + # scenario the cancel POST is meant to recover from. Run a + # tiny watcher that closes `resp` as soon as cancel fires, + # unblocking the iterator with a RemoteProtocolError caught + # in the except clause below. + cancel_watcher = asyncio.create_task( + _await_cancel_then_close(cancel_event, resp) + ) + try: + lines_iter = resp.aiter_lines() + async for raw_line in lines_iter: + if cancel_event.is_set(): + break + if await request.is_disconnected(): + cancel_event.set() + break + if not raw_line: + continue + if not raw_line.startswith("data: "): + continue + # Relay verbatim to preserve llama-server's native id, + # finish_reason, delta.tool_calls, and usage chunks. + yield raw_line + "\n\n" + if raw_line[6:].strip() == "[DONE]": + break + except (httpx.RemoteProtocolError, httpx.ReadError, httpx.CloseError): + # Watcher closed resp on cancel. Emit nothing extra; the + # client either initiated the cancel or already disconnected. + if not cancel_event.is_set(): + raise + except Exception as e: + # 200 headers are already flushed; errors must be in the SSE body. + logger.error("openai passthrough stream error: %s", e) + err = { + "error": { + "message": _friendly_error(e), + "type": "server_error", + }, + } + yield f"data: {json.dumps(err)}\n\n" + finally: + cancel_watcher.cancel() + try: + await cancel_watcher + except (asyncio.CancelledError, Exception): + pass + if lines_iter is not None: + try: + await lines_iter.aclose() + except Exception: + pass + try: + await resp.aclose() + except Exception: + pass + try: + await client.aclose() + except Exception: + pass + _tracker.__exit__(None, None, None) + + return StreamingResponse( + _stream(), + media_type = "text/event-stream", + headers = { + "Cache-Control": "no-cache", + "Connection": "keep-alive", + "X-Accel-Buffering": "no", + }, + ) + except BaseException: + _tracker.__exit__(None, None, None) + raise async def _openai_passthrough_non_streaming( @@ -3914,7 +4456,9 @@ async def _openai_passthrough_non_streaming( token counts. """ target_url = f"{llama_backend.base_url}/v1/chat/completions" - body = _build_openai_passthrough_body(payload) + body = _build_openai_passthrough_body( + payload, backend_ctx = llama_backend.context_length + ) try: async with httpx.AsyncClient() as client: @@ -3935,6 +4479,41 @@ async def _openai_passthrough_non_streaming( detail = f"llama-server error: {resp.text[:500]}", ) + # Guided-decoding fence wrap. llama-server returns raw JSON that matches + # the schema (no surrounding markdown) because the GBNF grammar only + # emits the JSON object itself. data_designer's llm-structured parser + # looks for a ```json ... ``` markdown fence and discards unfenced + # output, which collapses a 100%-valid guided-decoding run to 0/N. + # Wrap each choice's content in the expected fence when the caller + # asked for guided decoding, leaving already-fenced content alone. + if _extract_response_format(payload) is not None: + try: + data = resp.json() + changed = False + for choice in data.get("choices", []): + if not isinstance(choice, dict): + continue + msg = choice.get("message") + if not isinstance(msg, dict): + continue + content = msg.get("content") + if not isinstance(content, str): + continue + stripped = content.strip() + if not stripped or stripped.startswith("```"): + continue + msg["content"] = f"```json\n{stripped}\n```" + changed = True + if changed: + return JSONResponse(content = data) + except Exception as exc: + # Wrap is best-effort; fall through to the verbatim body if + # the response is not JSON-shaped or the structure is unusual. + logger.warning( + "response_format fence wrap skipped: %s", + exc, + ) + # Pass the upstream body through as raw bytes — skips a redundant # parse+re-serialize round-trip and keeps the response truly # verbatim (matches the docstring). Status is guaranteed 200 by diff --git a/studio/backend/routes/models.py b/studio/backend/routes/models.py index db27ce1907..d01e94b0c9 100644 --- a/studio/backend/routes/models.py +++ b/studio/backend/routes/models.py @@ -8,6 +8,7 @@ Model Management API routes import hashlib import json import os +import shutil import sys import uuid from pathlib import Path @@ -623,7 +624,7 @@ def _scan_ollama_dir( gguf_link_path: Optional[str] = None quant = f"-{file_type}" if file_type else "" safe_name = repo_name.replace("/", "-") - for layer in manifest.get("layers", []): + for layer in manifest.get("layers") or []: media = layer.get("mediaType", "") digest = layer.get("digest", "") if not digest: @@ -1684,6 +1685,338 @@ async def scan_loras( ) +def _is_path_under(path: Path, root: Path) -> bool: + try: + path.resolve().relative_to(root.resolve()) + return True + except ValueError: + return False + + +def _is_path_under_lexically(path: Path, root: Path) -> bool: + """Check containment without resolving the final path's symlink target.""" + try: + absolute_path = Path(os.path.abspath(str(path))) + absolute_root = Path(os.path.abspath(str(root))) + absolute_path.relative_to(absolute_root) + return True + except ValueError: + return False + + +def _loaded_model_matches_deleted_path(active_model: str, deleted_path: Path) -> bool: + try: + active = Path(active_model).expanduser().resolve() + target = deleted_path.resolve() + return active == target or (target.is_dir() and active.is_relative_to(target)) + except (OSError, RuntimeError, ValueError) as e: + logger.debug( + "Could not resolve loaded/deleted model paths; falling back to string comparison: %s", + e, + ) + active_lower = active_model.lower() + target_lower = str(deleted_path).lower() + return active_lower == target_lower or active_lower.startswith( + f"{target_lower}{os.sep}" + ) + + +def _loading_model_matches_deleted_path( + loading_model: object, + deleted_path: Path, +) -> bool: + if not loading_model: + return False + return _loaded_model_matches_deleted_path(str(loading_model), deleted_path) + + +def _prune_empty_parents(start: Path, stop_at: Path) -> None: + """Remove empty ancestor directories of ``start`` up to (but not including) ``stop_at``. + + Used after deleting a model checkpoint so the enclosing run directory does + not linger as an empty entry in scan results. + """ + try: + stop_resolved = stop_at.resolve() + except OSError: + return + parent = start.parent + while True: + try: + parent_resolved = parent.resolve() + except OSError: + return + if parent_resolved == stop_resolved: + return + try: + parent_resolved.relative_to(stop_resolved) + except ValueError: + return + try: + parent.rmdir() + except OSError: + return + parent = parent.parent + + +def _delete_gguf_variant_files(root: Path, variant: str) -> tuple[int, int]: + deleted_count = 0 + deleted_bytes = 0 + for path in root.rglob("*"): + if not path.is_file() or not _is_main_gguf_filename(path.name): + continue + if _extract_quant_label(path.name).lower() != variant.lower(): + continue + try: + deleted_bytes += path.stat().st_size + except OSError: + pass + path.unlink() + deleted_count += 1 + return deleted_count, deleted_bytes + + +@router.delete("/delete-finetuned") +async def delete_finetuned_model( + model_path: str = Body(...), + source: str = Body(...), + export_type: Optional[str] = Body(None), + gguf_variant: Optional[str] = Body(None), + current_subject: str = Depends(get_current_subject), +): + """Delete a Studio-trained or exported model from disk. + + Only paths under Studio's outputs/exports roots are accepted. Exported + GGUF entries can delete one quantization variant at a time. + """ + if source not in {"training", "exported"}: + raise HTTPException( + status_code = 400, + detail = "Only trained or exported Studio models can be deleted", + ) + + if not model_path or not model_path.strip(): + raise HTTPException(status_code = 400, detail = "model_path is required") + + if export_type == "gguf" and not gguf_variant: + raise HTTPException( + status_code = 400, + detail = "gguf_variant is required when export_type is 'gguf'", + ) + + raw_path = Path(model_path).expanduser() + if source == "training": + target_path = raw_path + allowed_root = outputs_root() + else: + allowed_root = exports_root() + target_path = ( + raw_path.parent + if export_type == "gguf" and raw_path.suffix.lower() == ".gguf" + else raw_path + ) + + allowed_root = allowed_root.resolve() + delete_path = Path(os.path.abspath(str(target_path))) + delete_path_is_symlink = delete_path.is_symlink() + + if delete_path_is_symlink: + if not _is_path_under_lexically(delete_path, allowed_root): + raise HTTPException( + status_code = 400, + detail = "Model path is outside Studio storage", + ) + if export_type == "gguf" and gguf_variant: + target_path = delete_path.resolve() + if not _is_path_under(target_path, allowed_root): + raise HTTPException( + status_code = 400, + detail = "Model path is outside Studio storage", + ) + else: + target_path = delete_path + else: + target_path = target_path.resolve() + + should_check_resolved_path = not delete_path_is_symlink or ( + export_type == "gguf" and gguf_variant + ) + if should_check_resolved_path and not _is_path_under(target_path, allowed_root): + raise HTTPException( + status_code = 400, + detail = "Model path is outside Studio storage", + ) + if target_path == allowed_root: + raise HTTPException( + status_code = 400, + detail = "Refusing to delete storage root", + ) + if not target_path.exists() and not target_path.is_symlink(): + raise HTTPException(status_code = 404, detail = "Model not found on disk") + + if source == "training": + try: + from core.training import get_training_backend + + training_backend = get_training_backend() + if training_backend.is_training_active(): + raise HTTPException( + status_code = 409, + detail = "Cannot delete trained models while training is running", + ) + except HTTPException: + raise + except Exception as e: + logger.warning("Could not check training status before delete: %s", e) + raise HTTPException( + status_code = 500, + detail = "Could not verify training status before deleting", + ) from e + + try: + from routes.inference import get_llama_cpp_backend + + llama_backend = get_llama_cpp_backend() + if ( + llama_backend.is_active + and not llama_backend.is_loaded + and llama_backend.model_identifier + and _loaded_model_matches_deleted_path( + llama_backend.model_identifier, + target_path, + ) + and ( + not gguf_variant + or not llama_backend.hf_variant + or llama_backend.hf_variant.lower() == gguf_variant.lower() + ) + ): + raise HTTPException( + status_code = 409, + detail = "Cannot delete a model while it is loading", + ) + if ( + llama_backend.is_loaded + and llama_backend.model_identifier + and _loaded_model_matches_deleted_path( + llama_backend.model_identifier, + target_path, + ) + and ( + not gguf_variant + or not llama_backend.hf_variant + or llama_backend.hf_variant.lower() == gguf_variant.lower() + ) + ): + raise HTTPException( + status_code = 400, + detail = "Unload the model before deleting", + ) + except HTTPException: + raise + except Exception as e: + logger.warning("Could not check llama.cpp loaded model before delete: %s", e) + raise HTTPException( + status_code = 503, + detail = "Could not verify model load status before deleting", + ) from e + + try: + inference_backend = get_inference_backend() + loading_models = getattr(inference_backend, "loading_models", set()) + if any( + _loading_model_matches_deleted_path(loading_model, target_path) + for loading_model in loading_models + ): + raise HTTPException( + status_code = 409, + detail = "Cannot delete a model while it is loading", + ) + if inference_backend.active_model_name: + if _loaded_model_matches_deleted_path( + inference_backend.active_model_name, + target_path, + ): + raise HTTPException( + status_code = 400, + detail = "Unload the model before deleting", + ) + except HTTPException: + raise + except Exception as e: + logger.warning( + "Could not check inference backend loaded model before delete: %s", e + ) + raise HTTPException( + status_code = 503, + detail = "Could not verify model load status before deleting", + ) from e + + try: + if export_type == "gguf" and gguf_variant: + if not target_path.is_dir(): + raise HTTPException( + status_code = 400, + detail = "GGUF variant deletion requires an export directory", + ) + deleted_count, deleted_bytes = _delete_gguf_variant_files( + target_path, + gguf_variant, + ) + if deleted_count == 0: + raise HTTPException( + status_code = 404, + detail = f"Variant {gguf_variant} not found on disk", + ) + try: + if not any(target_path.iterdir()): + target_path.rmdir() + _prune_empty_parents(target_path, allowed_root) + except OSError: + pass + logger.info( + "Deleted %s GGUF file(s) for exported model at %s variant %s (%0.1f MB freed)", + deleted_count, + target_path, + gguf_variant, + deleted_bytes / (1024 * 1024), + ) + return { + "status": "deleted", + "path": str(target_path), + "gguf_variant": gguf_variant, + } + + if target_path.is_symlink() or target_path.is_file(): + target_path.unlink() + else: + shutil.rmtree(target_path) + + if target_path.exists() or target_path.is_symlink(): + raise HTTPException( + status_code = 500, + detail = "Deletion incomplete; some files could not be removed", + ) + + _prune_empty_parents(target_path, allowed_root) + + logger.info("Deleted fine-tuned model at %s", target_path) + return {"status": "deleted", "path": str(target_path)} + except HTTPException: + raise + except Exception as e: + logger.error( + "Error deleting fine-tuned model %s: %s", + target_path, + e, + exc_info = True, + ) + raise HTTPException( + status_code = 500, + detail = f"Failed to delete fine-tuned model: {str(e)}", + ) + + @router.get("/loras/{lora_path:path}/base-model", response_model = LoRABaseModelResponse) async def get_lora_base_model( lora_path: str, diff --git a/studio/backend/routes/training.py b/studio/backend/routes/training.py index e625408bad..e5195bb337 100644 --- a/studio/backend/routes/training.py +++ b/studio/backend/routes/training.py @@ -25,6 +25,12 @@ if str(backend_path) not in sys.path: # Import backend functions try: from core.training import get_training_backend + from core.training.resume import ( + can_resume_run, + get_resume_checkpoint_path, + normalize_resume_output_dir, + ) + from storage.studio_db import get_resumable_run_by_output_dir from utils.models.model_config import load_model_defaults from utils.paths import resolve_dataset_path except ImportError: @@ -33,6 +39,12 @@ except ImportError: if str(parent_backend) not in sys.path: sys.path.insert(0, str(parent_backend)) from core.training import get_training_backend + from core.training.resume import ( + can_resume_run, + get_resume_checkpoint_path, + normalize_resume_output_dir, + ) + from storage.studio_db import get_resumable_run_by_output_dir from utils.models.model_config import load_model_defaults from utils.paths import resolve_dataset_path @@ -152,6 +164,28 @@ async def start_training( request.local_eval_datasets = _validate_local_dataset_paths( request.local_eval_datasets, "Local eval dataset" ) + resume_output_dir: Optional[str] = None + if request.resume_from_checkpoint: + try: + resume_output_dir = normalize_resume_output_dir( + request.resume_from_checkpoint + ) + except ValueError as e: + raise HTTPException(status_code = 400, detail = str(e)) + + resume_run = get_resumable_run_by_output_dir(resume_output_dir) + if not resume_run or not can_resume_run(resume_run): + raise HTTPException( + status_code = 400, + detail = "Resume checkpoint must belong to a stopped run with saved trainer state.", + ) + resume_checkpoint = get_resume_checkpoint_path(resume_output_dir) + if not resume_checkpoint: + raise HTTPException( + status_code = 400, + detail = "Resume checkpoint must include saved trainer state.", + ) + request.resume_from_checkpoint = resume_checkpoint # Convert request to kwargs for backend training_kwargs = { @@ -209,6 +243,8 @@ async def start_training( "wandb_project": request.wandb_project or "", "enable_tensorboard": request.enable_tensorboard, "tensorboard_dir": request.tensorboard_dir or "", + "output_dir": resume_output_dir, + "resume_from_checkpoint": request.resume_from_checkpoint, "trust_remote_code": request.trust_remote_code, "gpu_ids": request.gpu_ids, } @@ -437,6 +473,9 @@ async def get_training_status( "loss": getattr(progress, "loss", None), "learning_rate": getattr(progress, "learning_rate", None), } + output_dir = getattr(backend, "_output_dir", None) + if output_dir: + details["output_dir"] = output_dir # Build metric history for chart recovery after SSE reconnection metric_history = None diff --git a/studio/backend/routes/training_history.py b/studio/backend/routes/training_history.py index 597c4424c0..6f34321959 100644 --- a/studio/backend/routes/training_history.py +++ b/studio/backend/routes/training_history.py @@ -11,6 +11,7 @@ from fastapi import APIRouter, Depends, HTTPException, Query from loggers import get_logger from auth.authentication import get_current_subject +from core.training.resume import can_resume_run from models import ( TrainingRunDeleteResponse, TrainingRunDetailResponse, @@ -34,7 +35,10 @@ async def list_training_runs( """List training runs, newest first.""" result = list_runs(limit = limit, offset = offset) return TrainingRunListResponse( - runs = [TrainingRunSummary(**r) for r in result["runs"]], + runs = [ + TrainingRunSummary(**{**r, "can_resume": can_resume_run(r)}) + for r in result["runs"] + ], total = result["total"], ) @@ -58,7 +62,12 @@ async def get_training_run_detail( metrics_data = get_run_metrics(run_id) return TrainingRunDetailResponse( - run = TrainingRunSummary(**{k: v for k, v in run.items() if k != "config_json"}), + run = TrainingRunSummary( + **{ + **{k: v for k, v in run.items() if k != "config_json"}, + "can_resume": can_resume_run(run), + } + ), config = config, metrics = TrainingRunMetrics(**metrics_data), ) diff --git a/studio/backend/run.py b/studio/backend/run.py index 7590ef1067..c5b103ff70 100644 --- a/studio/backend/run.py +++ b/studio/backend/run.py @@ -244,7 +244,7 @@ _shutdown_event = None def run_server( - host: str = "0.0.0.0", + host: str = "127.0.0.1", port: int = 8888, frontend_path: Path = Path(__file__).resolve().parent.parent / "frontend" / "dist", silent: bool = False, @@ -392,7 +392,11 @@ if __name__ == "__main__": pass parser = argparse.ArgumentParser(description = "Run Unsloth UI Backend server") - parser.add_argument("--host", default = "0.0.0.0", help = "Host to bind to") + parser.add_argument( + "--host", + default = "127.0.0.1", + help = "Host to bind to (default: 127.0.0.1; use 0.0.0.0 for network/cloud access)", + ) parser.add_argument("--port", type = int, default = 8888, help = "Port to bind to") parser.add_argument( "--frontend", diff --git a/studio/backend/state/tool_policy.py b/studio/backend/state/tool_policy.py new file mode 100644 index 0000000000..9343a39806 --- /dev/null +++ b/studio/backend/state/tool_policy.py @@ -0,0 +1,33 @@ +# SPDX-License-Identifier: AGPL-3.0-only +# Copyright 2026-present the Unsloth AI Inc. team. All rights reserved. + +"""Process-level server-side tool policy. + +Set by `unsloth run` at startup; consulted by the inference route gates. + + None -> no CLI override (default). Per-request `enable_tools` is honored. + True -> CLI forced tools on for every request. + False -> CLI forced tools off for every request. +""" + +from typing import Optional + +_tool_policy: Optional[bool] = None + + +def get_tool_policy() -> Optional[bool]: + return _tool_policy + + +def set_tool_policy(value: Optional[bool]) -> None: + if value is not None and not isinstance(value, bool): + raise TypeError( + f"tool_policy must be Optional[bool], got {type(value).__name__}" + ) + global _tool_policy + _tool_policy = value + + +def reset_tool_policy() -> None: + global _tool_policy + _tool_policy = None diff --git a/studio/backend/storage/studio_db.py b/studio/backend/storage/studio_db.py index 89f75632ef..29e787c196 100644 --- a/studio/backend/storage/studio_db.py +++ b/studio/backend/storage/studio_db.py @@ -267,10 +267,23 @@ def list_runs(limit: int = 50, offset: int = 0) -> dict: total = conn.execute("SELECT COUNT(*) FROM training_runs").fetchone()[0] rows = conn.execute( """ - SELECT id, status, model_name, dataset_name, started_at, ended_at, - total_steps, final_step, final_loss, output_dir, - duration_seconds, error_message, loss_sparkline - FROM training_runs + SELECT r.id, r.status, r.model_name, r.dataset_name, r.started_at, + r.ended_at, r.total_steps, r.final_step, r.final_loss, + r.output_dir, r.duration_seconds, r.error_message, + r.loss_sparkline, + CASE + WHEN r.status = 'stopped' + AND r.output_dir IS NOT NULL + AND EXISTS ( + SELECT 1 + FROM training_runs newer + WHERE newer.output_dir = r.output_dir + AND newer.status IN ('stopped', 'completed') + AND newer.started_at > r.started_at + ) + THEN 1 ELSE 0 + END AS resumed_later + FROM training_runs r ORDER BY started_at DESC LIMIT ? OFFSET ? """, @@ -297,7 +310,26 @@ def list_runs(limit: int = 50, offset: int = 0) -> dict: def get_run(id: str) -> Optional[dict]: conn = get_connection() try: - row = conn.execute("SELECT * FROM training_runs WHERE id = ?", (id,)).fetchone() + row = conn.execute( + """ + SELECT r.*, + CASE + WHEN r.status = 'stopped' + AND r.output_dir IS NOT NULL + AND EXISTS ( + SELECT 1 + FROM training_runs newer + WHERE newer.output_dir = r.output_dir + AND newer.status IN ('stopped', 'completed') + AND newer.started_at > r.started_at + ) + THEN 1 ELSE 0 + END AS resumed_later + FROM training_runs r + WHERE r.id = ? + """, + (id,), + ).fetchone() if row is None: return None run = dict(row) @@ -313,6 +345,45 @@ def get_run(id: str) -> Optional[dict]: conn.close() +def get_resumable_run_by_output_dir(output_dir: str) -> Optional[dict]: + conn = get_connection() + try: + row = conn.execute( + """ + SELECT r.*, + 0 AS resumed_later + FROM training_runs r + WHERE r.output_dir = ? + AND r.status = 'stopped' + AND NOT EXISTS ( + SELECT 1 + FROM training_runs newer + WHERE newer.output_dir = r.output_dir + AND newer.status IN ('stopped', 'completed') + AND newer.started_at > r.started_at + ) + ORDER BY r.started_at DESC + LIMIT 1 + """, + (output_dir,), + ).fetchone() + if row is None: + return None + run = dict(row) + sparkline = run.get("loss_sparkline") + if sparkline: + try: + run["loss_sparkline"] = json.loads(sparkline) + except (json.JSONDecodeError, TypeError): + logger.debug( + "Failed to parse loss_sparkline for output_dir %s", output_dir + ) + run["loss_sparkline"] = None + return run + finally: + conn.close() + + def get_run_metrics(id: str) -> dict: """Return metric arrays for a run, using paired step arrays per metric.""" conn = get_connection() diff --git a/studio/backend/tests/test_data_recipe_github_progress.py b/studio/backend/tests/test_data_recipe_github_progress.py new file mode 100644 index 0000000000..8e8c3995f4 --- /dev/null +++ b/studio/backend/tests/test_data_recipe_github_progress.py @@ -0,0 +1,91 @@ +# SPDX-License-Identifier: AGPL-3.0-only +# Copyright 2026-present the Unsloth AI Inc. team. All rights reserved. See /studio/LICENSE.AGPL-3.0 + +from core.data_recipe.jobs.parse import apply_update, parse_log_message +from core.data_recipe.jobs.types import Job +from routes.data_recipe.validate import _GITHUB_VALIDATE_NOTE, validate +from models.data_recipe import RecipePayload + + +def test_github_page_log_updates_source_progress_without_cursor(): + job = Job(job_id = "job-1") + job.source_progress_estimated_total = 200 + + update = parse_log_message( + "[unslothai/unsloth] issues page 2 (+15) cursor=abc123 remaining=2960" + ) + + assert update is not None + apply_update(job, update) + + progress = job.source_progress + assert progress is not None + assert progress.source == "github" + assert progress.status == "fetching" + assert progress.repo == "unslothai/unsloth" + assert progress.resource == "issues" + assert progress.page == 2 + assert progress.page_items == 15 + assert progress.fetched_items == 15 + assert progress.estimated_total == 200 + assert progress.rate_remaining == 2960 + assert progress.message is not None + assert "cursor" not in progress.message + assert "abc123" not in progress.message + + +def test_github_rate_limit_log_updates_source_progress(): + job = Job(job_id = "job-1") + + update = parse_log_message("Rate limit hit. Sleeping 123s until reset.") + + assert update is not None + apply_update(job, update) + + progress = job.source_progress + assert progress is not None + assert progress.status == "rate_limited" + assert progress.retry_after_sec == 123 + assert "resume automatically" in (progress.message or "") + + +def test_github_real_sample_prs_and_trial_limit_are_parsed(): + job = Job(job_id = "job-1") + + for message in ( + "[unslothai/unsloth] PRs page 4 (+25) cursor=abc123 remaining=4983", + "Trial limit reached for PRs (100)", + ): + update = parse_log_message(message) + assert update is not None + apply_update(job, update) + + progress = job.source_progress + assert progress is not None + assert progress.repo == "unslothai/unsloth" + assert progress.resource == "pulls" + assert progress.page == 4 + assert progress.fetched_items == 25 + assert progress.rate_remaining == 4983 + assert progress.message == "GitHub pulls trial limit reached (100)." + + +def test_github_validate_skips_live_access_with_honest_note(): + response = validate( + RecipePayload( + recipe = { + "seed_config": { + "source": { + "seed_type": "github_repo", + "repos": ["unslothai/unsloth"], + "item_types": ["issues"], + "limit": 1, + } + }, + "columns": [{"column_type": "expression", "name": "x", "expr": "1"}], + } + ) + ) + + assert response.valid is True + assert response.raw_detail == _GITHUB_VALIDATE_NOTE diff --git a/studio/backend/tests/test_desktop_auth.py b/studio/backend/tests/test_desktop_auth.py index c8cf1c7081..a5508c1c8b 100644 --- a/studio/backend/tests/test_desktop_auth.py +++ b/studio/backend/tests/test_desktop_auth.py @@ -246,7 +246,10 @@ def test_desktop_session_uses_real_admin_identity_for_api_keys(): assert [row["name"] for row in rows] == ["desktop"] -def test_local_recipe_token_preserves_desktop_marker(loaded_local_model): +def test_local_recipe_token_authenticates_as_admin_for_desktop_user(loaded_local_model): + # _inject_local_providers mints an internal sk-unsloth-* API key (not a + # forwarded JWT). The unified API-key path validates as the real admin + # user regardless of whether the incoming session was desktop or web. from auth.authentication import create_access_token, get_current_subject seed_user(must_change_password = True) @@ -260,13 +263,7 @@ def test_local_recipe_token_preserves_desktop_marker(loaded_local_model): jobs_route._inject_local_providers(recipe, local_recipe_request(incoming_token)) local_token = recipe["model_providers"][0]["api_key"] - payload = jwt.decode( - local_token, - storage.get_jwt_secret(storage.DEFAULT_ADMIN_USERNAME), - algorithms = ["HS256"], - ) - assert payload["sub"] == storage.DEFAULT_ADMIN_USERNAME - assert payload["desktop"] is True + assert local_token.startswith(storage.API_KEY_PREFIX) credentials = HTTPAuthorizationCredentials( scheme = "Bearer", credentials = local_token, @@ -276,8 +273,10 @@ def test_local_recipe_token_preserves_desktop_marker(loaded_local_model): ) -def test_local_recipe_token_keeps_web_marker_absent(loaded_local_model): - from auth.authentication import create_access_token +def test_local_recipe_token_authenticates_as_admin_for_web_user(loaded_local_model): + # Mirror of the desktop variant: API-key issuance is identical for web + # and desktop incoming tokens; auth via get_current_subject works the same. + from auth.authentication import create_access_token, get_current_subject seed_user(must_change_password = False) jobs_route = data_recipe_jobs_module() @@ -287,13 +286,14 @@ def test_local_recipe_token_keeps_web_marker_absent(loaded_local_model): jobs_route._inject_local_providers(recipe, local_recipe_request(incoming_token)) local_token = recipe["model_providers"][0]["api_key"] - payload = jwt.decode( - local_token, - storage.get_jwt_secret(storage.DEFAULT_ADMIN_USERNAME), - algorithms = ["HS256"], + assert local_token.startswith(storage.API_KEY_PREFIX) + credentials = HTTPAuthorizationCredentials( + scheme = "Bearer", + credentials = local_token, + ) + assert ( + asyncio.run(get_current_subject(credentials)) == storage.DEFAULT_ADMIN_USERNAME ) - assert payload["sub"] == storage.DEFAULT_ADMIN_USERNAME - assert "desktop" not in payload def test_desktop_login_rejects_invalid_secret(): @@ -381,6 +381,7 @@ def test_health_response_reports_desktop_capability_fields(monkeypatch): datasets_router = APIRouter(), export_router = APIRouter(), inference_router = APIRouter(), + inference_studio_router = APIRouter(), models_router = APIRouter(), training_history_router = APIRouter(), training_router = APIRouter(), diff --git a/studio/backend/tests/test_gpu_selection.py b/studio/backend/tests/test_gpu_selection.py index c6f26037af..a1fe5653ef 100644 --- a/studio/backend/tests/test_gpu_selection.py +++ b/studio/backend/tests/test_gpu_selection.py @@ -746,7 +746,15 @@ class TestRouteErrors(unittest.TestCase): ): with self.assertRaises(HTTPException) as exc_info: asyncio.run( - inference_route.load_model(request, current_subject = "test-user") + inference_route.load_model( + request, + SimpleNamespace( + app = SimpleNamespace( + state = SimpleNamespace(llama_parallel_slots = 1), + ), + ), + current_subject = "test-user", + ) ) self.assertEqual(exc_info.exception.status_code, 400) @@ -886,7 +894,15 @@ class TestRouteErrors(unittest.TestCase): ): with self.assertRaises(HTTPException) as exc_info: asyncio.run( - inference_route.load_model(request, current_subject = "test-user") + inference_route.load_model( + request, + SimpleNamespace( + app = SimpleNamespace( + state = SimpleNamespace(llama_parallel_slots = 1), + ), + ), + current_subject = "test-user", + ) ) self.assertEqual(exc_info.exception.status_code, 400) @@ -942,7 +958,15 @@ class TestRouteErrors(unittest.TestCase): ): with self.assertRaises(HTTPException) as exc_info: asyncio.run( - inference_route.load_model(request, current_subject = "test-user") + inference_route.load_model( + request, + SimpleNamespace( + app = SimpleNamespace( + state = SimpleNamespace(llama_parallel_slots = 1), + ), + ), + current_subject = "test-user", + ) ) self.assertEqual(exc_info.exception.status_code, 400) @@ -1025,6 +1049,182 @@ class TestMinGpuVram(unittest.TestCase): class TestPerGpuFitGuardAllCounts(unittest.TestCase): + def test_training_estimate_resolves_attention_without_raising(self): + with ( + patch("utils.hardware.hardware.get_device", return_value = DeviceType.CUDA), + patch( + "utils.hardware.hardware.estimate_fp16_model_size_bytes", + return_value = (8 * (1024**3), "config"), + ), + patch( + "utils.hardware.hardware._resolve_model_identifier_for_gpu_estimate", + return_value = "unsloth/test", + ), + patch( + "utils.hardware.hardware._load_config_for_gpu_estimate", + return_value = SimpleNamespace( + hidden_size = 4096, + num_hidden_layers = 32, + num_attention_heads = 32, + num_key_value_heads = 8, + intermediate_size = 14336, + vocab_size = 128256, + tie_word_embeddings = False, + ), + ), + patch( + "utils.hardware.hardware._determine_attention_impl_for_gpu_estimate", + return_value = "eager", + ), + patch("utils.hardware.hardware.get_visible_gpu_count", return_value = 1), + ): + _, metadata = estimate_required_model_memory_gb( + "unsloth/test", + training_type = "LoRA/QLoRA", + load_in_4bit = True, + ) + + self.assertEqual(metadata.get("estimation_mode"), "detailed") + self.assertEqual(metadata.get("attention_implementation"), "eager") + + def test_training_estimate_falls_back_when_attention_resolution_fails(self): + with ( + patch("utils.hardware.hardware.get_device", return_value = DeviceType.CUDA), + patch( + "utils.hardware.hardware.estimate_fp16_model_size_bytes", + return_value = (8 * (1024**3), "config"), + ), + patch( + "utils.hardware.hardware._resolve_model_identifier_for_gpu_estimate", + return_value = "unsloth/test", + ), + patch( + "utils.hardware.hardware._load_config_for_gpu_estimate", + return_value = SimpleNamespace( + hidden_size = 4096, + num_hidden_layers = 32, + num_attention_heads = 32, + num_key_value_heads = 8, + intermediate_size = 14336, + vocab_size = 128256, + tie_word_embeddings = False, + ), + ), + patch( + "utils.hardware.hardware._determine_attention_impl_for_gpu_estimate", + side_effect = RuntimeError("attention unavailable"), + ), + patch("utils.hardware.hardware.get_visible_gpu_count", return_value = 1), + ): + _, metadata = estimate_required_model_memory_gb( + "unsloth/test", + training_type = "LoRA/QLoRA", + load_in_4bit = True, + ) + + self.assertEqual(metadata.get("estimation_mode"), "detailed") + self.assertEqual( + metadata.get("attention_implementation"), + "eager", + ) + + def test_attention_resolver_does_not_mutate_loaded_config(self): + from utils.hardware import hardware as hardware_module + + config = SimpleNamespace( + hidden_size = 1024, + num_hidden_layers = 2, + num_attention_heads = 8, + num_key_value_heads = 8, + intermediate_size = 2048, + vocab_size = 1024, + tie_word_embeddings = True, + ) + + def _stub_resolver(model_class, cfg): + cfg._attn_implementation = "eager" + return "eager" + + with patch( + "unsloth.models._utils.resolve_attention_implementation", + side_effect = _stub_resolver, + ): + hardware_module._determine_attention_impl_for_gpu_estimate(config) + + self.assertFalse(hasattr(config, "_attn_implementation")) + + def test_attention_resolver_handles_missing_model_mapping(self): + from utils.hardware import hardware as hardware_module + + config = SimpleNamespace( + hidden_size = 1024, + num_hidden_layers = 2, + num_attention_heads = 8, + num_key_value_heads = 8, + intermediate_size = 2048, + vocab_size = 1024, + tie_word_embeddings = True, + ) + captured = {} + + def _stub_resolver(model_class, cfg): + captured["model_class"] = model_class + return "eager" + + from transformers import AutoModel, AutoModelForCausalLM + + with ( + patch.object(AutoModelForCausalLM, "_model_mapping", new = None), + patch.object(AutoModel, "_model_mapping", new = None), + patch( + "unsloth.models._utils.resolve_attention_implementation", + side_effect = _stub_resolver, + ), + ): + result = hardware_module._determine_attention_impl_for_gpu_estimate(config) + + self.assertEqual(result, "eager") + self.assertIsNone(captured["model_class"]) + + def test_attention_resolver_does_not_mutate_nested_text_config(self): + from utils.hardware import hardware as hardware_module + + text_config = SimpleNamespace( + hidden_size = 1024, + num_hidden_layers = 2, + num_attention_heads = 8, + num_key_value_heads = 8, + intermediate_size = 2048, + vocab_size = 1024, + tie_word_embeddings = True, + ) + config = SimpleNamespace( + hidden_size = 1024, + num_hidden_layers = 2, + num_attention_heads = 8, + num_key_value_heads = 8, + intermediate_size = 2048, + vocab_size = 1024, + tie_word_embeddings = True, + text_config = text_config, + ) + + def _stub_resolver(model_class, cfg): + cfg._attn_implementation = "eager" + inner = getattr(cfg, "text_config", None) + if inner is not None: + inner._attn_implementation = "eager" + return "eager" + + with patch( + "unsloth.models._utils.resolve_attention_implementation", + side_effect = _stub_resolver, + ): + hardware_module._determine_attention_impl_for_gpu_estimate(config) + + self.assertFalse(hasattr(config, "_attn_implementation")) + self.assertFalse(hasattr(text_config, "_attn_implementation")) + def test_min_per_gpu_generated_for_all_visible_counts(self): with ( patch("utils.hardware.hardware.get_device", return_value = DeviceType.CUDA), @@ -1101,3 +1301,123 @@ class TestXpuRejection(_GpuCacheResetMixin, unittest.TestCase): with patch("utils.hardware.hardware.get_device", return_value = DeviceType.XPU): with self.assertRaisesRegex(ValueError, "only supported on CUDA"): prepare_gpu_selection([0], model_name = "unsloth/test") + + +class TestEstimateFp16ModelSizeBytesPrefersLocalWeights(unittest.TestCase): + def _run( + self, + model_path, + *, + config_bytes, + local_bytes, + safetensors_params = None, + config = object(), + ): + from utils.hardware import hardware as hardware_module + + with ( + patch.object( + hardware_module, + "_resolve_model_identifier_for_gpu_estimate", + return_value = model_path, + ), + patch.object( + hardware_module, + "_get_hf_safetensors_total_params", + return_value = safetensors_params, + ), + patch.object( + hardware_module, + "_load_config_for_gpu_estimate", + return_value = config, + ), + patch.object( + hardware_module, + "_estimate_fp16_model_size_bytes_from_config", + return_value = config_bytes, + ), + patch.object( + hardware_module, + "_get_local_weight_size_bytes", + return_value = local_bytes, + ), + ): + return hardware_module.estimate_fp16_model_size_bytes(model_path) + + def test_local_weight_bytes_preferred_when_larger_than_config(self): + bytes_, src = self._run( + "/local/vlm", + config_bytes = 2 * (1 << 30), + local_bytes = 20 * (1 << 30), + ) + self.assertEqual(bytes_, 20 * (1 << 30)) + self.assertEqual(src, "weight_bytes") + + def test_config_bytes_preferred_when_larger_than_local(self): + bytes_, src = self._run( + "/local/text-only", + config_bytes = 20 * (1 << 30), + local_bytes = 2 * (1 << 30), + ) + self.assertEqual(bytes_, 20 * (1 << 30)) + self.assertEqual(src, "config") + + def test_config_bytes_returned_when_no_local_weights(self): + bytes_, src = self._run( + "/local/no-weights", + config_bytes = 5 * (1 << 30), + local_bytes = None, + ) + self.assertEqual(bytes_, 5 * (1 << 30)) + self.assertEqual(src, "config") + + def test_local_bytes_returned_when_config_resolution_fails(self): + bytes_, src = self._run( + "/local/no-config", + config_bytes = None, + local_bytes = 7 * (1 << 30), + config = None, + ) + self.assertEqual(bytes_, 7 * (1 << 30)) + self.assertEqual(src, "weight_bytes") + + def test_equal_local_and_config_keeps_config_label(self): + # why: tie-breaker is "local must be strictly larger" so an exact + # match keeps the config-derived path. + same = 8 * (1 << 30) + bytes_, src = self._run( + "/local/equal", + config_bytes = same, + local_bytes = same, + ) + self.assertEqual(bytes_, same) + self.assertEqual(src, "config") + + def test_remote_safetensors_path_unaffected_by_local_weights(self): + from utils.hardware import hardware as hardware_module + + with ( + patch.object( + hardware_module, + "_resolve_model_identifier_for_gpu_estimate", + return_value = "owner/repo", + ), + patch.object( + hardware_module, + "_get_hf_safetensors_total_params", + return_value = 1_000_000_000, + ), + patch.object( + hardware_module, + "_load_config_for_gpu_estimate", + ) as mock_load, + patch.object( + hardware_module, + "_get_local_weight_size_bytes", + ) as mock_local, + ): + bytes_, src = hardware_module.estimate_fp16_model_size_bytes("owner/repo") + self.assertEqual(bytes_, 2 * 1_000_000_000) + self.assertEqual(src, "safetensors") + mock_load.assert_not_called() + mock_local.assert_not_called() diff --git a/studio/backend/tests/test_host_defaults.py b/studio/backend/tests/test_host_defaults.py new file mode 100644 index 0000000000..8b81474e92 --- /dev/null +++ b/studio/backend/tests/test_host_defaults.py @@ -0,0 +1,98 @@ +# SPDX-License-Identifier: AGPL-3.0-only +# Copyright 2026-present the Unsloth AI Inc. team. All rights reserved. See /studio/LICENSE.AGPL-3.0 + +"""Tests that Unsloth Studio defaults to 127.0.0.1 (loopback) not 0.0.0.0. + +Uses AST parsing to inspect source-level defaults without requiring the +full studio venv (run.py has heavy dependencies like structlog/uvicorn). +""" + +import ast +from pathlib import Path + +_RUN_PY = Path(__file__).resolve().parent.parent / "run.py" + + +def _parse_function_param_defaults(source: str, func_name: str) -> dict: + """Return {param_name: default_value} for a named function in *source*. + + Only handles ast.Constant defaults (strings, ints, bools). + """ + tree = ast.parse(source) + for node in ast.walk(tree): + if ( + isinstance(node, (ast.FunctionDef, ast.AsyncFunctionDef)) + and node.name == func_name + ): + result = {} + all_args = node.args.args + defaults = node.args.defaults + # Defaults are right-aligned against the args list + offset = len(all_args) - len(defaults) + for i, default in enumerate(defaults): + arg_name = all_args[offset + i].arg + if isinstance(default, ast.Constant): + result[arg_name] = default.value + return result + return {} + + +def _parse_argparse_add_argument_default(source: str, option_name: str): + """Return the 'default' kwarg value for add_argument(option_name, ...) in *source*. + + Walks the entire module so the call can live in __main__ or in a helper + function — only handles ast.Constant defaults. + """ + tree = ast.parse(source) + for node in ast.walk(tree): + if not isinstance(node, ast.Call): + continue + func = node.func + if not (isinstance(func, ast.Attribute) and func.attr == "add_argument"): + continue + if not node.args: + continue + first_arg = node.args[0] + if not (isinstance(first_arg, ast.Constant) and first_arg.value == option_name): + continue + for kw in node.keywords: + if kw.arg == "default" and isinstance(kw.value, ast.Constant): + return kw.value.value + return None + + +def test_run_server_default_host_is_loopback(): + """run_server() parameter default for 'host' must be 127.0.0.1, not 0.0.0.0. + + Binding to 0.0.0.0 by default exposes the service on all network + interfaces, contradicting the documented "privacy first / 100% local" + guarantee. Loopback (127.0.0.1) is the least-permissive default; + users who need network access can pass -H 0.0.0.0 explicitly. + """ + source = _RUN_PY.read_text() + defaults = _parse_function_param_defaults(source, "run_server") + assert ( + "host" in defaults + ), "run_server() must have a 'host' parameter with a default" + host_default = defaults["host"] + assert host_default == "127.0.0.1", ( + f"run_server() host default must be '127.0.0.1' (loopback) " + f"but got '{host_default}'. Binding to '{host_default}' by default " + f"exposes the service beyond localhost." + ) + + +def test_argparse_default_host_is_loopback(): + """argparse --host add_argument default must be 127.0.0.1. + + When run.py is invoked directly (python run.py), the argparse default + should match the function default so direct execution is equally safe. + """ + source = _RUN_PY.read_text() + host_default = _parse_argparse_add_argument_default(source, "--host") + assert ( + host_default is not None + ), "Could not find add_argument('--host', ...) in run.py" + assert ( + host_default == "127.0.0.1" + ), f"run.py argparse --host default must be '127.0.0.1', got '{host_default}'" diff --git a/studio/backend/tests/test_kv_cache_estimation.py b/studio/backend/tests/test_kv_cache_estimation.py index 2640ded90d..29d87804ff 100644 --- a/studio/backend/tests/test_kv_cache_estimation.py +++ b/studio/backend/tests/test_kv_cache_estimation.py @@ -12,6 +12,7 @@ Cross-platform: Linux, macOS, Windows, WSL. """ import io +import json import struct import sys import types as _types @@ -37,35 +38,43 @@ sys.modules.setdefault("loggers", _loggers_stub) _structlog_stub = _types.ModuleType("structlog") sys.modules.setdefault("structlog", _structlog_stub) -# httpx -_httpx_stub = _types.ModuleType("httpx") -for _exc_name in ( - "ConnectError", - "TimeoutException", - "ReadTimeout", - "ReadError", - "RemoteProtocolError", - "CloseError", -): - setattr(_httpx_stub, _exc_name, type(_exc_name, (Exception,), {})) +# httpx -- only stub when the real library isn't installed. Stubbing +# unconditionally would shadow ``HTTPError`` / ``Response`` etc. that +# ``huggingface_hub.errors`` imports at module load time, which causes +# the transformers introspection tier to silently return None inside +# the test process. +try: + import httpx as _httpx_real # noqa: F401 +except ImportError: + _httpx_stub = _types.ModuleType("httpx") + for _exc_name in ( + "ConnectError", + "TimeoutException", + "ReadTimeout", + "ReadError", + "RemoteProtocolError", + "CloseError", + "HTTPError", + "RequestError", + ): + setattr(_httpx_stub, _exc_name, type(_exc_name, (Exception,), {})) + class _FakeTimeout: + def __init__(self, *a, **kw): + pass -class _FakeTimeout: - def __init__(self, *a, **kw): - pass - - -_httpx_stub.Timeout = _FakeTimeout -_httpx_stub.Client = type( - "Client", - (), - { - "__init__": lambda self, **kw: None, - "__enter__": lambda self: self, - "__exit__": lambda self, *a: None, - }, -) -sys.modules.setdefault("httpx", _httpx_stub) + _httpx_stub.Timeout = _FakeTimeout + _httpx_stub.Response = type("Response", (), {}) + _httpx_stub.Client = type( + "Client", + (), + { + "__init__": lambda self, **kw: None, + "__enter__": lambda self: self, + "__exit__": lambda self, *a: None, + }, + ) + sys.modules["httpx"] = _httpx_stub from core.inference.llama_cpp import LlamaCppBackend @@ -77,8 +86,7 @@ from core.inference.llama_cpp import LlamaCppBackend def _make_gguf_bytes(arch: str, kv_pairs: dict) -> bytes: """Build a minimal GGUF v3 binary blob with the given KV metadata. - Only supports UINT32 (type 4), UINT64 (type 10), and STRING (type 8) - values, which is all the metadata parser reads. + Supports the scalar and simple array metadata used by the parser. """ buf = io.BytesIO() # Header: magic, version, tensor_count, kv_count @@ -96,6 +104,17 @@ def _make_gguf_bytes(arch: str, kv_pairs: dict) -> bytes: val_bytes = val.encode("utf-8") buf.write(struct.pack(" bytes: return buf.getvalue() -def _backend_from_gguf(arch: str, fields: dict) -> LlamaCppBackend: - """Create a LlamaCppBackend with parsed GGUF metadata from given fields.""" +def _backend_from_gguf( + arch: str, fields: dict, general: dict | None = None +) -> LlamaCppBackend: + """Create a LlamaCppBackend with parsed GGUF metadata from given fields. + + `general` lets a test inject extra `general.*` metadata (used to + verify the dynamic SWA resolver picks up source-repo hints from + GGUFs that ship them). + """ kv = {"general.architecture": arch} + for k, v in (general or {}).items(): + kv[k] = v for k, v in fields.items(): kv[f"{arch}.{k}"] = v import tempfile, os @@ -133,7 +161,7 @@ def _backend_from_gguf(arch: str, fields: dict) -> LlamaCppBackend: class TestGGUFParserNewFields: - """Verify that the 8 new architecture-aware fields are correctly parsed.""" + """Verify that architecture-aware fields are correctly parsed.""" @pytest.mark.parametrize( "field,gguf_key,value", @@ -158,15 +186,189 @@ class TestGGUFParserNewFields: "_kv_key_length", "_kv_value_length", "_sliding_window", + "_sliding_window_pattern", "_full_attention_interval", "_kv_lora_rank", "_key_length_mla", + "_kv_key_length_swa", + "_kv_value_length_swa", "_ssm_inner_size", "_ssm_state_size", ]: assert getattr(b, attr) is None - def test_all_13_fields_parsed_together(self): + def test_array_fields_parsed(self): + b = _backend_from_gguf( + "gemma4", + { + "block_count": 6, + "attention.head_count_kv": [8, 8, 8, 8, 8, 2], + "attention.sliding_window_pattern": [ + True, + True, + True, + True, + True, + False, + ], + }, + ) + # Per-layer KV head count is preserved exactly... + assert b._n_kv_heads_by_layer == [8, 8, 8, 8, 8, 2] + # ...and mirrored into the scalar field as a conservative max so + # non-SWA estimator paths and any caller using + # `n_kv = self._n_kv_heads or ...` get a safe upper bound. + assert b._n_kv_heads == 8 + assert b._sliding_window_pattern == [True, True, True, True, True, False] + + +class TestArchSwaPatternDefaults: + """Bootstrap arch table fires when GGUF reports `sliding_window` but + no per-layer pattern (true for every Gemma 2/3/3n/gpt-oss GGUF today).""" + + @pytest.mark.parametrize( + "arch,n_layers,expected_period", + [ + ("gemma2", 26, 2), + ("gemma3", 18, 6), + ("gemma3n", 35, 5), + ("gpt_oss", 24, 2), + ("cohere2", 32, 4), + ], + ) + def test_arch_default_pattern_applied(self, arch, n_layers, expected_period): + b = _backend_from_gguf( + arch, + { + "block_count": n_layers, + "attention.head_count": 4, + "attention.head_count_kv": 1, + "attention.key_length": 256, + "attention.value_length": 256, + "attention.sliding_window": 512, + }, + ) + expected_pattern = [(i + 1) % expected_period != 0 for i in range(n_layers)] + assert ( + b._sliding_window_pattern == expected_pattern + ), f"{arch} should expand to period={expected_period}" + + def test_unknown_arch_no_default(self): + b = _backend_from_gguf( + "totallymadeupv7", + { + "block_count": 24, + "attention.head_count": 4, + "attention.head_count_kv": 1, + "attention.key_length": 128, + "attention.value_length": 128, + "attention.sliding_window": 1024, + }, + ) + assert b._sliding_window_pattern is None + + def test_explicit_pattern_overrides_arch_default(self): + # Period=6 is the gemma3 default; the explicit array must win. + b = _backend_from_gguf( + "gemma3", + { + "block_count": 6, + "attention.head_count": 4, + "attention.head_count_kv": 1, + "attention.key_length": 256, + "attention.value_length": 256, + "attention.sliding_window": 512, + "attention.sliding_window_pattern": [ + True, + False, + True, + False, + True, + False, + ], + }, + ) + assert b._sliding_window_pattern == [True, False, True, False, True, False] + + def test_no_sliding_window_no_pattern(self): + b = _backend_from_gguf( + "gemma3", + { + "block_count": 18, + "attention.head_count": 4, + "attention.head_count_kv": 1, + "attention.key_length": 256, + "attention.value_length": 256, + # no sliding_window key + }, + ) + assert b._sliding_window_pattern is None + + @pytest.mark.parametrize( + "arch", ["llama", "qwen2", "qwen3", "mistral", "mistral3", "glm4", "llama4"] + ) + def test_non_swa_arch_uses_full_attention_path(self, arch): + # Pure-GQA arches: GGUF has no sliding_window, no synthetic + # pattern, estimator hits Path 4. + b = _backend_from_gguf( + arch, + { + "block_count": 32, + "attention.head_count": 32, + "attention.head_count_kv": 8, + "attention.key_length": 128, + "attention.value_length": 128, + "embedding_length": 4096, + }, + ) + assert b._sliding_window_pattern is None + assert b._sliding_window is None + kv = b._estimate_kv_cache_bytes(8192, "f16") + gqa_expected = 32 * 8192 * 8 * (128 + 128) * 2 + assert kv == gqa_expected + + def test_arch_default_reduces_kv_estimate_vs_legacy(self): + common = { + "block_count": 62, + "attention.head_count": 32, + "attention.head_count_kv": 16, + "attention.key_length": 128, + "attention.value_length": 128, + "attention.sliding_window": 1024, + "embedding_length": 5376, + } + with_default = _backend_from_gguf("gemma3", common) + # Arch not in the table -> legacy 1/4 path. + without_default = _backend_from_gguf("totallymadeupv7", common) + + kv_default = with_default._estimate_kv_cache_bytes(131072, "f16") + kv_legacy = without_default._estimate_kv_cache_bytes(131072, "f16") + assert kv_default > 0 + assert kv_legacy > 0 + assert kv_default < kv_legacy, ( + f"arch fallback should under-shoot legacy estimate: " + f"{kv_default} >= {kv_legacy}" + ) + + def test_scalar_sliding_window_pattern_expanded(self): + block_count = 8 + b = _backend_from_gguf( + "gemma3", + { + "attention.sliding_window_pattern": 4, + "block_count": block_count, + "attention.head_count_kv": 4, + "attention.key_length": 256, + "attention.value_length": 256, + "attention.sliding_window": 1024, + }, + ) + expected = [(i + 1) % 4 != 0 for i in range(block_count)] + assert isinstance(b._sliding_window_pattern, list) + assert b._sliding_window_pattern == expected + assert b._estimate_kv_cache_bytes(4096, "f16") > 0 + + def test_all_fields_parsed_together(self): fields = { "context_length": 131072, "block_count": 62, @@ -176,9 +378,12 @@ class TestGGUFParserNewFields: "attention.key_length": 128, "attention.value_length": 128, "attention.sliding_window": 1024, + "attention.sliding_window_pattern": [True, False], "full_attention_interval": 6, "attention.kv_lora_rank": 512, "attention.key_length_mla": 256, + "attention.key_length_swa": 64, + "attention.value_length_swa": 64, "ssm.inner_size": 4096, "ssm.state_size": 128, } @@ -191,13 +396,294 @@ class TestGGUFParserNewFields: assert b._kv_key_length == 128 assert b._kv_value_length == 128 assert b._sliding_window == 1024 + assert b._sliding_window_pattern == [True, False] assert b._full_attention_interval == 6 assert b._kv_lora_rank == 512 assert b._key_length_mla == 256 + assert b._kv_key_length_swa == 64 + assert b._kv_value_length_swa == 64 assert b._ssm_inner_size == 4096 assert b._ssm_state_size == 128 +_SWA_FIELDS = { + "block_count": 12, + "attention.head_count": 4, + "attention.head_count_kv": 1, + "attention.key_length": 256, + "attention.value_length": 256, + "attention.sliding_window": 512, +} + + +class TestDynamicSwaResolver: + """4-tier resolver: GGUF metadata, on-disk cache, bootstrap, HF fetch.""" + + def _isolate_cache(self, monkeypatch, tmp_path): + from core.inference import llama_cpp as lc + + monkeypatch.setenv("UNSLOTH_STUDIO_HOME", str(tmp_path)) + monkeypatch.setattr(lc, "_SWA_CACHE", None) + return tmp_path + + def test_period_from_layer_types_finds_smallest_period(self): + from core.inference.llama_cpp import _period_from_layer_types + + # gemma3 (1 global per 6), gpt-oss (alternating), gemma3n (1 per 5). + assert ( + _period_from_layer_types( + (["sliding_attention"] * 5 + ["full_attention"]) * 4 + ) + == 6 + ) + assert ( + _period_from_layer_types(["sliding_attention", "full_attention"] * 12) == 2 + ) + assert ( + _period_from_layer_types( + (["sliding_attention"] * 4 + ["full_attention"]) * 7 + ) + == 5 + ) + + def test_period_from_layer_types_returns_none_for_aperiodic(self): + from core.inference.llama_cpp import _period_from_layer_types + + lt = [ + "sliding_attention", + "full_attention", + "sliding_attention", + "sliding_attention", + "full_attention", + "sliding_attention", + "sliding_attention", + "sliding_attention", + ] + assert _period_from_layer_types(lt) is None + + def test_hf_repo_from_url(self): + from core.inference.llama_cpp import _hf_repo_from_url + + assert ( + _hf_repo_from_url("https://huggingface.co/google/gemma-3-1b-it") + == "google/gemma-3-1b-it" + ) + assert ( + _hf_repo_from_url( + "https://huggingface.co/google/gemma-3-1b-it/blob/main/config.json" + ) + == "google/gemma-3-1b-it" + ) + for bad in [ + "https://huggingface.co/google", + "https://example.com/foo/bar", + None, + "", + ]: + assert _hf_repo_from_url(bad) is None + + def test_bootstrap_tier_used_when_no_cache(self, monkeypatch, tmp_path): + self._isolate_cache(monkeypatch, tmp_path) + from core.inference import llama_cpp as lc + + def boom(*a, **kw): + raise AssertionError("HF fetch must not run when bootstrap covers the arch") + + monkeypatch.setattr(lc, "_fetch_swa_entry_from_hf", boom) + b = _backend_from_gguf("gemma3", dict(_SWA_FIELDS, block_count = 18)) + assert b._sliding_window_pattern == [(i + 1) % 6 != 0 for i in range(18)] + + def test_disk_cache_takes_precedence_over_bootstrap(self, monkeypatch, tmp_path): + self._isolate_cache(monkeypatch, tmp_path) + # Override bootstrap=6 with a cached period=3. + with open(tmp_path / "swa_cache.json", "w") as f: + json.dump({"gemma3": 3}, f) + b = _backend_from_gguf("gemma3", dict(_SWA_FIELDS, block_count = 18)) + assert b._sliding_window_pattern == [(i + 1) % 3 != 0 for i in range(18)] + + def test_disk_cache_supports_array_entries(self, monkeypatch, tmp_path): + # Aperiodic mask gets tiled across n_layers. + self._isolate_cache(monkeypatch, tmp_path) + mask = [True, False, True, True, False, True, False, False] + with open(tmp_path / "swa_cache.json", "w") as f: + json.dump({"customarch": mask}, f) + b = _backend_from_gguf("customarch", dict(_SWA_FIELDS, block_count = 16)) + assert b._sliding_window_pattern == [bool(mask[i % 8]) for i in range(16)] + + def test_hf_fetch_populates_cache(self, monkeypatch, tmp_path): + self._isolate_cache(monkeypatch, tmp_path) + from core.inference import llama_cpp as lc + + calls = [] + + def fake_fetch(repo_id): + calls.append(repo_id) + return 4 if repo_id == "vendor/newmodel-1b-instruct" else None + + monkeypatch.setattr(lc, "_fetch_swa_entry_from_hf", fake_fetch) + b = _backend_from_gguf( + "newmodel", + _SWA_FIELDS, + general = { + "general.source.huggingface.repository": "vendor/newmodel-1b-instruct" + }, + ) + assert b._sliding_window_pattern == [(i + 1) % 4 != 0 for i in range(12)] + assert calls == ["vendor/newmodel-1b-instruct"] + with open(tmp_path / "swa_cache.json") as f: + assert json.load(f) == {"newmodel": 4} + + def test_hf_fetch_falls_back_to_other_candidates(self, monkeypatch, tmp_path): + self._isolate_cache(monkeypatch, tmp_path) + from core.inference import llama_cpp as lc + + monkeypatch.setattr( + lc, + "_fetch_swa_entry_from_hf", + lambda r: 6 if r == "vendor/newmodel-base" else None, + ) + b = _backend_from_gguf( + "newmodel", + _SWA_FIELDS, + general = { + "general.base_model.0.repo_url": "https://huggingface.co/vendor/newmodel-base" + }, + ) + assert b._sliding_window_pattern == [(i + 1) % 6 != 0 for i in range(12)] + + def test_offline_env_skips_network(self, monkeypatch, tmp_path): + self._isolate_cache(monkeypatch, tmp_path) + monkeypatch.setenv("UNSLOTH_STUDIO_OFFLINE", "1") + from core.inference import llama_cpp as lc + + def boom(*a, **kw): + raise AssertionError("HF fetch must not run when offline=1") + + monkeypatch.setattr(lc, "_fetch_swa_entry_from_hf", boom) + b = _backend_from_gguf( + "newmodel", + _SWA_FIELDS, + general = {"general.source.huggingface.repository": "vendor/newmodel"}, + ) + assert b._sliding_window_pattern is None + + def test_hf_fetch_failure_falls_through_silently(self, monkeypatch, tmp_path): + self._isolate_cache(monkeypatch, tmp_path) + from core.inference import llama_cpp as lc + + monkeypatch.setattr(lc, "_fetch_swa_entry_from_hf", lambda repo_id: None) + # Force the failure into the Tier 3 path; bypass Tier 2.5. + monkeypatch.setattr( + lc, "_resolve_swa_entry_from_transformers", lambda arch: None + ) + b = _backend_from_gguf( + "newmodel", + _SWA_FIELDS, + general = {"general.source.huggingface.repository": "vendor/does-not-exist"}, + ) + assert b._sliding_window_pattern is None + assert not (tmp_path / "swa_cache.json").exists() + + +class TestTransformersIntrospection: + """Tier 2.5: default-init the matching Config; on failure, parse via inspect.""" + + def _isolate_cache(self, monkeypatch, tmp_path): + from core.inference import llama_cpp as lc + + monkeypatch.setenv("UNSLOTH_STUDIO_HOME", str(tmp_path)) + monkeypatch.setattr(lc, "_SWA_CACHE", None) + return tmp_path + + def test_arch_aliases_normalises_hyphen_underscore(self): + from core.inference.llama_cpp import _arch_aliases + + aliases = _arch_aliases("falcon-h1") + assert aliases[0] == "falcon-h1" and "falcon_h1" in aliases + assert _arch_aliases("gemma3") == ("gemma3",) + assert _arch_aliases("") == () + + def test_resolves_real_transformers_arches(self): + from core.inference.llama_cpp import _resolve_swa_entry_from_transformers + + assert _resolve_swa_entry_from_transformers("gemma3") == 6 + assert _resolve_swa_entry_from_transformers("gemma2") == 2 + assert _resolve_swa_entry_from_transformers("cohere2") == 4 + + def test_falls_back_to_inspect_when_default_init_raises(self, monkeypatch): + from core.inference import llama_cpp as lc + + class _FakeBrokenConfig: + """Class with sliding_window_pattern: int = 7 in its docstring.""" + + def __init__(self, required_arg): + raise TypeError("requires an argument") + + class _FakeLazyMapping(dict): + def __getitem__(self, k): + return ( + _FakeBrokenConfig if k == "brokenarch" else super().__getitem__(k) + ) + + import sys, types as _types + + fake_auto = _types.ModuleType("transformers.models.auto.configuration_auto") + fake_auto.CONFIG_MAPPING_NAMES = {"brokenarch": "FakeBroken"} + fake_auto.CONFIG_MAPPING = _FakeLazyMapping({"brokenarch": "FakeBroken"}) + monkeypatch.setitem( + sys.modules, "transformers.models.auto.configuration_auto", fake_auto + ) + assert lc._resolve_swa_entry_from_transformers("brokenarch") == 7 + + def test_returns_none_when_transformers_unavailable(self, monkeypatch): + from core.inference import llama_cpp as lc + import sys + + orig_import = ( + __builtins__["__import__"] + if isinstance(__builtins__, dict) + else __builtins__.__import__ + ) + + def fake_import(name, *a, **kw): + if name.startswith("transformers"): + raise ImportError("transformers not installed") + return orig_import(name, *a, **kw) + + monkeypatch.setattr("builtins.__import__", fake_import) + for k in list(sys.modules): + if k.startswith("transformers"): + monkeypatch.delitem(sys.modules, k, raising = False) + assert lc._resolve_swa_entry_from_transformers("gemma3") is None + + def test_returns_none_for_arch_unknown_to_transformers(self): + from core.inference.llama_cpp import _resolve_swa_entry_from_transformers + + assert _resolve_swa_entry_from_transformers("totally-fake-arch-xyz") is None + + def test_full_resolver_uses_transformers_before_hf_fetch( + self, monkeypatch, tmp_path + ): + # With bootstrap empty, Tier 2.5 must answer before Tier 3 fires. + self._isolate_cache(monkeypatch, tmp_path) + from core.inference import llama_cpp as lc + + monkeypatch.setattr(lc, "_BOOTSTRAP_SWA_DEFAULTS", {}) + + def boom(repo_id): + raise AssertionError("Tier 3 must not run when Tier 2.5 has the answer") + + monkeypatch.setattr(lc, "_fetch_swa_entry_from_hf", boom) + b = _backend_from_gguf( + "gemma3", + dict(_SWA_FIELDS, block_count = 18), + general = {"general.source.huggingface.repository": "google/gemma-3-1b-it"}, + ) + assert b._sliding_window_pattern == [(i + 1) % 6 != 0 for i in range(18)] + with open(tmp_path / "swa_cache.json") as f: + assert json.load(f) == {"gemma3": 6} + + class TestGGUFParserReset: """Verify that fields are properly reset between parses.""" @@ -209,11 +695,19 @@ class TestGGUFParserReset: "block_count": 32, "attention.key_length": 128, "attention.kv_lora_rank": 512, + "attention.head_count_kv": [8, 2], + "attention.sliding_window_pattern": [True, False], + "attention.key_length_swa": 64, + "attention.value_length_swa": 64, "ssm.inner_size": 4096, }, ) assert b._kv_key_length == 128 assert b._kv_lora_rank == 512 + assert b._n_kv_heads_by_layer == [8, 2] + assert b._sliding_window_pattern == [True, False] + assert b._kv_key_length_swa == 64 + assert b._kv_value_length_swa == 64 assert b._ssm_inner_size == 4096 # Second parse without those fields -- they should be None @@ -230,6 +724,10 @@ class TestGGUFParserReset: os.unlink(path) assert b._kv_key_length is None assert b._kv_lora_rank is None + assert b._n_kv_heads_by_layer is None + assert b._sliding_window_pattern is None + assert b._kv_key_length_swa is None + assert b._kv_value_length_swa is None assert b._ssm_inner_size is None assert b._n_layers == 64 @@ -455,7 +953,9 @@ class TestSlidingWindowEstimation: n_global = max(1, 62 // 4) # 15 n_swa = 62 - n_global # 47 kv_per = 16 * (128 + 128) * 2 - expected = int(n_global * 131072 * kv_per + n_swa * min(131072, 1024) * kv_per) + # SWA cache is double-buffered: 2 * sliding_window cells, capped at n_ctx. + swa_cells = min(131072, 2 * 1024) + expected = int(n_global * 131072 * kv_per + n_swa * swa_cells * kv_per) assert b._estimate_kv_cache_bytes(131072, "f16") == expected def test_gpt_oss(self): @@ -472,27 +972,52 @@ class TestSlidingWindowEstimation: n_global = max(1, 24 // 4) # 6 n_swa = 24 - n_global # 18 kv_per = 8 * (64 + 64) * 2 - expected = int(n_global * 131072 * kv_per + n_swa * min(131072, 128) * kv_per) + swa_cells = min(131072, 2 * 128) + expected = int(n_global * 131072 * kv_per + n_swa * swa_cells * kv_per) assert b._estimate_kv_cache_bytes(131072, "f16") == expected + def test_gemma4_per_layer_swa_metadata(self): + b = self._swa_backend( + _n_layers = 30, + _n_kv_heads = None, + _n_kv_heads_by_layer = [8, 8, 8, 8, 8, 2] * 5, + _n_heads = 16, + _embedding_length = 2816, + _kv_key_length = 512, + _kv_value_length = 512, + _sliding_window = 1024, + _sliding_window_pattern = [True, True, True, True, True, False] * 5, + _kv_key_length_swa = 256, + _kv_value_length_swa = 256, + ) + + full_layers = 5 + sliding_layers = 25 + + def expected(ctx): + full = full_layers * ctx * 2 * (512 + 512) * 2 + sliding = sliding_layers * min(ctx, 2 * 1024) * 8 * (256 + 256) * 2 + return int(full + sliding) + + for ctx in (4096, 46500, 262144): + assert b._estimate_kv_cache_bytes(ctx, "f16") == expected(ctx) + def test_ctx_smaller_than_window(self): - """When context < sliding_window, SWA layers use full context anyway.""" + """When context < 2 * sliding_window, SWA cache caps at ctx.""" b = self._swa_backend(_sliding_window = 8192) n_global = max(1, 62 // 4) # 15 n_swa = 62 - n_global # 47 kv_per = 16 * (128 + 128) * 2 ctx = 4096 - expected = int(n_global * ctx * kv_per + n_swa * min(ctx, 8192) * kv_per) - # min(4096, 8192) = 4096, so both pools use full ctx + expected = int(n_global * ctx * kv_per + n_swa * min(ctx, 2 * 8192) * kv_per) assert b._estimate_kv_cache_bytes(ctx, "f16") == expected def test_odd_layer_count(self): - """Odd layer count: n_global = max(1, n//4), n_swa = n - n_global.""" b = self._swa_backend(_n_layers = 63) n_global = max(1, 63 // 4) # 15 n_swa = 63 - n_global # 48 kv_per = 16 * (128 + 128) * 2 - expected = int(n_global * 1000 * kv_per + n_swa * min(1000, 1024) * kv_per) + expected = int(n_global * 1000 * kv_per + n_swa * min(1000, 2 * 1024) * kv_per) assert b._estimate_kv_cache_bytes(1000, "f16") == expected @@ -785,6 +1310,686 @@ class TestEdgeCases: assert result == expected +# --------------------------------------------------------------------------- +# J2. Server-flag knobs (--swa-full, --kv-unified/--parallel, +# --ctx-checkpoints, --kv-offload) +# --------------------------------------------------------------------------- + + +class TestServerFlags: + """Estimator should mirror llama-server CLI flags that change KV size.""" + + def _swa_backend(self, **overrides): + defaults = { + "_n_layers": 26, + "_n_kv_heads": 4, + "_n_heads": 8, + "_embedding_length": 1152, + "_kv_key_length": 256, + "_kv_value_length": 256, + "_sliding_window": 512, + "_sliding_window_pattern": [True, True, True, True, True, False] * 4 + + [True, True], + } + defaults.update(overrides) + b = LlamaCppBackend() + for k, v in defaults.items(): + setattr(b, k, v) + return b + + def _gqa_backend(self, **overrides): + defaults = { + "_n_layers": 28, + "_n_kv_heads": 8, + "_n_heads": 16, + "_embedding_length": 1024, + "_kv_key_length": 128, + "_kv_value_length": 128, + } + defaults.update(overrides) + b = LlamaCppBackend() + for k, v in defaults.items(): + setattr(b, k, v) + return b + + # ── --swa-full ────────────────────────────────────────────────── + + def test_swa_full_collapses_pattern_path_to_full_ctx(self): + b = self._swa_backend() + ctx = 32_768 + flagged = b._estimate_kv_cache_bytes(ctx, "f16", swa_full = True) + # With swa_full, every layer caches n_ctx -- equals path 4 sizing. + kv_per_token = 4 * (256 + 256) * 2 # n_kv_heads * (k+v) * f16 + expected = 26 * ctx * kv_per_token + assert flagged == expected + assert flagged > b._estimate_kv_cache_bytes(ctx, "f16") + + def test_swa_full_collapses_legacy_path_to_full_ctx(self): + # No per-layer pattern -> 1/4-global heuristic; swa_full overrides. + b = self._swa_backend(_sliding_window_pattern = None) + ctx = 16_384 + flagged = b._estimate_kv_cache_bytes(ctx, "f16", swa_full = True) + n_global = max(1, 26 // 4) + n_swa = 26 - n_global + kv_per = 4 * (256 + 256) * 2 + # swa_cells == n_ctx when swa_full=True + expected = n_global * ctx * kv_per + n_swa * ctx * kv_per + assert flagged == expected + + def test_swa_full_no_op_for_non_swa_model(self): + b = self._gqa_backend() + baseline = b._estimate_kv_cache_bytes(8192, "f16") + flagged = b._estimate_kv_cache_bytes(8192, "f16", swa_full = True) + assert flagged == baseline + + def test_swa_full_suppresses_checkpoint_term(self): + b = self._swa_backend() + with_cp = b._estimate_kv_cache_bytes(8192, "f16", ctx_checkpoints = 8) + with_cp_full = b._estimate_kv_cache_bytes( + 8192, "f16", ctx_checkpoints = 8, swa_full = True + ) + no_cp_full = b._estimate_kv_cache_bytes(8192, "f16", swa_full = True) + # Checkpoints only matter when SWA layers don't already keep n_ctx. + assert with_cp_full == no_cp_full + assert with_cp > b._estimate_kv_cache_bytes(8192, "f16") + + # ── --parallel + --kv-unified ────────────────────────────────── + # Empirically verified against llama-server: non-SWA caches partition + # n_ctx across slots (total memory constant); SWA layers are the only + # portion that scales with --parallel. --kv-unified is currently a + # no-op for memory math (kept for API forward-compat). + + def test_gqa_kv_constant_across_parallel(self): + b = self._gqa_backend() + baseline = b._estimate_kv_cache_bytes(4096, "f16") + for slots in (1, 2, 4, 8): + for unified in (True, False): + assert ( + b._estimate_kv_cache_bytes( + 4096, "f16", n_parallel = slots, kv_unified = unified + ) + == baseline + ) + + def test_zero_parallel_floors_at_one(self): + b = self._gqa_backend() + baseline = b._estimate_kv_cache_bytes(4096, "f16") + for unified in (True, False): + assert ( + b._estimate_kv_cache_bytes( + 4096, "f16", n_parallel = 0, kv_unified = unified + ) + == baseline + ) + + def test_swa_path_scales_only_swa_portion(self): + b = self._swa_backend() + ctx = 8192 + baseline = b._estimate_kv_cache_bytes(ctx, "f16") + # Decompose baseline by walking the same loop the estimator does. + swa = b._sliding_window + per_token_global = 4 * (256 + 256) * 2 # n_kv * (k+v) * f16 + per_token_swa = 4 * (256 + 256) * 2 # k_swa/val_swa fall back + per_slot_swa_cells = min(ctx, 2 * swa) # not clamped at parallel=1 + global_bytes = sum( + ctx * per_token_global + for f in b._sliding_window_pattern[: b._n_layers] + if not f + ) + swa_bytes_per_slot = sum( + per_slot_swa_cells * per_token_swa + for f in b._sliding_window_pattern[: b._n_layers] + if f + ) + # Sanity: parallel=1 reproduces baseline exactly + assert global_bytes + swa_bytes_per_slot == baseline + # Only SWA portion scales by parallel + for slots in (1, 2, 3, 4): + scaled = b._estimate_kv_cache_bytes( + ctx, "f16", n_parallel = slots, kv_unified = False + ) + # SWA cells get clamped to per_slot_ctx when ctx/slots < 2*swa + per_slot_ctx = max(1, ctx // slots) + cells = min(ctx, 2 * swa, per_slot_ctx) + swa_bps = sum( + cells * per_token_swa + for f in b._sliding_window_pattern[: b._n_layers] + if f + ) + assert scaled == global_bytes + slots * swa_bps + + def test_mla_kv_constant_across_parallel(self): + b = LlamaCppBackend() + b._n_layers = 60 + b._n_kv_heads = 1 + b._kv_lora_rank = 512 + b._key_length_mla = 64 + b._kv_key_length = 576 + baseline = b._estimate_kv_cache_bytes(8192, "f16") + for slots in (1, 2, 4, 8): + for unified in (True, False): + assert ( + b._estimate_kv_cache_bytes( + 8192, "f16", n_parallel = slots, kv_unified = unified + ) + == baseline + ) + + # ── --ctx-checkpoints ────────────────────────────────────────── + + def test_ctx_checkpoints_zero_is_no_op(self): + b = self._swa_backend() + baseline = b._estimate_kv_cache_bytes(8192, "f16") + assert b._estimate_kv_cache_bytes(8192, "f16", ctx_checkpoints = 0) == baseline + + def test_ctx_checkpoints_no_op_for_non_swa(self): + b = self._gqa_backend() + baseline = b._estimate_kv_cache_bytes(8192, "f16") + assert b._estimate_kv_cache_bytes(8192, "f16", ctx_checkpoints = 32) == baseline + + def test_ctx_checkpoints_pattern_path_adds_known_bytes(self): + b = self._swa_backend() + ctx = 8192 + baseline = b._estimate_kv_cache_bytes(ctx, "f16") + flagged = b._estimate_kv_cache_bytes(ctx, "f16", ctx_checkpoints = 4) + # 22 SWA layers * 4 checkpoints * 512 cells * 4 heads * (256+256) * 2 bytes + n_swa_layers = sum( + 1 for f in [True, True, True, True, True, False] * 4 + [True, True] if f + ) + per_layer = 4 * 512 * 4 * (256 + 256) * 2 + assert flagged == baseline + n_swa_layers * per_layer + + def test_ctx_checkpoints_legacy_path_adds_known_bytes(self): + b = self._swa_backend(_sliding_window_pattern = None) + ctx = 8192 + baseline = b._estimate_kv_cache_bytes(ctx, "f16") + flagged = b._estimate_kv_cache_bytes(ctx, "f16", ctx_checkpoints = 4) + n_global = max(1, 26 // 4) + n_swa = 26 - n_global + kv_per = 4 * (256 + 256) * 2 + extra = 4 * n_swa * 512 * kv_per # ctx_checkpoints * n_swa * sliding * kv_per + assert flagged == baseline + extra + + def test_ctx_checkpoints_compose_with_n_parallel(self): + # Only the SWA + checkpoint portion scales by n_parallel; the + # global-layer portion stays constant. + b = self._swa_backend() + ctx = 8192 + swa = b._sliding_window + per_token = 4 * (256 + 256) * 2 + global_bytes = sum( + ctx * per_token for f in b._sliding_window_pattern[: b._n_layers] if not f + ) + n_swa_layers = sum(1 for f in b._sliding_window_pattern[: b._n_layers] if f) + slots = 3 + per_slot_ctx = max(1, ctx // slots) + swa_cells = min(ctx, 2 * swa, per_slot_ctx) + swa_bytes_per_slot = n_swa_layers * swa_cells * per_token + cp_extra_per_slot = n_swa_layers * 4 * swa * per_token # 4 checkpoints + flagged = b._estimate_kv_cache_bytes( + ctx, "f16", ctx_checkpoints = 4, n_parallel = slots, kv_unified = False + ) + assert flagged == global_bytes + slots * ( + swa_bytes_per_slot + cp_extra_per_slot + ) + + # ── --kv-offload (kv_on_gpu) ─────────────────────────────────── + + def test_fit_returns_requested_when_kv_off_gpu(self): + b = self._gqa_backend() + # Tiny VRAM budget -- normally would force a reduction. + fitted = b._fit_context_to_vram( + requested_ctx = 32_768, + available_mib = 1, + model_size_bytes = 100, + cache_type_kv = "f16", + kv_on_gpu = False, + ) + assert fitted == 32_768 + + def test_fit_reduces_when_kv_on_gpu(self): + b = self._gqa_backend() + fitted = b._fit_context_to_vram( + requested_ctx = 32_768, + available_mib = 64, + model_size_bytes = 1024 * 1024, # 1 MiB + cache_type_kv = "f16", + kv_on_gpu = True, + ) + assert fitted < 32_768 + + def test_fit_threads_swa_full_through_estimator(self): + # SWA model, generous budget; both should fit but cache size differs. + b = self._swa_backend() + ctx = 8192 + kv_default = b._estimate_kv_cache_bytes(ctx, "f16") + kv_full = b._estimate_kv_cache_bytes(ctx, "f16", swa_full = True) + assert kv_full > kv_default + # Budget = model + kv_default (rounded up) -- swa_full should not fit. + budget_mib = (1024 * 1024 + kv_default) / (1024 * 1024) / 0.90 + 1 + fitted_default = b._fit_context_to_vram( + requested_ctx = ctx, + available_mib = int(budget_mib), + model_size_bytes = 1024 * 1024, + cache_type_kv = "f16", + ) + fitted_full = b._fit_context_to_vram( + requested_ctx = ctx, + available_mib = int(budget_mib), + model_size_bytes = 1024 * 1024, + cache_type_kv = "f16", + swa_full = True, + ) + assert fitted_default == ctx + assert fitted_full < ctx + + +# --------------------------------------------------------------------------- +# J2.5. --parallel N memory accounting (per-layer-type scaling rule) +# --------------------------------------------------------------------------- + + +class TestParallelSWAScaling: + """Verifies the per-layer-type scaling rule against the closed form + measured from llama-server. Empirical formula on Gemma-3 270m at + ctx=8192: total_kv = 24 + parallel * 15 (MiB). + + Rule (verified vs ``llama-server`` log on real GGUFs): + * non-SWA layers: total cells = n_ctx, partitioned across slots, + memory CONSTANT in n_parallel. + * SWA layers: per-slot cells = 2 * sliding_window (clamped at + n_ctx and at per_slot_ctx); memory LINEAR in n_parallel. + * --kv-unified is a no-op for memory math; both modes yield the + same total in measured cases. + """ + + def _gqa_backend(self, **overrides): + defaults = { + "_n_layers": 28, + "_n_kv_heads": 8, + "_n_heads": 16, + "_embedding_length": 1024, + "_kv_key_length": 128, + "_kv_value_length": 128, + } + defaults.update(overrides) + b = LlamaCppBackend() + for k, v in defaults.items(): + setattr(b, k, v) + return b + + def _swa_backend(self, **overrides): + defaults = { + "_n_layers": 18, + "_n_kv_heads": 1, + "_n_heads": 4, + "_embedding_length": 1024, + "_kv_key_length": 256, + "_kv_value_length": 256, + "_sliding_window": 512, + # 15 SWA + 3 global, mirrors gemma-3-270m + "_sliding_window_pattern": [ + t == "swa" for t in (["swa"] * 5 + ["global"]) * 3 + ], + } + defaults.update(overrides) + b = LlamaCppBackend() + for k, v in defaults.items(): + setattr(b, k, v) + return b + + # ── non-SWA paths: constant ──────────────────────────────────── + + def test_pure_gqa_constant_across_parallel(self): + b = self._gqa_backend() + baseline = b._estimate_kv_cache_bytes(8192, "f16") + for slots in (1, 2, 4, 8): + for unified in (True, False): + assert ( + b._estimate_kv_cache_bytes( + 8192, "f16", n_parallel = slots, kv_unified = unified + ) + == baseline + ) + + def test_mla_constant_across_parallel(self): + b = LlamaCppBackend() + b._n_layers = 60 + b._n_kv_heads = 1 + b._kv_lora_rank = 512 + b._key_length_mla = 64 + b._kv_key_length = 576 + baseline = b._estimate_kv_cache_bytes(8192, "f16") + for slots in (1, 2, 4, 8): + assert b._estimate_kv_cache_bytes(8192, "f16", n_parallel = slots) == baseline + + def test_hybrid_constant_across_parallel(self): + b = LlamaCppBackend() + b._n_layers = 64 + b._n_kv_heads = 16 + b._n_heads = 32 + b._embedding_length = 4096 + b._kv_key_length = 128 + b._kv_value_length = 128 + b._ssm_inner_size = 4096 + b._full_attention_interval = 4 + baseline = b._estimate_kv_cache_bytes(8192, "f16") + for slots in (1, 2, 4, 8): + assert b._estimate_kv_cache_bytes(8192, "f16", n_parallel = slots) == baseline + + def test_legacy_constant_across_parallel(self): + b = LlamaCppBackend() + b._n_layers = 32 + b._n_kv_heads = 8 + b._n_heads = 8 + b._embedding_length = 4096 + baseline = b._estimate_kv_cache_bytes(8192, "f16") + for slots in (1, 2, 4, 8): + assert b._estimate_kv_cache_bytes(8192, "f16", n_parallel = slots) == baseline + + # ── SWA paths: scale only the SWA portion ────────────────────── + + def test_swa_pattern_scales_only_swa_portion(self): + b = self._swa_backend() + ctx = 8192 + swa = b._sliding_window + per_token = 1 * (256 + 256) * 2 # n_kv * (k+v) * f16 + n_global = sum(1 for f in b._sliding_window_pattern if not f) + n_swa = sum(1 for f in b._sliding_window_pattern if f) + global_bytes = n_global * ctx * per_token + for slots in (1, 2, 4, 8): + per_slot_ctx = max(1, ctx // slots) + cells = min(ctx, 2 * swa, per_slot_ctx) + swa_bps = n_swa * cells * per_token + for unified in (True, False): + got = b._estimate_kv_cache_bytes( + ctx, "f16", n_parallel = slots, kv_unified = unified + ) + assert got == global_bytes + slots * swa_bps + + def test_swa_fallback_scales_only_swa_portion(self): + # No per-layer pattern -> 1/4-global heuristic. + b = self._swa_backend(_sliding_window_pattern = None) + ctx = 8192 + swa = b._sliding_window + n_layers = 18 + n_global = max(1, n_layers // 4) + n_swa = n_layers - n_global + per_token = 1 * (256 + 256) * 2 + global_bytes = n_global * ctx * per_token + for slots in (1, 2, 4, 8): + per_slot_ctx = max(1, ctx // slots) + cells = min(ctx, 2 * swa, per_slot_ctx) + swa_bps = n_swa * cells * per_token + got = b._estimate_kv_cache_bytes(ctx, "f16", n_parallel = slots) + assert got == global_bytes + slots * swa_bps + + def test_swa_per_slot_clamped_when_ctx_lt_slots_x_2window(self): + # ctx=4096 / slots=8 -> per_slot_ctx=512, but 2*sliding=1024. + # SWA cells should clamp at per_slot_ctx (512), not 2*sliding. + b = self._swa_backend() + ctx = 4096 + per_slot_ctx_at_8 = ctx // 8 + assert per_slot_ctx_at_8 < 2 * b._sliding_window + # Build expected with the clamped formula + n_swa = sum(1 for f in b._sliding_window_pattern if f) + n_global = sum(1 for f in b._sliding_window_pattern if not f) + per_token = 1 * (256 + 256) * 2 + global_bytes = n_global * ctx * per_token + cells = min(ctx, 2 * b._sliding_window, per_slot_ctx_at_8) + assert cells == per_slot_ctx_at_8 + expected = global_bytes + 8 * (n_swa * cells * per_token) + assert b._estimate_kv_cache_bytes(ctx, "f16", n_parallel = 8) == expected + + def test_swa_full_does_not_scale_under_parallel(self): + # swa_full forces every layer to n_ctx; result is the all-global + # GQA-style total, which is constant in parallel. + b = self._swa_backend() + ctx = 8192 + baseline = b._estimate_kv_cache_bytes(ctx, "f16", swa_full = True) + for slots in (1, 2, 4, 8): + assert ( + b._estimate_kv_cache_bytes(ctx, "f16", swa_full = True, n_parallel = slots) + == baseline + ) + + # ── kv_unified: no-op for memory math ────────────────────────── + + def test_kv_unified_is_no_op_for_memory_math(self): + # Both unified=True and unified=False must produce the same + # total bytes for every backend type and every parallel value. + backends = [ + ("gqa", self._gqa_backend()), + ("swa", self._swa_backend()), + ] + for label, b in backends: + for slots in (1, 2, 4, 8): + u = b._estimate_kv_cache_bytes( + 8192, "f16", n_parallel = slots, kv_unified = True + ) + nu = b._estimate_kv_cache_bytes( + 8192, "f16", n_parallel = slots, kv_unified = False + ) + assert u == nu, f"{label} parallel={slots} unified-mismatch" + + # ── Empirical Gemma-3 270m formula ───────────────────────────── + + def test_matches_empirical_gemma3_270m_formula(self): + """Exact match against the formula measured from llama-server: + total_kv = 24 + parallel * 15 (MiB) at ctx=8192. + + Geometry: 18 layers (3 global + 15 SWA), n_kv=1, head_dim=256, + sliding=512, f16. + """ + b = LlamaCppBackend() + b._n_layers = 18 + b._n_kv_heads = 1 + b._n_heads = 4 + b._embedding_length = 1024 + b._kv_key_length = 256 + b._kv_value_length = 256 + b._sliding_window = 512 + # 5-period [swa,swa,swa,swa,full] * 3 + [swa,swa,swa]: mirrors the + # bootstrap-resolved pattern for gemma3 (period 6) on an 18-layer + # model (15 SWA, 3 global). + b._sliding_window_pattern = [(i + 1) % 6 != 0 for i in range(18)] + n_global = 3 + n_swa = 15 + # Confirm pattern shape + assert sum(b._sliding_window_pattern) == n_swa + for slots, expected_mib in [(1, 39), (2, 54), (4, 84)]: + got_bytes = b._estimate_kv_cache_bytes(8192, "f16", n_parallel = slots) + got_mib = got_bytes / (1024 * 1024) + assert ( + got_mib == expected_mib + ), f"slots={slots}: got {got_mib} MiB, expected {expected_mib} MiB" + + +# --------------------------------------------------------------------------- +# J3. shared_kv_layers (Gemma 3n / Gemma 4) +# --------------------------------------------------------------------------- + + +class TestSharedKVLayers: + """``.attention.shared_kv_layers`` reduces the layer count that + actually allocates KV. The trailing ``shared_kv_layers`` blocks reuse + earlier caches (Gemma 3n: 35 layers, 15 shared -> 20 allocate; Gemma 4 + same field). Unset on every other arch -> no behavioural change.""" + + def _gemma3n_backend(self, **overrides): + # Mirrors google/gemma-3n-E4B-it: 35 layers, 15 shared, + # SWA window 1024, period 5 (4 sliding + 1 full repeating). + defaults = { + "_n_layers": 35, + "_n_kv_heads": 4, + "_n_heads": 8, + "_embedding_length": 2048, + "_kv_key_length": 256, + "_kv_value_length": 256, + "_sliding_window": 1024, + "_sliding_window_pattern": [ + t == "sliding_attention" + for t in (["sliding_attention"] * 4 + ["full_attention"]) * 7 + ], + "_shared_kv_layers": 15, + } + defaults.update(overrides) + b = LlamaCppBackend() + for k, v in defaults.items(): + setattr(b, k, v) + return b + + def _gqa_backend(self, **overrides): + defaults = { + "_n_layers": 28, + "_n_kv_heads": 8, + "_n_heads": 16, + "_embedding_length": 1024, + "_kv_key_length": 128, + "_kv_value_length": 128, + } + defaults.update(overrides) + b = LlamaCppBackend() + for k, v in defaults.items(): + setattr(b, k, v) + return b + + def test_field_initialises_to_none(self): + b = LlamaCppBackend() + assert b._shared_kv_layers is None + + def test_unset_field_is_noop(self): + b = self._gqa_backend() + baseline = b._estimate_kv_cache_bytes(8192, "f16") + b._shared_kv_layers = None + assert b._estimate_kv_cache_bytes(8192, "f16") == baseline + b._shared_kv_layers = 0 + assert b._estimate_kv_cache_bytes(8192, "f16") == baseline + + def test_path4_drops_shared_layers(self): + b = self._gqa_backend(_shared_kv_layers = 4) + ctx = 4096 + kv_per = 8 * (128 + 128) * 2 + # 28 - 4 = 24 layers actually allocate + assert b._estimate_kv_cache_bytes(ctx, "f16") == 24 * ctx * kv_per + + def test_path5_drops_shared_layers(self): + b = LlamaCppBackend() + b._n_layers = 32 + b._n_kv_heads = 8 + b._n_heads = 8 + b._embedding_length = 4096 + b._shared_kv_layers = 8 + ctx = 4096 + head_dim = 4096 // 8 # 512 + # 32 - 8 = 24 layers + expected = 2 * 8 * head_dim * 24 * ctx * 2 + assert b._estimate_kv_cache_bytes(ctx, "f16") == expected + + def test_path1_mla_drops_shared_layers(self): + b = LlamaCppBackend() + b._n_layers = 60 + b._n_kv_heads = 1 + b._kv_lora_rank = 512 + b._key_length_mla = 64 + b._kv_key_length = 576 + b._shared_kv_layers = 10 + ctx = 8192 + # 60 - 10 = 50 + assert b._estimate_kv_cache_bytes(ctx, "f16") == 50 * ctx * 1 * 576 * 2 + + def test_path3_pattern_loops_only_unshared_layers(self): + b = self._gemma3n_backend() + ctx = 8192 + # First 20 layers contribute; layers 20..34 are skipped. + # Pattern: [s,s,s,s,F] repeated. In layers 0..19: + # sliding: 16, full: 4 + sliding_in_unshared = sum(b._sliding_window_pattern[:20]) + full_in_unshared = 20 - sliding_in_unshared + assert sliding_in_unshared == 16 + assert full_in_unshared == 4 + kv_per = 4 * (256 + 256) * 2 + swa_cells = min(ctx, 2 * 1024) + expected = ( + full_in_unshared * ctx * kv_per + sliding_in_unshared * swa_cells * kv_per + ) + assert b._estimate_kv_cache_bytes(ctx, "f16") == expected + + def test_shared_layers_reduces_estimate(self): + b = self._gemma3n_backend() + with_shared = b._estimate_kv_cache_bytes(8192, "f16") + b._shared_kv_layers = 0 + without_shared = b._estimate_kv_cache_bytes(8192, "f16") + # 20/35 = 0.571 of the work; expect ~43% reduction. + ratio = with_shared / without_shared + assert 0.5 < ratio < 0.65 + + def test_path3_pattern_with_swa_full_and_shared(self): + b = self._gemma3n_backend() + ctx = 8192 + flagged = b._estimate_kv_cache_bytes(ctx, "f16", swa_full = True) + # Every unshared layer caches n_ctx; equals path-4-style sizing + # over only the 20 unshared layers. + kv_per = 4 * (256 + 256) * 2 + assert flagged == 20 * ctx * kv_per + + def test_path3_fallback_uses_unshared_count(self): + # No per-layer pattern -> 1/4-global heuristic over n_layers_kv, + # not n_layers. + b = self._gemma3n_backend(_sliding_window_pattern = None) + ctx = 8192 + n_layers_kv = 35 - 15 # 20 + n_global = max(1, n_layers_kv // 4) # 5 + n_swa = n_layers_kv - n_global # 15 + kv_per = 4 * (256 + 256) * 2 + swa_cells = min(ctx, 2 * 1024) + expected = n_global * ctx * kv_per + n_swa * swa_cells * kv_per + assert b._estimate_kv_cache_bytes(ctx, "f16") == expected + + def test_shared_floors_at_one_layer(self): + # Pathological: shared >= n_layers should not zero out the cache. + b = self._gqa_backend(_shared_kv_layers = 99) + ctx = 4096 + kv_per = 8 * (128 + 128) * 2 + assert b._estimate_kv_cache_bytes(ctx, "f16") == 1 * ctx * kv_per + + def test_composes_with_n_parallel(self): + # Only the SWA portion of the unshared layers scales by n_parallel; + # the global portion stays constant. + b = self._gemma3n_backend() + ctx = 8192 + swa = b._sliding_window + per_token = 4 * (256 + 256) * 2 + unshared_pattern = b._sliding_window_pattern[:20] # 35 - 15 shared + sliding_in_unshared = sum(unshared_pattern) + global_in_unshared = len(unshared_pattern) - sliding_in_unshared + global_bytes = global_in_unshared * ctx * per_token + slots = 3 + per_slot_ctx = max(1, ctx // slots) + swa_cells = min(ctx, 2 * swa, per_slot_ctx) + swa_bytes_per_slot = sliding_in_unshared * swa_cells * per_token + flagged = b._estimate_kv_cache_bytes( + ctx, "f16", n_parallel = slots, kv_unified = False + ) + assert flagged == global_bytes + slots * swa_bytes_per_slot + + def test_composes_with_ctx_checkpoints(self): + b = self._gemma3n_backend() + ctx = 8192 + baseline = b._estimate_kv_cache_bytes(ctx, "f16") + with_cp = b._estimate_kv_cache_bytes(ctx, "f16", ctx_checkpoints = 4) + # Checkpoints only count over UNSHARED SWA layers (16 of them). + sliding_in_unshared = sum(b._sliding_window_pattern[:20]) + per_cp_layer = 4 * 1024 * 4 * (256 + 256) * 2 # cps * swa * heads * (k+v) * bpe + assert with_cp == baseline + sliding_in_unshared * per_cp_layer + + def test_unload_resets_shared_kv_layers(self): + b = LlamaCppBackend() + b._shared_kv_layers = 12 + b.unload_model() + assert b._shared_kv_layers is None + + # --------------------------------------------------------------------------- # K. Lifecycle Tests # --------------------------------------------------------------------------- @@ -799,13 +2004,18 @@ class TestLifecycle: "_kv_key_length", "_kv_value_length", "_sliding_window", + "_sliding_window_pattern", "_full_attention_interval", "_kv_lora_rank", "_key_length_mla", + "_kv_key_length_swa", + "_kv_value_length_swa", "_ssm_inner_size", "_ssm_state_size", + "_shared_kv_layers", ]: assert getattr(b, attr) is None + assert b._n_kv_heads_by_layer is None def test_unload_resets_fields(self): b = LlamaCppBackend() @@ -813,20 +2023,30 @@ class TestLifecycle: b._kv_key_length = 128 b._kv_lora_rank = 512 b._sliding_window = 1024 + b._sliding_window_pattern = [True, False] + b._n_kv_heads_by_layer = [8, 2] + b._kv_key_length_swa = 64 + b._kv_value_length_swa = 64 b._ssm_inner_size = 4096 b._full_attention_interval = 4 + b._shared_kv_layers = 8 b.unload_model() for attr in [ "_kv_key_length", "_kv_value_length", "_sliding_window", + "_sliding_window_pattern", "_full_attention_interval", "_kv_lora_rank", "_key_length_mla", + "_kv_key_length_swa", + "_kv_value_length_swa", "_ssm_inner_size", "_ssm_state_size", + "_shared_kv_layers", ]: assert getattr(b, attr) is None + assert b._n_kv_heads_by_layer is None def test_end_to_end_synthetic_mla(self): """Full round-trip: write GGUF -> parse -> estimate.""" @@ -887,12 +2107,46 @@ class TestLifecycle: ) assert b._can_estimate_kv() result = b._estimate_kv_cache_bytes(131072, "f16") - n_global = max(1, 62 // 4) # 15 - n_swa = 62 - n_global # 47 + # gemma3 -> period 6 from the bootstrap table, SWA cache + # double-buffered to 2 * sliding_window cells. + period = 6 kv_per = 16 * 256 * 2 - expected = int(n_global * 131072 * kv_per + n_swa * 1024 * kv_per) + expected = 0 + for i in range(62): + is_swa = (i + 1) % period != 0 + layer_ctx = min(131072, 2 * 1024) if is_swa else 131072 + expected += layer_ctx * kv_per assert result == expected + def test_end_to_end_synthetic_shared_kv_round_trip(self): + # Mirrors gemma3n_text: 35 layers, 15 shared, sliding_window=1024. + b = _backend_from_gguf( + "gemma3n_text", + { + "context_length": 32768, + "block_count": 35, + "attention.head_count_kv": 4, + "attention.head_count": 8, + "embedding_length": 2048, + "attention.key_length": 256, + "attention.value_length": 256, + "attention.sliding_window": 1024, + "attention.shared_kv_layers": 15, + }, + ) + assert b._can_estimate_kv() + assert b._shared_kv_layers == 15 + # Bootstrap table for gemma3n_text -> period 5; the resolver + # synthesises a 35-entry bool array. The first 20 entries + # (n_layers - shared) are the only ones that allocate KV. + result = b._estimate_kv_cache_bytes(8192, "f16") + assert result > 0 + # Sanity: setting shared back to 0 must produce a strictly larger + # estimate (more layers allocate). + b._shared_kv_layers = 0 + unshared = b._estimate_kv_cache_bytes(8192, "f16") + assert unshared > result + def test_end_to_end_synthetic_gqa(self): b = _backend_from_gguf( "qwen3", diff --git a/studio/backend/tests/test_llama_cpp_context_fit.py b/studio/backend/tests/test_llama_cpp_context_fit.py index f498655347..caa6397901 100644 --- a/studio/backend/tests/test_llama_cpp_context_fit.py +++ b/studio/backend/tests/test_llama_cpp_context_fit.py @@ -114,9 +114,13 @@ def _make_backend( inst._kv_value_length = kv_value_length inst._kv_lora_rank = None inst._sliding_window = None + inst._sliding_window_pattern = None inst._ssm_inner_size = None inst._full_attention_interval = None inst._key_length_mla = None + inst._n_kv_heads_by_layer = None + inst._kv_key_length_swa = None + inst._kv_value_length_swa = None return inst @@ -137,7 +141,7 @@ def _drive( model_size = int(model_gib * GIB) cache_type_kv = None - def fake_estimate(n_ctx_, _type = None): + def fake_estimate(n_ctx_, _type = None, **_kwargs): return 0 if n_ctx_ <= 0 else n_ctx_ * kv_per_token_bytes inst._estimate_kv_cache_bytes = fake_estimate diff --git a/studio/backend/tests/test_llama_cpp_max_context_threshold.py b/studio/backend/tests/test_llama_cpp_max_context_threshold.py index 5fd0243c9f..22e4cda7d1 100644 --- a/studio/backend/tests/test_llama_cpp_max_context_threshold.py +++ b/studio/backend/tests/test_llama_cpp_max_context_threshold.py @@ -99,9 +99,13 @@ def _make_backend(native_ctx = 131072): inst._kv_value_length = 128 inst._kv_lora_rank = None inst._sliding_window = None + inst._sliding_window_pattern = None inst._ssm_inner_size = None inst._full_attention_interval = None inst._key_length_mla = None + inst._n_kv_heads_by_layer = None + inst._kv_key_length_swa = None + inst._kv_value_length_swa = None return inst @@ -114,7 +118,7 @@ def _compute_max_available_ctx(native_ctx, model_gib, gpus, kv_per_token_bytes = model_size = int(model_gib * GIB) inst._estimate_kv_cache_bytes = ( - lambda n, _t = None: 0 if n <= 0 else n * kv_per_token_bytes + lambda n, _t = None, **_kw: 0 if n <= 0 else n * kv_per_token_bytes ) inst._can_estimate_kv = lambda: True diff --git a/studio/backend/tests/test_llama_server_args.py b/studio/backend/tests/test_llama_server_args.py new file mode 100644 index 0000000000..351fbd014d --- /dev/null +++ b/studio/backend/tests/test_llama_server_args.py @@ -0,0 +1,189 @@ +# SPDX-License-Identifier: AGPL-3.0-only +# Copyright 2026-present the Unsloth AI Inc. team. All rights reserved. See /studio/LICENSE.AGPL-3.0 + +"""Unit tests for the llama-server pass-through args validator. + +The validator is the security boundary between user-supplied CLI / HTTP +input and the llama-server subprocess command. These tests pin the +denylist behavior so the boundary doesn't quietly regress when new +managed flags are added. +""" + +from __future__ import annotations + +import pytest + +from core.inference.llama_server_args import ( + is_managed_flag, + validate_extra_args, +) + + +# ── Pass-through (allowed) ─────────────────────────────────────────── + + +@pytest.mark.parametrize( + "args", + [ + # Sampling + ["--top-k", "20"], + ["--top-p", "0.9", "--min-p", "0.05"], + ["--seed", "-1"], # negative value, not a flag + ["--temp", "0.0"], + ["--repeat-penalty", "1.05"], + ["--mirostat", "2", "--mirostat-lr", "0.1"], + ["--xtc-probability", "0.05", "--xtc-threshold", "0.1"], + ["--dry-multiplier", "0.5"], + # Tier-2 knobs that map to LoadRequest fields + ["--cache-type-k", "q8_0"], + ["--cache-type-v", "q8_0"], + ["--chat-template-file", "/tmp/tpl.jinja"], + ["--chat-template-kwargs", '{"reasoning_effort":"high"}'], + ["--spec-type", "ngram-mod"], + ["--spec-default"], + # Reasoning controls + ["--reasoning-format", "deepseek"], + ["-rea", "auto"], + # Soft-managed flags the user may want to override on the CLI; + # llama.cpp's last-wins parsing means these win over Studio's + # auto-set version. + ["-c", "131072"], + ["--ctx-size", "8192"], + ["--parallel", "1"], + ["-np", "8"], + ["--flash-attn", "off"], + ["-fa", "on"], + ["--no-context-shift"], + ["--context-shift"], + ["--jinja"], + ["--no-jinja"], + ["-ngl", "-1"], + ["--gpu-layers", "32"], + ["-t", "16"], + ["--threads", "32"], + ["-fit", "off"], + ["--fit", "on"], + ["--fit-ctx", "8192"], + ], +) +def test_pass_through_allowed(args): + assert validate_extra_args(args) == args + + +def test_none_returns_empty_list(): + assert validate_extra_args(None) == [] + + +def test_empty_list_returns_empty_list(): + assert validate_extra_args([]) == [] + + +def test_value_with_equals_form_passes_through(): + assert validate_extra_args(["--top-k=20"]) == ["--top-k=20"] + + +def test_non_flag_token_passes_through(): + # A bare positional value (not preceded by a flag) is preserved + # verbatim. llama-server may reject it, but that's not our job. + assert validate_extra_args(["foo"]) == ["foo"] + + +# ── Denylist (rejected) ────────────────────────────────────────────── + + +@pytest.mark.parametrize( + "denied", + [ + # Model identity + "-m", + "--model", + "-hf", + "-hfr", + "--hf-repo", + "-hff", + "--hf-file", + "-hft", + "--hf-token", + "-mm", + "--mmproj", + "--mmproj-url", + # Networking (Studio binds + proxies) + "--host", + "--port", + "--path", + "--api-prefix", + "--reuse-port", + # Auth / TLS + "--api-key", + "--api-key-file", + "--ssl-key-file", + "--ssl-cert-file", + # Single-model server + "--webui", + "--no-webui", + "--models-dir", + "--models-max", + ], +) +def test_denylist_rejects_all_aliases(denied): + with pytest.raises(ValueError, match = denied): + validate_extra_args([denied, "value"]) + + +def test_denylist_rejects_equals_form(): + with pytest.raises(ValueError, match = "--port"): + validate_extra_args(["--port=9000"]) + + +def test_denylist_rejects_short_form_when_long_is_denied(): + # -m is the short form of the hard-denied --model; rejecting only + # the long form would leave a trivial bypass. + with pytest.raises(ValueError, match = "-m"): + validate_extra_args(["-m", "/some/other/path.gguf"]) + + +def test_denylist_message_names_offending_flag(): + with pytest.raises(ValueError) as excinfo: + validate_extra_args(["--top-k", "20", "--api-key", "secret"]) + assert "--api-key" in str(excinfo.value) + + +def test_first_denied_flag_short_circuits(): + # Validation stops at the first denied flag; later denied flags + # in the same call don't matter for behaviour, but the message + # should name the first one we hit. + with pytest.raises(ValueError, match = "--port"): + validate_extra_args(["--port", "1", "--host", "x"]) + + +# ── Numeric values that look flag-ish ───────────────────────────────── + + +@pytest.mark.parametrize("value", ["-1", "-0.5", "-42", "-.5"]) +def test_negative_number_value_is_not_flag(value): + # ``--seed -1`` is a value, not a flag. Validator must not try + # to look up "-1" in the denylist. + assert validate_extra_args(["--seed", value]) == ["--seed", value] + + +# ── is_managed_flag helper ─────────────────────────────────────────── + + +def test_is_managed_flag_true_for_denied(): + assert is_managed_flag("--port") is True + assert is_managed_flag("--api-key") is True + assert is_managed_flag("-m") is True + assert is_managed_flag("--model") is True + + +def test_is_managed_flag_false_for_pass_through(): + assert is_managed_flag("--top-k") is False + assert is_managed_flag("--cache-type-k") is False + assert is_managed_flag("--chat-template-file") is False + # Soft-managed flags pass through (last-wins override) + assert is_managed_flag("-c") is False + assert is_managed_flag("--ctx-size") is False + assert is_managed_flag("--parallel") is False + assert is_managed_flag("--flash-attn") is False + assert is_managed_flag("-ngl") is False + assert is_managed_flag("--threads") is False diff --git a/studio/backend/tests/test_native_context_length.py b/studio/backend/tests/test_native_context_length.py index 7c69e56f89..60622c776d 100644 --- a/studio/backend/tests/test_native_context_length.py +++ b/studio/backend/tests/test_native_context_length.py @@ -320,11 +320,23 @@ class TestPydanticModels: """Field exists in InferenceStatusResponse.model_fields.""" assert "native_context_length" in InferenceStatusResponse.model_fields + def test_status_response_has_chat_template_field(self): + """Status includes chat_template so the UI can rehydrate after refresh.""" + assert "chat_template" in InferenceStatusResponse.model_fields + def test_status_response_defaults_none(self): """Omitting native_context_length defaults to None.""" resp = InferenceStatusResponse() assert resp.native_context_length is None + def test_status_response_chat_template_roundtrip(self): + """chat_template serializes and validates as part of status.""" + resp = InferenceStatusResponse(chat_template = "{{ messages }}") + roundtripped = InferenceStatusResponse.model_validate_json( + resp.model_dump_json() + ) + assert roundtripped.chat_template == "{{ messages }}" + def test_roundtrip_preserves_value(self): """model_validate_json(model_dump_json()) round-trips.""" resp = LoadResponse( diff --git a/studio/backend/tests/test_openai_tool_passthrough.py b/studio/backend/tests/test_openai_tool_passthrough.py index ccb0dba325..cdb7f5d270 100644 --- a/studio/backend/tests/test_openai_tool_passthrough.py +++ b/studio/backend/tests/test_openai_tool_passthrough.py @@ -144,13 +144,14 @@ class TestChatMessageToolRoles: # ── Role-aware content requirements ──────────────────────────── - def test_user_empty_content_rejected(self): - with pytest.raises(ValidationError): - ChatMessage(role = "user", content = "") + @pytest.mark.parametrize("role", ["user", "system"]) + def test_empty_string_content_allowed(self, role): + msg = ChatMessage(role = role, content = "") + assert msg.content == "" - def test_system_empty_content_rejected(self): + def test_user_missing_content_rejected(self): with pytest.raises(ValidationError): - ChatMessage(role = "system", content = "") + ChatMessage(role = "user") def test_user_empty_list_content_rejected(self): with pytest.raises(ValidationError): @@ -226,6 +227,14 @@ class TestChatCompletionRequestToolFields: assert len(req.tools) == 1 assert req.tools[0]["function"]["name"] == "get_weather" + def test_image_base64_allows_empty_user_text(self): + req = ChatCompletionRequest( + messages = [{"role": "user", "content": ""}], + image_base64 = "aW1hZ2U=", + ) + assert req.messages[0].content == "" + assert req.image_base64 == "aW1hZ2U=" + def test_tool_choice_string_auto(self): assert self._make(tool_choice = "auto").tool_choice == "auto" diff --git a/studio/backend/tests/test_tool_policy_gates.py b/studio/backend/tests/test_tool_policy_gates.py new file mode 100644 index 0000000000..01f6bbbc3f --- /dev/null +++ b/studio/backend/tests/test_tool_policy_gates.py @@ -0,0 +1,56 @@ +# SPDX-License-Identifier: AGPL-3.0-only +# Copyright 2026-present the Unsloth AI Inc. team. All rights reserved. + +""" +Tests for `_effective_enable_tools` -- the helper that folds the +process-level `tool_policy` over a request's `enable_tools` field. + +Truth table (policy x payload.enable_tools -> effective): + policy=None + payload=None -> None + policy=None + payload=True -> True + policy=None + payload=False -> False + policy=True + payload=* -> True + policy=False + payload=* -> False +""" + +import os +import sys +from types import SimpleNamespace + +_backend = os.path.join(os.path.dirname(__file__), "..") +sys.path.insert(0, _backend) + +import pytest + +from routes.inference import _effective_enable_tools +from state.tool_policy import reset_tool_policy, set_tool_policy + + +@pytest.fixture(autouse = True) +def _reset(): + reset_tool_policy() + yield + reset_tool_policy() + + +def _payload(value): + return SimpleNamespace(enable_tools = value) + + +class TestEffectiveEnableTools: + @pytest.mark.parametrize( + "payload_value,expected", + [(None, None), (True, True), (False, False)], + ) + def test_no_policy_falls_through_to_payload(self, payload_value, expected): + assert _effective_enable_tools(_payload(payload_value)) == expected + + @pytest.mark.parametrize("payload_value", [None, True, False]) + def test_policy_true_overrides_any_payload(self, payload_value): + set_tool_policy(True) + assert _effective_enable_tools(_payload(payload_value)) is True + + @pytest.mark.parametrize("payload_value", [None, True, False]) + def test_policy_false_overrides_any_payload(self, payload_value): + set_tool_policy(False) + assert _effective_enable_tools(_payload(payload_value)) is False diff --git a/studio/backend/tests/test_tool_policy_state.py b/studio/backend/tests/test_tool_policy_state.py new file mode 100644 index 0000000000..5f6b228281 --- /dev/null +++ b/studio/backend/tests/test_tool_policy_state.py @@ -0,0 +1,59 @@ +# SPDX-License-Identifier: AGPL-3.0-only +# Copyright 2026-present the Unsloth AI Inc. team. All rights reserved. + +""" +Tests for the process-level server-side tool policy used by `unsloth run`. + +The policy has three states: + None -> no CLI override (default; honor per-request enable_tools) + True -> CLI forced tools on + False -> CLI forced tools off +""" + +import os +import sys + +_backend = os.path.join(os.path.dirname(__file__), "..") +sys.path.insert(0, _backend) + +import pytest + +from state.tool_policy import ( + get_tool_policy, + reset_tool_policy, + set_tool_policy, +) + + +@pytest.fixture(autouse = True) +def _reset(): + reset_tool_policy() + yield + reset_tool_policy() + + +class TestToolPolicy: + def test_default_is_none(self): + assert get_tool_policy() is None + + def test_set_true_then_get(self): + set_tool_policy(True) + assert get_tool_policy() is True + + def test_set_false_then_get(self): + set_tool_policy(False) + assert get_tool_policy() is False + + def test_set_none_clears(self): + set_tool_policy(True) + set_tool_policy(None) + assert get_tool_policy() is None + + def test_reset_clears(self): + set_tool_policy(False) + reset_tool_policy() + assert get_tool_policy() is None + + def test_rejects_non_optional_bool(self): + with pytest.raises(TypeError): + set_tool_policy("true") # type: ignore[arg-type] diff --git a/studio/backend/tests/test_training_worker_flash_attn.py b/studio/backend/tests/test_training_worker_flash_attn.py index 986958408e..41a7c87df1 100644 --- a/studio/backend/tests/test_training_worker_flash_attn.py +++ b/studio/backend/tests/test_training_worker_flash_attn.py @@ -133,6 +133,22 @@ def test_causal_conv1d_fast_path_preserves_wheel_first_install_args(monkeypatch) ) +def test_causal_conv1d_fast_path_includes_qwen3_6_variants(monkeypatch): + install_mock = mock.Mock(return_value = True) + monkeypatch.setattr(worker, "_install_package_wheel_first", install_mock) + + worker._ensure_causal_conv1d_fast_path( + event_queue = [], + model_name = "unsloth/Qwen3.6-4B", + ) + worker._ensure_causal_conv1d_fast_path( + event_queue = [], + model_name = "unsloth/Qwen3_6-4B", + ) + + assert install_mock.call_count == 2 + + def test_mamba_ssm_path_preserves_wheel_first_install_args(monkeypatch): install_mock = mock.Mock(return_value = True) monkeypatch.setattr(worker, "_install_package_wheel_first", install_mock) diff --git a/studio/backend/tests/test_vram_estimation.py b/studio/backend/tests/test_vram_estimation.py index 0be067310d..e54ae6dcf8 100644 --- a/studio/backend/tests/test_vram_estimation.py +++ b/studio/backend/tests/test_vram_estimation.py @@ -2,7 +2,9 @@ # Copyright 2026-present the Unsloth AI Inc. team. All rights reserved. import unittest +from dataclasses import replace from types import SimpleNamespace +from unittest.mock import patch from utils.hardware.vram_estimation import ( ModelArchConfig, @@ -116,6 +118,55 @@ GPT_OSS = ModelArchConfig( num_dense_layers = 0, ) +STRUCTURED_MIXED = ModelArchConfig( + hidden_size = 256, + num_hidden_layers = 6, + num_attention_heads = 4, + num_key_value_heads = 2, + intermediate_size = 512, + vocab_size = 1024, + tie_word_embeddings = True, + head_dim = 80, + global_head_dim = 96, + num_global_key_value_heads = 1, + attention_k_eq_v = True, + layer_types = [ + "sliding_attention", + "full_attention", + "sliding_attention", + "full_attention", + "sliding_attention", + "full_attention", + ], +) + +STRUCTURED_SHARED = ModelArchConfig( + hidden_size = 192, + num_hidden_layers = 4, + num_attention_heads = 6, + num_key_value_heads = 2, + intermediate_size = 384, + vocab_size = 512, + tie_word_embeddings = True, + head_dim = 32, + num_kv_shared_layers = 2, + use_double_wide_mlp = True, + vocab_size_per_layer_input = 128, + hidden_size_per_layer_input = 48, + quant_4bit_factor = 3.6, +) + +QUANT_SKIP_STRUCTURED = replace( + STRUCTURED_SHARED, + quantization_skip_modules = [ + "model.layers.0.self_attn.q_proj", + "language_model.model.layers.1.mlp", + "layers.2", + "vision_tower", + "embed_tokens", + ], +) + class TestExtractArchConfig(unittest.TestCase): def test_basic_config(self): @@ -182,6 +233,42 @@ class TestExtractArchConfig(unittest.TestCase): arch = extract_arch_config(hf_config) self.assertEqual(arch.intermediate_size, 8192) + def test_structural_and_quantization_fields_are_config_derived(self): + hf_config = SimpleNamespace( + hidden_size = 256, + num_hidden_layers = 2, + num_attention_heads = 4, + num_key_value_heads = 2, + intermediate_size = 512, + vocab_size = 1024, + tie_word_embeddings = True, + head_dim = 80, + global_head_dim = 96, + num_global_key_value_heads = 1, + attention_k_eq_v = True, + layer_types = ["sliding_attention", "full_attention"], + num_kv_shared_layers = 1, + use_double_wide_mlp = True, + vocab_size_per_layer_input = 128, + hidden_size_per_layer_input = 48, + quantization_config = { + "bnb_4bit_use_double_quant": True, + "llm_int8_skip_modules": ["model.layers.0.self_attn"], + }, + ) + arch = extract_arch_config(hf_config) + self.assertEqual(arch.head_dim, 80) + self.assertEqual(arch.global_head_dim, 96) + self.assertEqual(arch.num_global_key_value_heads, 1) + self.assertTrue(arch.attention_k_eq_v) + self.assertEqual(arch.layer_types, ["sliding_attention", "full_attention"]) + self.assertEqual(arch.num_kv_shared_layers, 1) + self.assertTrue(arch.use_double_wide_mlp) + self.assertEqual(arch.vocab_size_per_layer_input, 128) + self.assertEqual(arch.hidden_size_per_layer_input, 48) + self.assertEqual(arch.quantization_skip_modules, ["model.layers.0.self_attn"]) + self.assertEqual(arch.quant_4bit_factor, 3.6) + class TestModelWeightsBytes(unittest.TestCase): def test_llama_8b_fp16(self): @@ -238,6 +325,18 @@ class TestLoraParams(unittest.TestCase): ratio = moe_lora / dense_lora self.assertAlmostEqual(ratio, 8.0, delta = 0.5) + def test_structured_moe_mlp_modules_scale_with_experts(self): + structured_moe = replace(QWEN3_MOE_30B, head_dim = 128) + dense_like = replace( + structured_moe, + num_experts = None, + moe_intermediate_size = None, + ) + target_modules = ["gate_proj", "up_proj", "down_proj"] + dense_lora = compute_lora_params(dense_like, 16, target_modules) + moe_lora = compute_lora_params(structured_moe, 16, target_modules) + self.assertGreater(moe_lora, dense_lora * 20) + def test_attention_modules_same_for_moe(self): dense_attn = compute_lora_params( LLAMA_8B, 16, ["q_proj", "k_proj", "v_proj", "o_proj"] @@ -247,6 +346,41 @@ class TestLoraParams(unittest.TestCase): ) self.assertEqual(dense_attn, moe_attn) + def test_all_linear_uses_default_text_modules(self): + text_only = compute_lora_params(STRUCTURED_MIXED, 16, DEFAULT_TARGET_MODULES) + all_linear = compute_lora_params(STRUCTURED_MIXED, 16, ["all-linear"]) + self.assertEqual(all_linear, text_only) + + def test_structural_layer_shapes_are_config_driven(self): + unstructured_arch = replace( + STRUCTURED_MIXED, + head_dim = None, + global_head_dim = None, + num_global_key_value_heads = None, + attention_k_eq_v = False, + layer_types = None, + ) + self.assertNotEqual( + compute_lora_params(unstructured_arch, 16, ["all-linear"]), + compute_lora_params(STRUCTURED_MIXED, 16, ["all-linear"]), + ) + self.assertNotEqual( + compute_model_weights_bytes(unstructured_arch, "qlora", True), + compute_model_weights_bytes(STRUCTURED_MIXED, "qlora", True), + ) + + def test_shared_kv_and_per_layer_inputs_change_weight_count(self): + unstructured_arch = replace( + STRUCTURED_SHARED, + head_dim = None, + num_kv_shared_layers = 0, + use_double_wide_mlp = False, + ) + self.assertNotEqual( + compute_model_weights_bytes(unstructured_arch, "qlora", True), + compute_model_weights_bytes(STRUCTURED_SHARED, "qlora", True), + ) + class TestOptimizerBytes(unittest.TestCase): def test_adamw_8bit(self): @@ -293,6 +427,163 @@ class TestActivationBytes(unittest.TestCase): act_4k = compute_activation_bytes(LLAMA_8B, 2, 4096, "unsloth") self.assertAlmostEqual(act_4k / act_2k, 2.0, delta = 0.1) + def test_flash_attention_uses_linear_path(self): + flash = compute_activation_bytes( + STRUCTURED_MIXED, + 1, + 4096, + "unsloth", + is_lora = True, + attention_implementation = "flash_attention_2", + ) + default = compute_activation_bytes( + STRUCTURED_MIXED, + 1, + 4096, + "unsloth", + is_lora = True, + ) + self.assertEqual(flash, default) + + def test_sdpa_attention_uses_linear_path(self): + flash = compute_activation_bytes( + STRUCTURED_MIXED, + 1, + 4096, + "unsloth", + is_lora = True, + attention_implementation = "flash_attention_2", + ) + sdpa = compute_activation_bytes( + STRUCTURED_MIXED, + 1, + 4096, + "unsloth", + is_lora = True, + attention_implementation = "sdpa", + ) + self.assertEqual(sdpa, flash) + + def test_non_flash_attention_uses_quadratic_path(self): + seq_len = 4096 + expected_quadratic = ( + 1 * STRUCTURED_MIXED.num_attention_heads * seq_len * seq_len * 2 * 12.0 + ) + for attention_implementation in ("eager", "unknown_impl", None): + with self.subTest(attention_implementation = attention_implementation): + non_flash = compute_activation_bytes( + STRUCTURED_MIXED, + 1, + seq_len, + "unsloth", + is_lora = True, + attention_implementation = attention_implementation, + ) + self.assertEqual(non_flash, int(expected_quadratic)) + + def test_non_flash_attention_without_gc_scales_quadratic_path_by_layers(self): + seq_len = 4096 + one_layer = ( + 1 * STRUCTURED_MIXED.num_attention_heads * seq_len * seq_len * 2 * 12.0 + ) + non_flash = compute_activation_bytes( + STRUCTURED_MIXED, + 1, + seq_len, + "none", + is_lora = True, + attention_implementation = "eager", + ) + self.assertEqual(non_flash, int(one_layer * STRUCTURED_MIXED.num_hidden_layers)) + self.assertGreater(non_flash, int(one_layer)) + + +class TestQuantizationSkips(unittest.TestCase): + def test_skipped_language_layers_stay_fp16(self): + no_skips = replace(QUANT_SKIP_STRUCTURED, quantization_skip_modules = []) + skipped = compute_model_weights_bytes(QUANT_SKIP_STRUCTURED, "qlora", True) + quantized = compute_model_weights_bytes(no_skips, "qlora", True) + self.assertGreater(skipped, quantized) + + def test_non_language_skips_do_not_double_count_text_weights(self): + arch = replace( + QUANT_SKIP_STRUCTURED, + quantization_skip_modules = ["vision_tower", "embed_tokens"], + ) + no_skips = replace(QUANT_SKIP_STRUCTURED, quantization_skip_modules = []) + self.assertEqual( + compute_model_weights_bytes(arch, "qlora", True), + compute_model_weights_bytes(no_skips, "qlora", True), + ) + + def test_double_quant_factor_reduces_quantized_weight_storage(self): + default_quant = replace(STRUCTURED_MIXED, quant_4bit_factor = 16 / 5) + double_quant = replace(STRUCTURED_MIXED, quant_4bit_factor = 3.6) + self.assertLess( + compute_model_weights_bytes(double_quant, "qlora", True), + compute_model_weights_bytes(default_quant, "qlora", True), + ) + + def test_prefixed_parent_and_child_skips_do_not_double_count(self): + parent_only = replace( + QUANT_SKIP_STRUCTURED, + quantization_skip_modules = ["language_model.model.layers.1.mlp"], + ) + parent_and_child = replace( + QUANT_SKIP_STRUCTURED, + quantization_skip_modules = [ + "language_model.model.layers.1.mlp", + "language_model.model.layers.1.mlp.gate_proj", + "model.layers.1.mlp.up_proj", + ], + ) + self.assertEqual( + compute_model_weights_bytes(parent_and_child, "qlora", True), + compute_model_weights_bytes(parent_only, "qlora", True), + ) + + def test_vlm_prefix_skip_module_does_not_match_text_alias(self): + # vision_tower-prefixed skips must not shadow text aliases sharing the + # same suffix. + baseline = replace(QUANT_SKIP_STRUCTURED, quantization_skip_modules = []) + vlm_skip = replace( + QUANT_SKIP_STRUCTURED, + quantization_skip_modules = [ + "vision_tower.model.layers.0.self_attn.q_proj", + "vision_tower.model.layers.1.mlp", + ], + ) + self.assertEqual( + compute_model_weights_bytes(vlm_skip, "qlora", True), + compute_model_weights_bytes(baseline, "qlora", True), + ) + + def test_mla_skip_module_uses_authoritative_attn_total(self): + from utils.hardware.vram_estimation import ( + _build_text_module_elements, + _compute_attn_elements, + ) + + mla = ModelArchConfig( + hidden_size = 2048, + num_hidden_layers = 4, + num_attention_heads = 16, + num_key_value_heads = 16, + intermediate_size = 8192, + vocab_size = 32000, + tie_word_embeddings = False, + q_lora_rank = 512, + kv_lora_rank = 128, + qk_nope_head_dim = 64, + qk_rope_head_dim = 32, + v_head_dim = 64, + ) + elements, _ = _build_text_module_elements(mla) + self.assertEqual( + elements["text.layers.0.self_attn"], + _compute_attn_elements(mla), + ) + class TestEstimateTrainingVram(unittest.TestCase): def test_llama_8b_qlora_reasonable_total(self): @@ -430,6 +721,90 @@ class TestEstimateTrainingVram(unittest.TestCase): v32.optimizer_states / v8.optimizer_states, 1.5, delta = 0.1 ) + def test_min_gpu_vram_treats_activations_as_per_gpu_fixed(self): + config = TrainingVramConfig(training_method = "qlora", load_in_4bit = True) + breakdown = estimate_training_vram(LLAMA_8B, config) + shardable = ( + breakdown.model_weights + + breakdown.lora_adapters + + breakdown.optimizer_states + + breakdown.gradients + ) + per_gpu_fixed = breakdown.activations + breakdown.cuda_overhead + for n_gpus in (1, 2, 4): + self.assertEqual( + breakdown.min_gpu_vram(n_gpus), + shardable // n_gpus + per_gpu_fixed, + ) + + def test_qlora_gradient_floor_is_capped_by_trainable_scale(self): + config = TrainingVramConfig( + training_method = "qlora", + batch_size = 1, + max_seq_length = 512, + lora_rank = 16, + target_modules = ["all-linear"], + gradient_checkpointing = "unsloth", + optimizer = "adamw_8bit", + load_in_4bit = True, + ) + breakdown = estimate_training_vram(LLAMA_8B, config) + lora_params = compute_lora_params(LLAMA_8B, 16, DEFAULT_TARGET_MODULES) + optimizer_bytes = compute_optimizer_bytes(lora_params, "adamw_8bit") + weight_floor = int(breakdown.model_weights * 0.15) + + self.assertEqual( + breakdown.gradients, + max(breakdown.activations_computed, optimizer_bytes), + ) + self.assertLess(breakdown.gradients, weight_floor) + self.assertEqual(breakdown.activations, breakdown.activations_computed) + + def test_full_finetuning_gradient_floor_remains_uncapped(self): + config = TrainingVramConfig( + training_method = "full", + batch_size = 1, + max_seq_length = 512, + gradient_checkpointing = "unsloth", + optimizer = "adamw_8bit", + load_in_4bit = False, + ) + expected_floor = int( + compute_model_weights_bytes(LLAMA_8B, "full", False) * 0.15 + ) + with patch( + "utils.hardware.vram_estimation.compute_gradient_bytes", + return_value = 1, + ): + breakdown = estimate_training_vram(LLAMA_8B, config) + self.assertEqual(breakdown.gradients, expected_floor) + + def test_non_flash_attention_flows_into_training_estimate(self): + config = TrainingVramConfig( + training_method = "qlora", + batch_size = 1, + max_seq_length = 4096, + lora_rank = 16, + target_modules = ["all-linear"], + gradient_checkpointing = "unsloth", + optimizer = "adamw_8bit", + load_in_4bit = True, + attention_implementation = "eager", + ) + breakdown = estimate_training_vram(STRUCTURED_MIXED, config) + self.assertEqual(breakdown.activations, breakdown.activations_computed) + self.assertGreater( + breakdown.activations, + compute_activation_bytes( + STRUCTURED_MIXED, + 1, + 4096, + "unsloth", + is_lora = True, + attention_implementation = "flash_attention_2", + ), + ) + class TestExtractArchConfigMoE(unittest.TestCase): def test_deepseek_v3_shared_experts(self): @@ -471,11 +846,16 @@ class TestExtractArchConfigMoE(unittest.TestCase): moe_intermediate_size = 768, decoder_sparse_step = 1, mlp_only_layers = [], + head_dim = 128, ) arch = extract_arch_config(hf_config) self.assertEqual(arch.num_experts, 128) self.assertEqual(arch.num_dense_layers, 0) + self.assertEqual(arch.head_dim, 128) self.assertIsNone(arch.q_lora_rank) + total_b = compute_total_params(arch) / 1e9 + self.assertGreater(total_b, 20) + self.assertLess(total_b, 50) def test_qwen3_moe_with_mlp_only_layers(self): hf_config = SimpleNamespace( @@ -542,6 +922,343 @@ class TestExtractArchConfigMoE(unittest.TestCase): self.assertEqual(arch.n_shared_experts, 0) self.assertEqual(arch.num_dense_layers, 0) self.assertIsNone(arch.q_lora_rank) + self.assertFalse(arch.moe_has_dense_mlp) + + def test_enable_moe_block_extracted_as_moe_has_dense_mlp(self): + hf_config = SimpleNamespace( + hidden_size = 2048, + num_hidden_layers = 8, + num_attention_heads = 16, + num_key_value_heads = 4, + intermediate_size = 4096, + vocab_size = 32000, + tie_word_embeddings = True, + num_experts = 8, + moe_intermediate_size = 1024, + head_dim = 128, + layer_types = ["full_attention"] * 8, + enable_moe_block = True, + ) + arch = extract_arch_config(hf_config) + self.assertTrue(arch.moe_has_dense_mlp) + + +class TestParallelDenseMoE(unittest.TestCase): + def _arch(self, **overrides): + base = ModelArchConfig( + hidden_size = 512, + num_hidden_layers = 4, + num_attention_heads = 8, + num_key_value_heads = 2, + intermediate_size = 1024, + vocab_size = 1024, + tie_word_embeddings = True, + num_experts = 8, + moe_intermediate_size = 512, + num_dense_layers = 0, + head_dim = 64, + layer_types = ["full_attention"] * 4, + ) + return replace(base, **overrides) + + def test_total_params_includes_parallel_dense_when_enable_moe_block(self): + without_parallel = self._arch(moe_has_dense_mlp = False) + with_parallel = self._arch(moe_has_dense_mlp = True) + self.assertGreater( + compute_total_params(with_parallel), + compute_total_params(without_parallel), + ) + + def test_lora_params_includes_parallel_dense_when_enable_moe_block(self): + without_parallel = self._arch(moe_has_dense_mlp = False) + with_parallel = self._arch(moe_has_dense_mlp = True) + target = ["gate_proj", "up_proj", "down_proj"] + self.assertGreater( + compute_lora_params(with_parallel, 16, target), + compute_lora_params(without_parallel, 16, target), + ) + + def test_activation_bytes_includes_parallel_dense_when_enable_moe_block(self): + without_parallel = self._arch(moe_has_dense_mlp = False) + with_parallel = self._arch(moe_has_dense_mlp = True) + self.assertGreater( + compute_activation_bytes( + with_parallel, + 1, + 2048, + "unsloth", + is_lora = True, + ), + compute_activation_bytes( + without_parallel, + 1, + 2048, + "unsloth", + is_lora = True, + ), + ) + + def test_layer_aggregates_split_dense_mlp_from_experts(self): + from utils.hardware.vram_estimation import _build_text_module_elements + + with_parallel = self._arch(moe_has_dense_mlp = True) + elements, _ = _build_text_module_elements(with_parallel) + moe_only = ( + with_parallel.hidden_size + * with_parallel.moe_intermediate_size + * 3 + * with_parallel.num_experts + + with_parallel.num_experts * with_parallel.hidden_size + ) + dense_only = with_parallel.hidden_size * with_parallel.intermediate_size * 3 + # why: under gemma4 enable_moe_block, the layer's `self.experts` is a + # sibling of `self.mlp`; the `text.layers..mlp` aggregate must + # cover the dense path only, with experts in their own aggregate. + self.assertEqual(elements["text.layers.0.mlp"], dense_only) + self.assertEqual(elements["text.layers.0.experts"], moe_only) + + +class TestDenseLayerIndices(unittest.TestCase): + def test_non_prefix_mlp_only_layers_preserve_position(self): + hf_config = SimpleNamespace( + hidden_size = 1024, + num_hidden_layers = 8, + num_attention_heads = 16, + num_key_value_heads = 4, + intermediate_size = 2048, + vocab_size = 32000, + tie_word_embeddings = True, + num_local_experts = 4, + moe_intermediate_size = 512, + decoder_sparse_step = 1, + mlp_only_layers = [3, 5], + ) + arch = extract_arch_config(hf_config) + self.assertEqual(arch.num_dense_layers, 2) + self.assertIn(3, arch.dense_layer_indices) + self.assertIn(5, arch.dense_layer_indices) + self.assertNotIn(0, arch.dense_layer_indices) + + def test_first_k_dense_replace_indices_are_prefix(self): + hf_config = SimpleNamespace( + hidden_size = 1024, + num_hidden_layers = 6, + num_attention_heads = 16, + num_key_value_heads = 4, + intermediate_size = 2048, + vocab_size = 32000, + tie_word_embeddings = False, + n_routed_experts = 8, + moe_intermediate_size = 512, + first_k_dense_replace = 2, + ) + arch = extract_arch_config(hf_config) + self.assertEqual(tuple(arch.dense_layer_indices), (0, 1)) + + +class TestKvSharedLayer(unittest.TestCase): + def test_fully_shared_kv_returns_false_matching_upstream(self): + from utils.hardware.vram_estimation import _is_kv_shared_layer + + arch = ModelArchConfig( + hidden_size = 512, + num_hidden_layers = 4, + num_attention_heads = 8, + num_key_value_heads = 2, + intermediate_size = 1024, + vocab_size = 1024, + num_kv_shared_layers = 4, + ) + for i in range(arch.num_hidden_layers): + self.assertFalse(_is_kv_shared_layer(arch, i)) + + def test_partial_share_returns_true_for_tail_layers(self): + from utils.hardware.vram_estimation import _is_kv_shared_layer + + arch = ModelArchConfig( + hidden_size = 512, + num_hidden_layers = 4, + num_attention_heads = 8, + num_key_value_heads = 2, + intermediate_size = 1024, + vocab_size = 1024, + num_kv_shared_layers = 2, + ) + self.assertFalse(_is_kv_shared_layer(arch, 0)) + self.assertFalse(_is_kv_shared_layer(arch, 1)) + self.assertTrue(_is_kv_shared_layer(arch, 2)) + self.assertTrue(_is_kv_shared_layer(arch, 3)) + + +class TestFlexAttentionLinear(unittest.TestCase): + def test_flex_attention_treated_as_linear(self): + flash = compute_activation_bytes( + STRUCTURED_MIXED, + 1, + 4096, + "unsloth", + is_lora = True, + attention_implementation = "flash_attention_2", + ) + flex = compute_activation_bytes( + STRUCTURED_MIXED, + 1, + 4096, + "unsloth", + is_lora = True, + attention_implementation = "flex_attention", + ) + self.assertEqual(flex, flash) + + +class TestNonStructuredParallelDense(unittest.TestCase): + def _arch(self, **overrides): + base = ModelArchConfig( + hidden_size = 1024, + num_hidden_layers = 4, + num_attention_heads = 16, + num_key_value_heads = 4, + intermediate_size = 4096, + vocab_size = 32000, + tie_word_embeddings = False, + num_experts = 8, + moe_intermediate_size = 768, + num_dense_layers = 0, + moe_has_dense_mlp = True, + ) + return replace(base, **overrides) + + def test_skip_module_uses_intermediate_size_for_parallel_dense(self): + from utils.hardware.vram_estimation import _build_text_module_elements + + arch = self._arch() + elements, _ = _build_text_module_elements(arch) + gate_proj = elements["text.layers.0.mlp.gate_proj"] + self.assertEqual(gate_proj, arch.hidden_size * arch.intermediate_size) + + +class TestPerLayerInputAccounting(unittest.TestCase): + def _arch(self, **overrides): + base = ModelArchConfig( + hidden_size = 1024, + num_hidden_layers = 4, + num_attention_heads = 16, + num_key_value_heads = 4, + intermediate_size = 2048, + vocab_size = 32000, + tie_word_embeddings = False, + head_dim = 64, + layer_types = ["full_attention"] * 4, + vocab_size_per_layer_input = 256, + hidden_size_per_layer_input = 96, + ) + return replace(base, **overrides) + + def test_per_layer_input_increases_total_params(self): + with_ple = self._arch() + without_ple = replace(with_ple, hidden_size_per_layer_input = 0) + self.assertGreater( + compute_total_params(with_ple), + compute_total_params(without_ple), + ) + + def test_per_layer_input_modules_count_quantizable_block(self): + with_ple = self._arch() + without_ple = replace(with_ple, hidden_size_per_layer_input = 0) + # The PLE block adds: model_projection (hd*nl*pli), per_layer_input_gate + # (hd*pli per layer) + per_layer_projection (pli*hd per layer) as + # quantizable text linears. + n_layers = with_ple.num_hidden_layers + hd = with_ple.hidden_size + pli = with_ple.hidden_size_per_layer_input + expected_quantizable_extra = ( + hd * (n_layers * pli) + (hd * pli) * n_layers + (pli * hd) * n_layers + ) + delta = compute_total_params(with_ple) - compute_total_params(without_ple) + self.assertGreaterEqual(delta, expected_quantizable_extra) + + def test_all_linear_lora_excludes_per_layer_input_modules(self): + # why: Unsloth's get_peft_regex requires module names to contain a + # component tag (mlp/attn/...); PLE module names (per_layer_input_gate, + # per_layer_projection, per_layer_model_projection) lack any tag, so + # all-linear training does NOT attach LoRA to them. + arch = self._arch() + without_ple = replace(arch, hidden_size_per_layer_input = 0) + self.assertEqual( + compute_lora_params(arch, 16, ["all-linear"]), + compute_lora_params(without_ple, 16, ["all-linear"]), + ) + + def test_explicit_target_modules_does_not_add_per_layer_input(self): + arch = self._arch() + without_ple = replace(arch, hidden_size_per_layer_input = 0) + self.assertEqual( + compute_lora_params(arch, 16, ["q_proj", "v_proj"]), + compute_lora_params(without_ple, 16, ["q_proj", "v_proj"]), + ) + + +class TestDenseMlpLayerFallback(unittest.TestCase): + def test_falls_back_to_count_when_indices_empty(self): + from utils.hardware.vram_estimation import _is_dense_mlp_layer + + arch = ModelArchConfig( + hidden_size = 512, + num_hidden_layers = 4, + num_attention_heads = 8, + num_key_value_heads = 2, + intermediate_size = 1024, + vocab_size = 1024, + num_experts = 4, + moe_intermediate_size = 256, + num_dense_layers = 2, + ) + self.assertTrue(_is_dense_mlp_layer(arch, 0)) + self.assertTrue(_is_dense_mlp_layer(arch, 1)) + self.assertFalse(_is_dense_mlp_layer(arch, 2)) + self.assertFalse(_is_dense_mlp_layer(arch, 3)) + + +class TestExpertsSkipGranularity(unittest.TestCase): + def _arch(self): + return ModelArchConfig( + hidden_size = 512, + num_hidden_layers = 4, + num_attention_heads = 8, + num_key_value_heads = 2, + intermediate_size = 1024, + vocab_size = 1024, + tie_word_embeddings = True, + num_experts = 8, + moe_intermediate_size = 512, + num_dense_layers = 0, + head_dim = 64, + layer_types = ["full_attention"] * 4, + moe_has_dense_mlp = True, + ) + + def test_experts_skip_excludes_parallel_dense_projections(self): + no_skip = self._arch() + skip_experts = replace( + no_skip, + quantization_skip_modules = ["model.layers.0.mlp.experts"], + ) + skip_full_mlp = replace( + no_skip, + quantization_skip_modules = ["model.layers.0.mlp"], + ) + bytes_no_skip = compute_model_weights_bytes(no_skip, "qlora", True) + bytes_skip_experts = compute_model_weights_bytes(skip_experts, "qlora", True) + bytes_skip_mlp = compute_model_weights_bytes(skip_full_mlp, "qlora", True) + # why: under gemma4 enable_moe_block, `self.experts` is a sibling of + # `self.mlp`; skipping `model.layers.0.mlp` should cover only the + # dense MLP, while `model.layers.0.mlp.experts` covers the routed + # experts. Routed experts have far more params than the dense MLP, + # so skipping experts must add more bytes than skipping the dense + # path. + self.assertGreater(bytes_skip_experts, bytes_no_skip) + self.assertGreater(bytes_skip_mlp, bytes_no_skip) + self.assertGreater(bytes_skip_experts, bytes_skip_mlp) class TestSharedExperts(unittest.TestCase): @@ -608,6 +1325,16 @@ class TestMLA(unittest.TestCase): lora_p = compute_lora_params(DEEPSEEK_V3, 16, ["q_proj", "v_proj", "o_proj"]) self.assertGreater(lora_p, 0) + def test_mla_with_head_dim_does_not_route_through_structured(self): + from utils.hardware.vram_estimation import _uses_structured_layer_shapes + + mla_with_head_dim = replace(DEEPSEEK_V3, head_dim = 128) + self.assertFalse(_uses_structured_layer_shapes(mla_with_head_dim)) + self.assertEqual( + compute_lora_params(DEEPSEEK_V3, 16, ["q_proj", "v_proj", "o_proj"]), + compute_lora_params(mla_with_head_dim, 16, ["q_proj", "v_proj", "o_proj"]), + ) + class TestDenseMoEMix(unittest.TestCase): def test_dense_layers_change_total(self): @@ -691,5 +1418,952 @@ class TestDenseMoEMix(unittest.TestCase): self.assertNotEqual(lora_all, lora_mix) +class TestMlpLayerTypesDispatch(unittest.TestCase): + def _hf(self, **fields): + text_config = SimpleNamespace( + hidden_size = 64, + num_hidden_layers = 4, + num_attention_heads = 4, + num_key_value_heads = 4, + intermediate_size = 128, + vocab_size = 1000, + tie_word_embeddings = True, + num_local_experts = 4, + moe_intermediate_size = 32, + **fields, + ) + return SimpleNamespace(text_config = text_config, quantization_config = {}) + + def test_mlp_layer_types_drives_dense_indices(self): + hf = self._hf(mlp_layer_types = ["sparse", "dense", "sparse", "dense"]) + arch = extract_arch_config(hf) + self.assertIsNotNone(arch) + self.assertEqual(arch.dense_layer_indices, (1, 3)) + self.assertEqual(arch.num_dense_layers, 2) + + def test_mlp_layer_types_takes_priority_over_first_k_dense_replace(self): + hf = self._hf( + mlp_layer_types = ["dense", "sparse", "dense", "sparse"], + first_k_dense_replace = 3, + ) + arch = extract_arch_config(hf) + self.assertEqual(arch.dense_layer_indices, (0, 2)) + + def test_mlp_layer_types_ignores_unknown_entries(self): + hf = self._hf(mlp_layer_types = ["dense", "moe", "dense", "linear"]) + arch = extract_arch_config(hf) + self.assertEqual(arch.dense_layer_indices, (0, 2)) + + def test_mlp_layer_types_shorter_than_layers_only_marks_present(self): + hf = self._hf(mlp_layer_types = ["dense", "sparse"]) + arch = extract_arch_config(hf) + self.assertEqual(arch.dense_layer_indices, (0,)) + + def test_empty_mlp_layer_types_falls_through_to_first_k(self): + hf = self._hf(mlp_layer_types = [], first_k_dense_replace = 2) + arch = extract_arch_config(hf) + self.assertEqual(arch.dense_layer_indices, (0, 1)) + + +class TestPerLayerInputSkipAlias(unittest.TestCase): + def _hf(self, skip): + text_config = SimpleNamespace( + hidden_size = 64, + num_hidden_layers = 2, + num_attention_heads = 4, + num_key_value_heads = 4, + intermediate_size = 128, + vocab_size = 1000, + tie_word_embeddings = True, + hidden_size_per_layer_input = 8, + vocab_size_per_layer_input = 256, + ) + return SimpleNamespace( + text_config = text_config, + quantization_config = {"llm_int8_skip_modules": list(skip)}, + ) + + def test_per_layer_input_gate_skip_pulls_nonzero_delta(self): + from utils.hardware.vram_estimation import _compute_skipped_quantizable_elements + + arch = extract_arch_config(self._hf(["model.layers.0.per_layer_input_gate"])) + delta = _compute_skipped_quantizable_elements(arch) + self.assertEqual(delta, arch.hidden_size * arch.hidden_size_per_layer_input) + + def test_per_layer_model_projection_skip_pulls_global_delta(self): + from utils.hardware.vram_estimation import _compute_skipped_quantizable_elements + + arch = extract_arch_config(self._hf(["model.per_layer_model_projection"])) + delta = _compute_skipped_quantizable_elements(arch) + self.assertEqual( + delta, + arch.hidden_size + * arch.num_hidden_layers + * arch.hidden_size_per_layer_input, + ) + + def test_layer_aggregate_skip_includes_per_layer_input_modules(self): + from utils.hardware.vram_estimation import ( + _compute_skipped_quantizable_elements, + ) + + arch_with = extract_arch_config(self._hf(["model.layers.0"])) + # The text.layers.0 aggregate must include the PLE per-layer modules, + # so the same skip on a config without PLE produces a smaller value. + arch_without = extract_arch_config( + SimpleNamespace( + text_config = SimpleNamespace( + hidden_size = 64, + num_hidden_layers = 2, + num_attention_heads = 4, + num_key_value_heads = 4, + intermediate_size = 128, + vocab_size = 1000, + tie_word_embeddings = True, + hidden_size_per_layer_input = 0, + vocab_size_per_layer_input = 0, + ), + quantization_config = {"llm_int8_skip_modules": ["model.layers.0"]}, + ) + ) + self.assertGreater( + _compute_skipped_quantizable_elements(arch_with), + _compute_skipped_quantizable_elements(arch_without), + ) + + +class TestAllLinearStringHandling(unittest.TestCase): + def test_compute_lora_params_accepts_bare_all_linear_string(self): + list_form = compute_lora_params(LLAMA_8B, 16, ["all-linear"]) + str_form = compute_lora_params(LLAMA_8B, 16, "all-linear") + self.assertEqual(list_form, str_form) + self.assertGreater(list_form, 0) + + def test_compute_lora_params_string_with_underscores_normalized(self): + list_form = compute_lora_params(LLAMA_8B, 16, ["all_linear"]) + str_form = compute_lora_params(LLAMA_8B, 16, "all_linear") + self.assertEqual(list_form, str_form) + self.assertGreater(str_form, 0) + + +class TestSharedExpertVariants(unittest.TestCase): + def _hf(self, **fields): + text_config = SimpleNamespace( + hidden_size = 256, + num_hidden_layers = 4, + num_attention_heads = 8, + num_key_value_heads = 4, + intermediate_size = 1024, + vocab_size = 1000, + tie_word_embeddings = False, + num_local_experts = 8, + moe_intermediate_size = 128, + **fields, + ) + return SimpleNamespace(text_config = text_config, quantization_config = {}) + + def test_shared_expert_intermediate_size_extracted_and_infers_count(self): + arch = extract_arch_config(self._hf(shared_expert_intermediate_size = 64)) + self.assertEqual(arch.shared_expert_intermediate_size, 64) + self.assertEqual(arch.n_shared_experts, 1) + + def test_num_shared_experts_alias_extracted(self): + arch = extract_arch_config(self._hf(num_shared_experts = 2)) + self.assertEqual(arch.n_shared_experts, 2) + + def test_n_shared_experts_takes_priority_over_alias(self): + arch = extract_arch_config(self._hf(n_shared_experts = 3, num_shared_experts = 99)) + self.assertEqual(arch.n_shared_experts, 3) + + def test_shared_expert_size_separate_from_routed_changes_weight_count(self): + from utils.hardware.vram_estimation import _compute_moe_mlp_elements + + arch_separate = extract_arch_config( + self._hf(shared_expert_intermediate_size = 64) + ) + arch_implicit = extract_arch_config(self._hf(n_shared_experts = 1)) + # Different shared sizes (64 vs default moe_intermediate_size=128) must + # produce different MoE element counts. + self.assertNotEqual( + _compute_moe_mlp_elements(arch_separate), + _compute_moe_mlp_elements(arch_implicit), + ) + + def test_shared_expert_gate_counted_only_for_qwen_style(self): + from utils.hardware.vram_estimation import _compute_moe_mlp_elements + + # Qwen-style: shared_expert_intermediate_size set -> shared_expert_gate counted. + qwen_arch = extract_arch_config(self._hf(shared_expert_intermediate_size = 64)) + hd = qwen_arch.hidden_size + ms = qwen_arch.moe_intermediate_size + ne = qwen_arch.num_experts + ss = qwen_arch.shared_expert_intermediate_size + expected = hd * ms * 3 * ne + ne * hd + hd * ss * 3 * 1 + 1 * hd + self.assertEqual(_compute_moe_mlp_elements(qwen_arch), expected) + + # Non-Qwen shared experts (e.g. Exaone-MoE) -> no shared_expert_gate. + plain_arch = extract_arch_config(self._hf(n_shared_experts = 1)) + hd = plain_arch.hidden_size + ms = plain_arch.moe_intermediate_size + ne = plain_arch.num_experts + expected_plain = hd * ms * 3 * ne + ne * hd + hd * ms * 3 * 1 + self.assertEqual(_compute_moe_mlp_elements(plain_arch), expected_plain) + + +class TestSharedExpertActivation(unittest.TestCase): + def _make(self, **fields): + text_config = SimpleNamespace( + hidden_size = 512, + num_hidden_layers = 4, + num_attention_heads = 8, + num_key_value_heads = 4, + intermediate_size = 1024, + vocab_size = 1000, + tie_word_embeddings = False, + num_local_experts = 4, + moe_intermediate_size = 64, + **fields, + ) + return extract_arch_config( + SimpleNamespace(text_config = text_config, quantization_config = {}) + ) + + def test_shared_expert_increases_activation_bytes(self): + with_shared = self._make(shared_expert_intermediate_size = 64) + without = self._make() + self.assertGreater( + compute_activation_bytes( + with_shared, + 2, + 1024, + "none", + is_lora = True, + attention_implementation = "flash_attention_2", + ), + compute_activation_bytes( + without, + 2, + 1024, + "none", + is_lora = True, + attention_implementation = "flash_attention_2", + ), + ) + + def test_shared_expert_plus_dense_block_compose(self): + # gemma4 enable_moe_block with hypothetical shared expert: dense + routed + # + shared all live per layer; mlp_size should sum all three terms. + from utils.hardware.vram_estimation import _layer_qkv_mlp_sizes + + arch = self._make( + enable_moe_block = True, + shared_expert_intermediate_size = 32, + head_dim = 64, + layer_types = ["full_attention"] * 4, + ) + _, mlp_size = _layer_qkv_mlp_sizes(arch, 0) + # routed (64) + shared (32) + parallel dense intermediate (1024) + self.assertEqual(mlp_size, 64 + 32 + 1024) + + +class TestPerLayerInputActivation(unittest.TestCase): + def _make(self, **fields): + text_config = SimpleNamespace( + hidden_size = 512, + num_hidden_layers = 4, + num_attention_heads = 8, + num_key_value_heads = 4, + intermediate_size = 1024, + vocab_size = 1000, + tie_word_embeddings = False, + **fields, + ) + return extract_arch_config( + SimpleNamespace(text_config = text_config, quantization_config = {}) + ) + + def test_ple_increases_activation_bytes(self): + with_ple = self._make( + hidden_size_per_layer_input = 64, + vocab_size_per_layer_input = 256, + ) + without = self._make() + self.assertGreater( + compute_activation_bytes( + with_ple, + 2, + 1024, + "none", + is_lora = True, + attention_implementation = "flash_attention_2", + ), + compute_activation_bytes( + without, + 2, + 1024, + "none", + is_lora = True, + attention_implementation = "flash_attention_2", + ), + ) + + def test_ple_zero_does_not_inflate_activations(self): + without = self._make(hidden_size_per_layer_input = 0) + baseline = self._make() + self.assertEqual( + compute_activation_bytes( + without, + 2, + 512, + "none", + is_lora = True, + attention_implementation = "flash_attention_2", + ), + compute_activation_bytes( + baseline, + 2, + 512, + "none", + is_lora = True, + attention_implementation = "flash_attention_2", + ), + ) + + +class TestKvSharedActivation(unittest.TestCase): + def _make(self, kv_shared): + text_config = SimpleNamespace( + hidden_size = 512, + num_hidden_layers = 4, + num_attention_heads = 8, + num_key_value_heads = 4, + intermediate_size = 1024, + vocab_size = 1000, + tie_word_embeddings = False, + head_dim = 64, + num_kv_shared_layers = kv_shared, + layer_types = ["full_attention"] * 4, + ) + return extract_arch_config( + SimpleNamespace(text_config = text_config, quantization_config = {}) + ) + + def test_kv_shared_layers_keep_activation_bytes(self): + shared = self._make(kv_shared = 2) + full = self._make(kv_shared = 0) + self.assertEqual( + compute_activation_bytes( + shared, + 2, + 1024, + "none", + is_lora = True, + attention_implementation = "flash_attention_2", + ), + compute_activation_bytes( + full, + 2, + 1024, + "none", + is_lora = True, + attention_implementation = "flash_attention_2", + ), + ) + + +class TestSparseMoeSkipAliases(unittest.TestCase): + def _hf(self, skip, **fields): + text_config = SimpleNamespace( + hidden_size = 128, + num_hidden_layers = 2, + num_attention_heads = 4, + num_key_value_heads = 4, + intermediate_size = 256, + vocab_size = 1000, + tie_word_embeddings = False, + num_local_experts = 4, + moe_intermediate_size = 64, + **fields, + ) + return SimpleNamespace( + text_config = text_config, + quantization_config = {"llm_int8_skip_modules": list(skip)}, + ) + + def test_gemma4_layers_experts_alias_pulls_routed(self): + from utils.hardware.vram_estimation import _compute_skipped_quantizable_elements + + arch = extract_arch_config( + self._hf(["model.layers.0.experts"], enable_moe_block = True) + ) + self.assertGreater(_compute_skipped_quantizable_elements(arch), 0) + + def test_qwen_shared_expert_skip_pulls_only_shared(self): + from utils.hardware.vram_estimation import _compute_skipped_quantizable_elements + + arch = extract_arch_config( + self._hf( + ["model.layers.0.mlp.shared_expert"], + shared_expert_intermediate_size = 32, + ) + ) + # shared_expert delta only -- routed mlp.experts is NOT skipped. + delta = _compute_skipped_quantizable_elements(arch) + self.assertGreater(delta, 0) + full_layer = extract_arch_config( + self._hf( + ["model.layers.0.mlp"], + shared_expert_intermediate_size = 32, + ) + ) + self.assertGreater( + _compute_skipped_quantizable_elements(full_layer), + delta, + ) + + def test_exaone_shared_experts_plural_alias(self): + from utils.hardware.vram_estimation import _compute_skipped_quantizable_elements + + arch = extract_arch_config( + self._hf( + ["model.layers.0.mlp.shared_experts"], + num_shared_experts = 1, + ) + ) + self.assertGreater(_compute_skipped_quantizable_elements(arch), 0) + + +class TestAllLinearMoELoraExclusion(unittest.TestCase): + def _arch(self, **fields): + text_config = SimpleNamespace( + hidden_size = 256, + num_hidden_layers = 2, + num_attention_heads = 4, + num_key_value_heads = 4, + intermediate_size = 512, + vocab_size = 1000, + tie_word_embeddings = False, + num_local_experts = 8, + moe_intermediate_size = 64, + **fields, + ) + return extract_arch_config( + SimpleNamespace(text_config = text_config, quantization_config = {}) + ) + + def test_all_linear_drops_routed_moe_expert_lora(self): + arch = self._arch() + all_linear = compute_lora_params(arch, 8, "all-linear") + explicit = compute_lora_params(arch, 8, ["gate_proj", "up_proj", "down_proj"]) + self.assertLess(all_linear, explicit) + + def test_all_linear_drops_shared_expert_lora(self): + arch = self._arch(shared_expert_intermediate_size = 32) + all_linear = compute_lora_params(arch, 8, "all-linear") + explicit = compute_lora_params(arch, 8, ["gate_proj", "up_proj", "down_proj"]) + # explicit includes routed + shared MoE; all-linear includes neither. + self.assertLess(all_linear, explicit) + + def test_all_linear_includes_attention_lora(self): + arch = self._arch() + all_linear = compute_lora_params(arch, 8, "all-linear") + attn_only = compute_lora_params( + arch, 8, ["q_proj", "k_proj", "v_proj", "o_proj"] + ) + # all-linear still attaches to attention nn.Linear modules. + self.assertGreaterEqual(all_linear, attn_only) + + +class TestExplicitPerLayerInputLora(unittest.TestCase): + def _arch(self): + text_config = SimpleNamespace( + hidden_size = 256, + num_hidden_layers = 3, + num_attention_heads = 4, + num_key_value_heads = 4, + intermediate_size = 512, + vocab_size = 1000, + tie_word_embeddings = False, + hidden_size_per_layer_input = 32, + vocab_size_per_layer_input = 128, + ) + return extract_arch_config( + SimpleNamespace(text_config = text_config, quantization_config = {}) + ) + + def test_explicit_per_layer_input_gate_returns_nonzero(self): + arch = self._arch() + result = compute_lora_params(arch, 16, ["per_layer_input_gate"]) + self.assertGreater(result, 0) + + def test_explicit_per_layer_projection_returns_nonzero(self): + arch = self._arch() + result = compute_lora_params(arch, 16, ["per_layer_projection"]) + self.assertGreater(result, 0) + + def test_explicit_per_layer_model_projection_returns_nonzero(self): + arch = self._arch() + result = compute_lora_params(arch, 16, ["per_layer_model_projection"]) + self.assertGreater(result, 0) + + def test_explicit_ple_string_target_handled(self): + # Bare-string target with a PLE name should not be iterated char-by-char. + arch = self._arch() + list_form = compute_lora_params(arch, 16, ["per_layer_input_gate"]) + str_form = compute_lora_params(arch, 16, "per_layer_input_gate") + self.assertEqual(list_form, str_form) + + +class TestTopKExpertActivation(unittest.TestCase): + def _make(self, **fields): + text_config = SimpleNamespace( + hidden_size = 512, + num_hidden_layers = 4, + num_attention_heads = 8, + num_key_value_heads = 4, + intermediate_size = 1024, + vocab_size = 1000, + tie_word_embeddings = False, + num_local_experts = 8, + moe_intermediate_size = 64, + **fields, + ) + return extract_arch_config( + SimpleNamespace(text_config = text_config, quantization_config = {}) + ) + + def test_num_experts_per_tok_extracted(self): + arch = self._make(num_experts_per_tok = 4) + self.assertEqual(arch.num_experts_per_tok, 4) + + def test_top_k_experts_alias_extracted(self): + arch = self._make(top_k_experts = 8) + self.assertEqual(arch.num_experts_per_tok, 8) + + def test_default_top_k_one_unchanged(self): + arch = self._make() + self.assertEqual(arch.num_experts_per_tok, 1) + + def test_top_k_scales_moe_activation(self): + single = self._make() + multi = self._make(num_experts_per_tok = 8) + single_act = compute_activation_bytes( + single, + 2, + 512, + "none", + is_lora = True, + attention_implementation = "flash_attention_2", + ) + multi_act = compute_activation_bytes( + multi, + 2, + 512, + "none", + is_lora = True, + attention_implementation = "flash_attention_2", + ) + self.assertGreater(multi_act, single_act) + + +class TestErnieMoEListConfig(unittest.TestCase): + def _hf(self, **fields): + text_config = SimpleNamespace( + hidden_size = 256, + num_hidden_layers = 4, + num_attention_heads = 4, + num_key_value_heads = 4, + intermediate_size = 1024, + vocab_size = 1000, + tie_word_embeddings = False, + **fields, + ) + return SimpleNamespace(text_config = text_config, quantization_config = {}) + + def test_list_moe_intermediate_size_scalarized(self): + arch = extract_arch_config( + self._hf( + moe_num_experts = 32, + moe_intermediate_size = [1536, 512], + ) + ) + # why: ERNIE 4.5 VL MoE encodes [text_routed, vision_routed]; the + # second element is the vision-routed expert width, not the shared + # expert width. Shared experts are sized from the text-routed width + # (= moe_intermediate_size[0]) when moe_num_shared_experts is set. + self.assertEqual(arch.moe_intermediate_size, 1536) + self.assertIsNone(arch.shared_expert_intermediate_size) + self.assertEqual(arch.n_shared_experts, 0) + + def test_moe_num_experts_alias_extracted(self): + arch = extract_arch_config( + self._hf( + moe_num_experts = 64, + moe_intermediate_size = 1024, + ) + ) + self.assertEqual(arch.num_experts, 64) + + def test_moe_num_shared_experts_alias_extracted(self): + arch = extract_arch_config( + self._hf( + moe_num_experts = 16, + moe_num_shared_experts = 2, + moe_intermediate_size = 1024, + ) + ) + self.assertEqual(arch.n_shared_experts, 2) + + def test_explicit_shared_size_overrides_list_second_element(self): + arch = extract_arch_config( + self._hf( + moe_num_experts = 8, + moe_intermediate_size = [1536, 512], + shared_expert_intermediate_size = 256, + ) + ) + # Explicit shared size wins over moe_intermediate_size[1]. + self.assertEqual(arch.shared_expert_intermediate_size, 256) + + +class TestSuffixSkipModuleMatch(unittest.TestCase): + def _hf(self, skip): + text_config = SimpleNamespace( + hidden_size = 128, + num_hidden_layers = 2, + num_attention_heads = 4, + num_key_value_heads = 4, + intermediate_size = 256, + vocab_size = 1000, + tie_word_embeddings = False, + ) + return SimpleNamespace( + text_config = text_config, + quantization_config = {"llm_int8_skip_modules": list(skip)}, + ) + + def test_q_proj_suffix_skip_matches_all_layers(self): + from utils.hardware.vram_estimation import _compute_skipped_quantizable_elements + + arch = extract_arch_config(self._hf(["q_proj"])) + delta = _compute_skipped_quantizable_elements(arch) + # 2 layers * hd * hd of q_proj weight elements. + self.assertEqual(delta, 2 * arch.hidden_size * arch.hidden_size) + + def test_self_attn_aggregate_skip_matches_aggregate(self): + from utils.hardware.vram_estimation import _compute_skipped_quantizable_elements + + arch = extract_arch_config(self._hf(["self_attn"])) + # The aggregate text.layers..self_attn matches; total covers both layers. + delta = _compute_skipped_quantizable_elements(arch) + self.assertGreater(delta, 0) + + def test_vision_prefix_skip_does_not_match_text_alias(self): + from utils.hardware.vram_estimation import _module_path_matches + + # vision_tower-prefixed full path must NOT match text-tower aliases. + self.assertFalse( + _module_path_matches( + "vision_tower.model.layers.0.self_attn.q_proj", + "model.layers.0.self_attn.q_proj", + ) + ) + + +class TestMultimodalFullModelBytes(unittest.TestCase): + def test_extra_bytes_added_when_safetensors_exceeds_text_arch(self): + from utils.hardware import hardware as hardware_module + + config = SimpleNamespace( + hidden_size = 1024, + num_hidden_layers = 4, + num_attention_heads = 8, + num_key_value_heads = 4, + intermediate_size = 2048, + vocab_size = 32000, + tie_word_embeddings = False, + ) + # Force safetensors size >>> arch text-only bytes. + big_safetensors = 20 * 1024**3 + with ( + patch.object( + hardware_module, + "_load_config_for_gpu_estimate", + return_value = config, + ), + patch.object( + hardware_module, + "estimate_fp16_model_size_bytes", + return_value = (big_safetensors, "safetensors"), + ), + patch.object( + hardware_module, + "_determine_attention_impl_for_gpu_estimate", + return_value = "flash_attention_2", + ), + patch.object( + hardware_module, + "get_visible_gpu_count", + return_value = 1, + ), + ): + _, metadata = hardware_module.estimate_required_model_memory_gb( + "fake/model", + training_type = "LoRA/QLoRA", + load_in_4bit = True, + ) + self.assertEqual(metadata.get("estimation_mode"), "detailed") + # model_weights_gb must reflect the extra non-text bytes (>5 GB + # since text-only arch_fp16 is small for these dims). + self.assertGreater(metadata["vram_breakdown"]["model_weights_gb"], 5.0) + + def test_no_extra_when_safetensors_smaller_than_text_arch(self): + from utils.hardware import hardware as hardware_module + + config = SimpleNamespace( + hidden_size = 4096, + num_hidden_layers = 32, + num_attention_heads = 32, + num_key_value_heads = 8, + intermediate_size = 11008, + vocab_size = 32000, + tie_word_embeddings = False, + ) + tiny_safetensors = 100 # bytes, deliberately absurdly small + with ( + patch.object( + hardware_module, + "_load_config_for_gpu_estimate", + return_value = config, + ), + patch.object( + hardware_module, + "estimate_fp16_model_size_bytes", + return_value = (tiny_safetensors, "safetensors"), + ), + patch.object( + hardware_module, + "_determine_attention_impl_for_gpu_estimate", + return_value = "flash_attention_2", + ), + patch.object( + hardware_module, + "get_visible_gpu_count", + return_value = 1, + ), + ): + required, metadata = hardware_module.estimate_required_model_memory_gb( + "fake/model", + training_type = "LoRA/QLoRA", + load_in_4bit = True, + ) + # No negative extra; required_gb stays a positive finite number. + self.assertGreater(required, 0) + + +class TestLlama4ArchExtraction(unittest.TestCase): + def _llama4_text_config(self, **fields): + base = dict( + hidden_size = 2048, + num_hidden_layers = 4, + num_attention_heads = 16, + num_key_value_heads = 4, + intermediate_size = 8192, + intermediate_size_mlp = 16384, + vocab_size = 32000, + tie_word_embeddings = True, + num_local_experts = 4, + num_experts_per_tok = 2, + ) + base.update(fields) + return SimpleNamespace(**base) + + def test_llama4_moe_layers_dispatch_uses_explicit_indices(self): + from utils.hardware.vram_estimation import _compute_dense_layer_indices + + cfg = SimpleNamespace(num_hidden_layers = 4, moe_layers = [1, 3]) + self.assertEqual(_compute_dense_layer_indices(cfg, 4), (0, 2)) + + def test_llama4_moe_layers_takes_priority_over_first_k_dense_replace(self): + from utils.hardware.vram_estimation import _compute_dense_layer_indices + + cfg = SimpleNamespace( + num_hidden_layers = 6, + moe_layers = [2, 4], + first_k_dense_replace = 4, + ) + self.assertEqual(_compute_dense_layer_indices(cfg, 6), (0, 1, 3, 5)) + + def test_dense_intermediate_size_picks_up_intermediate_size_mlp(self): + from utils.hardware.vram_estimation import _dense_mlp_size + + arch = extract_arch_config(self._llama4_text_config(moe_layers = [1, 3])) + self.assertIsNotNone(arch) + self.assertEqual(arch.intermediate_size, 8192) + self.assertEqual(arch.dense_intermediate_size, 16384) + self.assertEqual(_dense_mlp_size(arch), 16384) + + def test_auto_attaches_one_shared_expert_at_routed_width(self): + from utils.hardware.vram_estimation import _shared_expert_size + + arch = extract_arch_config(self._llama4_text_config(moe_layers = [1, 3])) + self.assertIsNotNone(arch) + self.assertEqual(arch.n_shared_experts, 1) + self.assertIsNone(arch.shared_expert_intermediate_size) + self.assertEqual(_shared_expert_size(arch), arch.intermediate_size) + + def test_non_llama4_config_leaves_dense_intermediate_size_none(self): + from utils.hardware.vram_estimation import _dense_mlp_size + + cfg = SimpleNamespace( + hidden_size = 1024, + num_hidden_layers = 4, + num_attention_heads = 8, + num_key_value_heads = 2, + intermediate_size = 4096, + vocab_size = 32000, + tie_word_embeddings = True, + ) + arch = extract_arch_config(cfg) + self.assertIsNotNone(arch) + self.assertIsNone(arch.dense_intermediate_size) + self.assertEqual(_dense_mlp_size(arch), 4096) + + def test_intermediate_size_mlp_without_moe_does_not_force_shared_expert(self): + cfg = SimpleNamespace( + hidden_size = 2048, + num_hidden_layers = 4, + num_attention_heads = 16, + num_key_value_heads = 4, + intermediate_size = 8192, + intermediate_size_mlp = 16384, + vocab_size = 32000, + tie_word_embeddings = True, + ) + arch = extract_arch_config(cfg) + self.assertIsNotNone(arch) + self.assertEqual(arch.dense_intermediate_size, 16384) + self.assertEqual(arch.n_shared_experts, 0) + + +class TestDbrxFfnConfigExtraction(unittest.TestCase): + def test_extracts_moe_fields_from_ffn_subconfig(self): + ffn = SimpleNamespace(moe_num_experts = 4, moe_top_k = 2, ffn_hidden_size = 1024) + cfg = SimpleNamespace( + hidden_size = 2048, + num_hidden_layers = 4, + num_attention_heads = 16, + num_key_value_heads = 4, + intermediate_size = 2048, + vocab_size = 32000, + tie_word_embeddings = False, + ffn_config = ffn, + ) + arch = extract_arch_config(cfg) + self.assertIsNotNone(arch) + self.assertEqual(arch.num_experts, 4) + self.assertEqual(arch.num_experts_per_tok, 2) + self.assertEqual(arch.moe_intermediate_size, 1024) + + def test_top_level_attrs_take_precedence_over_ffn_config(self): + ffn = SimpleNamespace(moe_num_experts = 4, moe_top_k = 2, ffn_hidden_size = 1024) + cfg = SimpleNamespace( + hidden_size = 2048, + num_hidden_layers = 4, + num_attention_heads = 16, + num_key_value_heads = 4, + intermediate_size = 2048, + vocab_size = 32000, + tie_word_embeddings = False, + ffn_config = ffn, + num_local_experts = 16, + num_experts_per_tok = 8, + ) + arch = extract_arch_config(cfg) + self.assertIsNotNone(arch) + self.assertEqual(arch.num_experts, 16) + self.assertEqual(arch.num_experts_per_tok, 8) + + +class TestErniePhaseModuloDispatch(unittest.TestCase): + def test_phase_modulo_with_interval_two_matches_decoder(self): + from utils.hardware.vram_estimation import _compute_dense_layer_indices + + cfg = SimpleNamespace( + num_hidden_layers = 10, + moe_layer_start_index = 2, + moe_layer_end_index = 8, + moe_layer_interval = 2, + ) + # Decoder gates by ((i + 1) % 2 == 0) AND 2 <= i <= 8 -> MoE = {3, 5, 7}. + self.assertEqual(_compute_dense_layer_indices(cfg, 10), (0, 1, 2, 4, 6, 8, 9)) + + def test_phase_modulo_with_interval_three(self): + from utils.hardware.vram_estimation import _compute_dense_layer_indices + + cfg = SimpleNamespace( + num_hidden_layers = 9, + moe_layer_start_index = 0, + moe_layer_end_index = -1, + moe_layer_interval = 3, + ) + self.assertEqual(_compute_dense_layer_indices(cfg, 9), (0, 1, 3, 4, 6, 7)) + + +class TestErnieVlSharedExpertWidth(unittest.TestCase): + def test_shared_expert_width_uses_text_routed_not_vision(self): + from utils.hardware.vram_estimation import ( + _compute_shared_moe_elements, + _shared_expert_size, + ) + + cfg = SimpleNamespace( + text_config = SimpleNamespace( + hidden_size = 1024, + num_hidden_layers = 4, + num_attention_heads = 8, + num_key_value_heads = 4, + intermediate_size = 2048, + vocab_size = 32000, + tie_word_embeddings = False, + moe_num_experts = 8, + moe_num_shared_experts = 2, + moe_intermediate_size = [1536, 512], + ), + quantization_config = {}, + ) + arch = extract_arch_config(cfg) + self.assertIsNotNone(arch) + self.assertIsNone(arch.shared_expert_intermediate_size) + self.assertEqual(arch.moe_intermediate_size, 1536) + self.assertEqual(arch.n_shared_experts, 2) + self.assertEqual(_shared_expert_size(arch), 1536) + self.assertEqual(_compute_shared_moe_elements(arch), 1024 * 1536 * 3 * 2) + + def test_qwen_style_explicit_shared_expert_size_still_adds_gate(self): + from utils.hardware.vram_estimation import _compute_shared_moe_elements + + cfg = SimpleNamespace( + hidden_size = 1024, + num_hidden_layers = 4, + num_attention_heads = 8, + num_key_value_heads = 4, + intermediate_size = 2048, + vocab_size = 32000, + tie_word_embeddings = False, + num_local_experts = 8, + moe_intermediate_size = 256, + shared_expert_intermediate_size = 768, + ) + arch = extract_arch_config(cfg) + self.assertIsNotNone(arch) + self.assertEqual(arch.shared_expert_intermediate_size, 768) + self.assertEqual(arch.n_shared_experts, 1) + self.assertEqual( + _compute_shared_moe_elements(arch), + 1024 * 768 * 3 + 1 * 1024, + ) + + if __name__ == "__main__": unittest.main() diff --git a/studio/backend/utils/datasets/model_mappings.py b/studio/backend/utils/datasets/model_mappings.py index 36f6886ef6..21e8566ac5 100644 --- a/studio/backend/utils/datasets/model_mappings.py +++ b/studio/backend/utils/datasets/model_mappings.py @@ -364,6 +364,10 @@ TEMPLATE_TO_MODEL_MAPPER = { "unsloth/Qwen3-4B-Thinking-2507-bnb-4bit", "unsloth/Qwen3-30B-A3B-Thinking-2507", "Qwen/Qwen3-30B-A3B-Thinking-2507", + "Qwen/Qwen3.6-35B-A3B", + "unsloth/Qwen3.6-35B-A3B", + "Qwen/Qwen3.6-27B", + "unsloth/Qwen3.6-27B", ), "qwen3.5": ( "unsloth/Qwen3.5-0.8B", diff --git a/studio/backend/utils/hardware/VRAM_ESTIMATION.md b/studio/backend/utils/hardware/VRAM_ESTIMATION.md index 26072b208f..a6b4de29d2 100644 --- a/studio/backend/utils/hardware/VRAM_ESTIMATION.md +++ b/studio/backend/utils/hardware/VRAM_ESTIMATION.md @@ -33,7 +33,13 @@ Non-quantizable = 2*H*L + V*H + (V*H if not tie_embeddings else 0) | QLoRA 4-bit | `Quantizable * 2 / 3.2 + Non-quantizable * 2` | | LoRA / Full fp16 | `(Quantizable + Non-quantizable) * 2` | -The 3.2 factor (`16/5`) accounts for BNB NF4 blockwise scales. +The 3.2 factor (`16/5`) accounts for BNB NF4 blockwise scales. Repos whose +quantization config enables `bnb_4bit_use_double_quant` use a tighter, still +conservative 3.6 factor for the quantized portion of the weights. +When a 4-bit config has `llm_int8_skip_modules` entries that point to language +model layers or submodules, those quantizable weights are charged at fp16 +instead of NF4. Generic embedding and multimodal skip names are already covered +by non-quantizable terms or excluded from text training weights. ## 2. LoRA Adapters @@ -53,6 +59,18 @@ MLP modules multiply by `E` for MoE. LoRA_bytes = sum(A + B per selected module) * L * 2 ``` +`all-linear` is treated as all known text linear modules in the table above. +The estimator deliberately does not infer multimodal or vision-tower LoRA +modules from config shapes; those modules vary too much across VLM families for +a generic config formula. + +Some decoder configs expose layer-shape fields such as `layer_types`, +`head_dim`, `global_head_dim`, `num_global_key_value_heads`, `attention_k_eq_v`, +`num_kv_shared_layers`, `use_double_wide_mlp`, `vocab_size_per_layer_input`, and +`hidden_size_per_layer_input`. When those fields are present, the estimator +derives text weight and LoRA counts from the per-layer shapes instead of +assuming every layer has the same seven projection modules. + ## 3. Optimizer States (calibrated) | Optimizer | Bytes/param | Notes | @@ -77,6 +95,21 @@ Per-layer (from `unsloth_zoo/vllm_utils.py`): Per_layer = (S*B*(H+K+K) + S*B*2 + S*B*(M+M)) * 2 * 1.25 ``` +When the resolved attention implementation is none of `flash_attention_2`, +`sdpa`, or `flex_attention` (PyTorch SDPA dispatches to flash or +memory-efficient kernels and FlexAttention is also a memory-efficient +kernel, all of which are O(n) in memory), activation memory also includes +a quadratic attention-score/workspace estimate: + +``` +Non_flash_attention = B * num_attention_heads * S^2 * 2 * 12.0 * effective_layers +Activations = max(Per_layer_with_gc, Non_flash_attention) +``` + +Studio resolves the attention implementation with Unsloth's +`resolve_attention_implementation` helper and uses that result directly. The +estimator does not duplicate model-family attention policy. + | GC Mode | Full FT | LoRA/QLoRA | |---------|---------|------------| | none | `L` layers | `L` layers | @@ -85,13 +118,33 @@ Per_layer = (S*B*(H+K+K) + S*B*2 + S*B*(M+M)) * 2 * 1.25 ## 6. Floors -Gradients and activations have minimum floors at **15% of model weight memory** to account for autograd overhead, attention score matrices, NCCL buffers, mixed-precision scaling, and PyTorch fragmentation. +Activations use the computed formula directly: ``` -gradient_bytes = max(computed, weights * 0.15) -activation_bytes = max(computed, weights * 0.15 * B/2) +activation_bytes = computed_activation_bytes ``` +Full fine-tuning keeps the gradient floor at **15% of model weight memory** to +account for autograd overhead, NCCL buffers, mixed-precision scaling, and +PyTorch fragmentation: + +``` +gradient_bytes = max(computed_gradient_bytes, weights * 0.15) +``` + +For LoRA/QLoRA, the base model is frozen, so the weight-derived gradient floor +is capped by trainable-state and live-activation scale: + +``` +raw_gradient_bytes = trainable_params * 2 +gradient_floor = min(weights * 0.15, max(computed_activation_bytes, optimizer_bytes)) +gradient_bytes = max(raw_gradient_bytes, gradient_floor) +``` + +This prevents frozen quantized model size from dominating gradient/state +overhead when the measured runtime footprint is governed by LoRA optimizer +states and live activations. + ## 7. CUDA Overhead **1.4 GB** fixed — CUDA driver + PyTorch runtime, calibrated on RTX 5070 Ti. @@ -106,34 +159,6 @@ usable_gb = free[gpu_0] + sum(free[gpu_i] * 0.85 for i in 1..N) --- -## Reference Table (bsz=2, seq=2048, rank=16, GC=unsloth, adamw_8bit) - -| Model | Weights | LoRA | Optim | Grad | Act | CUDA | Total | -|-------|---------|------|-------|------|-----|------|-------| -| 0.5B QLoRA | 0.5 | 0.0 | 0.0 | 0.1 | 0.1 | 1.4 | **2.1** | -| 1B QLoRA | 1.1 | 0.0 | 0.0 | 0.2 | 0.2 | 1.4 | **2.9** | -| 3B QLoRA | 2.4 | 0.0 | 0.1 | 0.5 | 0.5 | 1.4 | **4.9** | -| 8B QLoRA | 6.0 | 0.1 | 0.2 | 1.2 | 1.2 | 1.4 | **10.1** | -| 8B LoRA fp16 | 15.0 | 0.1 | 0.2 | 3.0 | 3.0 | 1.4 | **22.6** | -| 8B Full FT | 15.0 | — | 29.9 | 15.0 | 3.0 | 1.4 | **64.2** | -| 32B LoRA fp16 | 61.0 | 0.2 | 0.5 | 12.2 | 12.2 | 1.4 | **87.6** | -| 72B QLoRA | 45.5 | 0.4 | 0.8 | 9.1 | 9.1 | 1.4 | **66.3** | - -## E2E Validation (Llama-3.2-1B, B200 emulating 24GB) - -| Config | Estimated | Actual (nvsmi) | Error | -|--------|----------|----------------|-------| -| QLoRA bsz=2 seq=512 | 2.55 GB | 2.65 GB | -3.7% | -| QLoRA bsz=2 seq=2048 | 2.60 GB | 2.65 GB | -1.8% | -| QLoRA bsz=4 seq=2048 | 2.65 GB | 2.65 GB | +0.0% | -| LoRA fp16 bsz=2 | 3.84 GB | 3.88 GB | -1.0% | -| Full FT adamw_8bit | 10.89 GB | 10.80 GB | +0.8% | -| Full FT adamw_torch | 13.19 GB | 12.93 GB | +2.0% | - -*Note: e2e numbers predate the 15% floors, which add safety margin on top.* - ---- - ## Parameter Flow ``` diff --git a/studio/backend/utils/hardware/amd.py b/studio/backend/utils/hardware/amd.py index 755314ca3a..fdb1ab4520 100644 --- a/studio/backend/utils/hardware/amd.py +++ b/studio/backend/utils/hardware/amd.py @@ -16,6 +16,7 @@ import subprocess from typing import Any, Optional from loggers import get_logger +from utils.native_path_leases import child_env_without_native_path_secret logger = get_logger(__name__) @@ -28,6 +29,7 @@ def _run_amd_smi(*args: str, timeout: int = 5) -> Optional[Any]: capture_output = True, text = True, timeout = timeout, + env = child_env_without_native_path_secret(), ) except (OSError, subprocess.TimeoutExpired) as e: logger.warning("amd-smi query failed: %s", e) diff --git a/studio/backend/utils/hardware/hardware.py b/studio/backend/utils/hardware/hardware.py index be31c00a78..c218b7b4b9 100644 --- a/studio/backend/utils/hardware/hardware.py +++ b/studio/backend/utils/hardware/hardware.py @@ -774,6 +774,34 @@ def _load_config_for_gpu_estimate(model_name: str, hf_token: Optional[str] = Non return None +def _determine_attention_impl_for_gpu_estimate(config) -> str: + import copy as _copy + + from unsloth.models._utils import resolve_attention_implementation + from transformers import AutoModel, AutoModelForCausalLM + + # why: resolve_attention_implementation calls _set_attn_impl which writes + # _attn_implementation onto the config; PreTrainedConfig's setter walks + # `sub_configs` and propagates to nested text_config / sub-configs, so a + # shallow copy still mutates those shared inner objects on the cached + # config returned by _load_config_for_gpu_estimate. Deepcopy isolates them. + config_copy = _copy.deepcopy(config) + + model_class = None + for auto_model in (AutoModelForCausalLM, AutoModel): + mapping = getattr(auto_model, "_model_mapping", None) + if mapping is None: + continue + try: + if config_copy.__class__ in mapping: + model_class = mapping[config_copy.__class__] + break + except Exception: + continue + + return resolve_attention_implementation(model_class, config_copy) + + def _estimate_fp16_model_size_bytes_from_config(config) -> Optional[int]: from .vram_estimation import extract_arch_config, compute_total_params @@ -844,12 +872,21 @@ def estimate_fp16_model_size_bytes( return int(total_params * 2), "safetensors" config = _load_config_for_gpu_estimate(estimate_model, hf_token = hf_token) + config_bytes: Optional[int] = None if config is not None: config_bytes = _estimate_fp16_model_size_bytes_from_config(config) - if config_bytes is not None: - return config_bytes, "config" local_bytes = _get_local_weight_size_bytes(estimate_model) + + # why: config-derived bytes cover only the text tower; local safetensors + # include vision/audio towers. Take the larger so the multimodal + # extra_bytes correction can fire. + if config_bytes is not None and local_bytes is not None: + if local_bytes > config_bytes: + return local_bytes, "weight_bytes" + return config_bytes, "config" + if config_bytes is not None: + return config_bytes, "config" if local_bytes is not None: return local_bytes, "weight_bytes" @@ -877,6 +914,9 @@ def estimate_required_model_memory_gb( TrainingVramConfig, extract_arch_config, estimate_training_vram, + compute_total_params, + compute_optimizer_bytes, + compute_gradient_bytes, CUDA_OVERHEAD_BYTES, QUANT_4BIT_FACTOR, DEFAULT_TARGET_MODULES, @@ -926,13 +966,44 @@ def estimate_required_model_memory_gb( model_name, hf_token = hf_token ) config = _load_config_for_gpu_estimate(estimate_model, hf_token = hf_token) + if config is not None: + try: + vram_config.attention_implementation = ( + _determine_attention_impl_for_gpu_estimate(config) + ) + except Exception as e: + logger.warning( + "Could not resolve attention implementation for '%s': %s", + estimate_model, + e, + ) + # why: if we cannot prove flash attention is usable, charge the + # quadratic non-flash activation path so GPU selection stays + # conservative. + vram_config.attention_implementation = "eager" arch = extract_arch_config(config) if config is not None else None if arch is not None: breakdown = estimate_training_vram(arch, vram_config) + # why: extract_arch_config only sees text_config; safetensors include + # vision/audio tower bytes that the text-arch fp16 total misses. + arch_fp16_bytes = compute_total_params(arch) * 2 + extra_bytes = max(0, int(model_size_bytes) - arch_fp16_bytes) + if extra_bytes > 0: + breakdown.model_weights += extra_bytes + if training_method == "full": + # why: full fine-tuning makes the extra (vision/audio) params + # trainable; optimizer + gradient bytes scale with them too. + extra_params = extra_bytes // 2 + breakdown.optimizer_states += compute_optimizer_bytes( + extra_params, + vram_config.optimizer, + ) + breakdown.gradients += compute_gradient_bytes(extra_params) required_gb = breakdown.total / (1024**3) metadata["required_gb"] = round(required_gb, 3) metadata["estimation_mode"] = "detailed" + metadata["attention_implementation"] = vram_config.attention_implementation metadata["vram_breakdown"] = breakdown.to_gb_dict() max_gpus = max(1, get_visible_gpu_count()) for n_gpus in range(1, max_gpus + 1): diff --git a/studio/backend/utils/hardware/nvidia.py b/studio/backend/utils/hardware/nvidia.py index 274d9beb48..099c5fa3a5 100644 --- a/studio/backend/utils/hardware/nvidia.py +++ b/studio/backend/utils/hardware/nvidia.py @@ -6,6 +6,7 @@ from typing import Any, Optional from loggers import get_logger +from utils.native_path_leases import child_env_without_native_path_secret from utils.subprocess_compat import ( windows_hidden_subprocess_kwargs as _windows_hidden_subprocess_kwargs, ) @@ -65,6 +66,7 @@ def get_physical_gpu_count() -> Optional[int]: capture_output = True, text = True, timeout = 5, + env = child_env_without_native_path_secret(), **_windows_hidden_subprocess_kwargs(), ) if result.returncode == 0 and result.stdout.strip(): @@ -90,6 +92,7 @@ def get_primary_gpu_utilization() -> dict[str, Any]: capture_output = True, text = True, timeout = 5, + env = child_env_without_native_path_secret(), **_windows_hidden_subprocess_kwargs(), ) except (OSError, subprocess.TimeoutExpired) as e: @@ -141,6 +144,7 @@ def get_visible_gpu_utilization( capture_output = True, text = True, timeout = 5, + env = child_env_without_native_path_secret(), **_windows_hidden_subprocess_kwargs(), ) except (OSError, subprocess.TimeoutExpired) as e: @@ -227,6 +231,7 @@ def get_backend_visible_gpu_info( capture_output = True, text = True, timeout = 10, + env = child_env_without_native_path_secret(), **_windows_hidden_subprocess_kwargs(), ) except (OSError, subprocess.TimeoutExpired) as e: diff --git a/studio/backend/utils/hardware/vram_estimation.py b/studio/backend/utils/hardware/vram_estimation.py index e03665374d..ba1b1dfe61 100644 --- a/studio/backend/utils/hardware/vram_estimation.py +++ b/studio/backend/utils/hardware/vram_estimation.py @@ -16,7 +16,26 @@ from dataclasses import dataclass, field from typing import Dict, Optional QUANT_4BIT_FACTOR = 16 / 5 +DOUBLE_QUANT_4BIT_FACTOR = ( + 3.6 # bnb_4bit_use_double_quant; see VRAM_ESTIMATION.md section 1 +) CUDA_OVERHEAD_BYTES = int(1.4 * 1024**3) # calibrated on RTX 5070 Ti +NON_FLASH_ATTENTION_FACTOR = ( + 12.0 # eager attention score+workspace overhead; see VRAM_ESTIMATION.md section 5 +) + +LINEAR_ATTENTION_IMPLS = frozenset({"flash_attention_2", "sdpa", "flex_attention"}) + +_SKIP_MODULE_TEXT_PREFIXES = frozenset( + { + "model", + "model.model", + "language_model", + "language_model.model", + "model.language_model", + "model.language_model.model", + } +) DEFAULT_TARGET_MODULES = [ "q_proj", @@ -27,6 +46,8 @@ DEFAULT_TARGET_MODULES = [ "up_proj", "down_proj", ] +ATTENTION_TARGET_MODULES = {"q_proj", "k_proj", "v_proj", "o_proj"} +MLP_TARGET_MODULES = {"gate_proj", "up_proj", "down_proj"} # Empirically calibrated bytes/param — see VRAM_ESTIMATION.md for rationale. OPTIMIZER_BYTES_PER_PARAM: Dict[str, int] = { @@ -61,12 +82,28 @@ class ModelArchConfig: num_experts: Optional[int] = None moe_intermediate_size: Optional[int] = None n_shared_experts: int = 0 + shared_expert_intermediate_size: Optional[int] = None + num_experts_per_tok: int = 1 num_dense_layers: int = 0 q_lora_rank: Optional[int] = None kv_lora_rank: Optional[int] = None qk_nope_head_dim: Optional[int] = None qk_rope_head_dim: Optional[int] = None v_head_dim: Optional[int] = None + head_dim: Optional[int] = None + global_head_dim: Optional[int] = None + num_global_key_value_heads: Optional[int] = None + attention_k_eq_v: bool = False + layer_types: Optional[list] = None + num_kv_shared_layers: int = 0 + use_double_wide_mlp: bool = False + vocab_size_per_layer_input: int = 0 + hidden_size_per_layer_input: int = 0 + quantization_skip_modules: list = field(default_factory = list) + quant_4bit_factor: float = QUANT_4BIT_FACTOR + moe_has_dense_mlp: bool = False + dense_layer_indices: tuple = () + dense_intermediate_size: Optional[int] = None @dataclass @@ -79,6 +116,7 @@ class TrainingVramConfig: gradient_checkpointing: str = "unsloth" optimizer: str = "adamw_8bit" load_in_4bit: bool = True + attention_implementation: str = "flash_attention_2" @dataclass @@ -89,8 +127,8 @@ class VramBreakdown: gradients: int activations: int cuda_overhead: int - # The computed (formula-based) activation cost before floors. - # This is the true per-layer cost that doesn't shard across GPUs. + # Equals `activations`; retained for backward compatibility with + # consumers that read this field. activations_computed: int = 0 @property @@ -108,17 +146,15 @@ class VramBreakdown: """Minimum VRAM a single GPU needs: its shard + non-shardable costs. Weights/LoRA/optimizer/gradients shard across GPUs. - The computed activation cost does NOT shard (one GPU runs the layer). - The floor portion (activations - computed) is overhead that shards. + Activations do NOT shard (the GPU running a layer holds them). """ shardable = ( self.model_weights + self.lora_adapters + self.optimizer_states + self.gradients - + (self.activations - self.activations_computed) # floor overhead shards ) - per_gpu_fixed = self.activations_computed + self.cuda_overhead + per_gpu_fixed = self.activations + self.cuda_overhead return shardable // max(n_gpus, 1) + per_gpu_fixed def to_gb_dict(self) -> Dict[str, float]: @@ -133,28 +169,88 @@ class VramBreakdown: } -def _compute_num_dense_layers(text_config, total_layers: int) -> int: - """Count how many layers use dense MLP instead of MoE.""" +def _first_scalar(value): + # why: ERNIE MoE configs ship moe_intermediate_size / moe_num_experts as + # [routed, shared] lists; downstream arithmetic needs the routed scalar. + if isinstance(value, (list, tuple)): + return value[0] if value else None + return value + + +def _max_scalar(value): + # why: Hunyuan-V1-MoE moe_topk can be a per-layer list; activation + # accounting uses the max top-k as a conservative upper bound. + if isinstance(value, (list, tuple)): + items = [v for v in value if v is not None] + return max(items) if items else None + return value + + +def _compute_dense_layer_indices(text_config, total_layers: int) -> tuple: + """Layer indices that use dense MLP instead of MoE. Position matters.""" + # why: transformers Exaone-MoE / Laguna / Hy_v3 / GLM-MoE-DSA / GLM4-MoE-Lite / + # Ernie4_5_VL_MoE prefer per-position `mlp_layer_types` over the prefix-style + # `first_k_dense_replace` and may omit `decoder_sparse_step` entirely. + layer_types = getattr(text_config, "mlp_layer_types", None) + if layer_types: + return tuple( + i + for i, t in enumerate(layer_types[:total_layers]) + if str(t).lower() == "dense" + ) + + # why: Llama4TextConfig.__init__ auto-populates self.moe_layers from + # interleave_moe_layer_step; Llama4TextDecoderLayer dispatches via + # `layer_idx in config.moe_layers` (modeling_llama4.py). + llama4_moe_layers = getattr(text_config, "moe_layers", None) + if llama4_moe_layers is not None: + moe_indices = {int(i) for i in llama4_moe_layers} + return tuple(i for i in range(total_layers) if i not in moe_indices) + + # why: transformers ERNIE 4.5 MoE / ERNIE 4.5 VL MoE declare MoE layers + # via moe_layer_start_index / moe_layer_end_index / moe_layer_interval; + # the model's per-layer guard is `(layer_idx + 1) % interval == 0` with + # start <= layer_idx <= end (modeling_ernie4_5_moe.py). + moe_start = getattr(text_config, "moe_layer_start_index", None) + moe_interval = getattr(text_config, "moe_layer_interval", None) + if moe_start is not None and moe_interval is not None and int(moe_interval) > 0: + moe_end_raw = getattr(text_config, "moe_layer_end_index", None) + end = ( + total_layers + if moe_end_raw is None or int(moe_end_raw) == -1 + else min(int(moe_end_raw) + 1, total_layers) + ) + start = max(0, int(moe_start)) + interval = int(moe_interval) + moe_indices = {i for i in range(start, end) if (i + 1) % interval == 0} + return tuple(i for i in range(total_layers) if i not in moe_indices) + first_k = getattr(text_config, "first_k_dense_replace", None) if first_k is not None: - return min(int(first_k), total_layers) + return tuple(range(min(int(first_k), total_layers))) sparse_step = getattr(text_config, "decoder_sparse_step", None) mlp_only = getattr(text_config, "mlp_only_layers", None) or [] if sparse_step is not None and sparse_step > 0: - mlp_only_set = set(mlp_only) - moe_count = sum( - 1 + mlp_only_set = {int(i) for i in mlp_only} + return tuple( + i for i in range(total_layers) - if i not in mlp_only_set and (i + 1) % sparse_step == 0 + if i in mlp_only_set or (i + 1) % sparse_step != 0 ) - return total_layers - moe_count - - return 0 + return () def extract_arch_config(hf_config) -> Optional[ModelArchConfig]: text_config = getattr(hf_config, "text_config", None) or hf_config + quantization_config = getattr(hf_config, "quantization_config", None) or {} + if not isinstance(quantization_config, dict): + quantization_config = getattr(quantization_config, "to_dict", lambda: {})() + quant_4bit_factor = ( + DOUBLE_QUANT_4BIT_FACTOR + if quantization_config.get("bnb_4bit_use_double_quant", False) + else QUANT_4BIT_FACTOR + ) hidden_size = getattr(text_config, "hidden_size", None) num_layers = getattr(text_config, "num_hidden_layers", None) @@ -177,18 +273,75 @@ def extract_arch_config(hf_config) -> Optional[ModelArchConfig]: num_kv_heads = getattr(text_config, "num_key_value_heads", num_heads) + # why: DBRX places its MoE attrs on the DbrxFFNConfig sub-config; probe + # ffn_config as a secondary source so DBRX is not misclassified as dense. + ffn_config = getattr(text_config, "ffn_config", None) + + def _moe_attr(name): + value = getattr(text_config, name, None) + if value is None and ffn_config is not None: + value = getattr(ffn_config, name, None) + return value + num_experts = None - for attr in ("num_local_experts", "num_experts", "n_routed_experts"): - num_experts = getattr(text_config, attr, None) + for attr in ( + "num_local_experts", + "num_experts", + "n_routed_experts", + "moe_num_experts", + ): + num_experts = _first_scalar(_moe_attr(attr)) if num_experts is not None: break - moe_intermediate = getattr(text_config, "moe_intermediate_size", None) - n_shared_experts = getattr(text_config, "n_shared_experts", None) or 0 + moe_intermediate_raw = _moe_attr("moe_intermediate_size") + if moe_intermediate_raw is None: + moe_intermediate_raw = _moe_attr("ffn_hidden_size") + moe_intermediate = _first_scalar(moe_intermediate_raw) + # why: Exaone-MoE / ERNIE families alias num_shared_experts / + # moe_num_shared_experts to the canonical n_shared_experts. + n_shared_experts = ( + _first_scalar(_moe_attr("n_shared_experts")) + or _first_scalar(_moe_attr("num_shared_experts")) + or _first_scalar(_moe_attr("moe_num_shared_experts")) + or 0 + ) + shared_expert_intermediate_size = _moe_attr("shared_expert_intermediate_size") + if shared_expert_intermediate_size and n_shared_experts == 0: + n_shared_experts = 1 + # why: DBRX exposes moe_top_k, Hunyuan-V1-MoE exposes moe_topk (which can + # be a per-layer list); _max_scalar normalizes list values to the worst + # case so int(...) below cannot crash on the canonical attribute_map path. + num_experts_per_tok = ( + _max_scalar(_moe_attr("num_experts_per_tok")) + or _max_scalar(_moe_attr("top_k_experts")) + or _max_scalar(_moe_attr("moe_top_k")) + or _max_scalar(_moe_attr("moe_topk")) + or 1 + ) - num_dense_layers = 0 + dense_layer_indices: tuple = () if num_experts is not None and num_experts > 1: - num_dense_layers = _compute_num_dense_layers(text_config, num_layers) + dense_layer_indices = _compute_dense_layer_indices(text_config, num_layers) + num_dense_layers = len(dense_layer_indices) + + # why: Llama4 dense layers use intermediate_size_mlp; routed and shared + # experts use intermediate_size. Llama4TextMoe builds one shared_expert + # per MoE layer (modeling_llama4.py). + intermediate_size_mlp_raw = _first_scalar(_moe_attr("intermediate_size_mlp")) + dense_intermediate_size = ( + int(intermediate_size_mlp_raw) + if intermediate_size_mlp_raw is not None + else None + ) + if ( + intermediate_size_mlp_raw is not None + and num_experts is not None + and num_experts > 1 + and shared_expert_intermediate_size is None + and n_shared_experts == 0 + ): + n_shared_experts = 1 q_lora_rank = getattr(text_config, "q_lora_rank", None) kv_lora_rank = getattr(text_config, "kv_lora_rank", None) @@ -207,15 +360,418 @@ def extract_arch_config(hf_config) -> Optional[ModelArchConfig]: num_experts = num_experts, moe_intermediate_size = moe_intermediate, n_shared_experts = n_shared_experts, + shared_expert_intermediate_size = shared_expert_intermediate_size, + num_experts_per_tok = int(num_experts_per_tok), num_dense_layers = num_dense_layers, q_lora_rank = q_lora_rank, kv_lora_rank = kv_lora_rank, qk_nope_head_dim = qk_nope_head_dim, qk_rope_head_dim = qk_rope_head_dim, v_head_dim = v_head_dim, + head_dim = getattr(text_config, "head_dim", None), + global_head_dim = getattr(text_config, "global_head_dim", None), + num_global_key_value_heads = getattr( + text_config, + "num_global_key_value_heads", + None, + ), + attention_k_eq_v = bool(getattr(text_config, "attention_k_eq_v", False)), + layer_types = getattr(text_config, "layer_types", None), + num_kv_shared_layers = getattr(text_config, "num_kv_shared_layers", None) or 0, + use_double_wide_mlp = bool(getattr(text_config, "use_double_wide_mlp", False)), + vocab_size_per_layer_input = getattr( + text_config, + "vocab_size_per_layer_input", + None, + ) + or 0, + hidden_size_per_layer_input = getattr( + text_config, + "hidden_size_per_layer_input", + None, + ) + or 0, + quantization_skip_modules = list( + quantization_config.get("llm_int8_skip_modules", []) or [] + ), + quant_4bit_factor = quant_4bit_factor, + moe_has_dense_mlp = bool(getattr(text_config, "enable_moe_block", False)), + dense_layer_indices = dense_layer_indices, + dense_intermediate_size = dense_intermediate_size, ) +def _targets_all_linear(target_modules) -> bool: + # why: peft LoraConfig accepts target_modules="all-linear" as a bare + # string; iterating a string yields chars and never matches the set. + if isinstance(target_modules, str): + target_modules = [target_modules] + normalized = {str(module).lower().replace("_", "-") for module in target_modules} + return normalized == {"all-linear"} + + +def _head_dim(arch: ModelArchConfig) -> int: + return arch.head_dim or arch.hidden_size // arch.num_attention_heads + + +def _layer_types(arch: ModelArchConfig) -> list: + if arch.layer_types and len(arch.layer_types) == arch.num_hidden_layers: + return arch.layer_types + return ["full_attention"] * arch.num_hidden_layers + + +def _uses_structured_layer_shapes(arch: ModelArchConfig) -> bool: + # MLA configs have their own q/kv low-rank projection shape formulas in + # _compute_attn_elements / _lora_attn_elements; do not let head_dim or + # other structured fields override that path. + if arch.q_lora_rank is not None: + return False + return bool( + arch.layer_types + or arch.head_dim is not None + or arch.global_head_dim is not None + or arch.num_global_key_value_heads is not None + or arch.attention_k_eq_v + or arch.num_kv_shared_layers > 0 + or arch.use_double_wide_mlp + ) + + +def _is_kv_shared_layer(arch: ModelArchConfig, layer_idx: int) -> bool: + if arch.num_kv_shared_layers <= 0: + return False + first_shared = arch.num_hidden_layers - arch.num_kv_shared_layers + # why: transformers Gemma4 (modeling_gemma4.py:1031, modular_gemma4.py:863) + # uses the same `> 0` guard so a fully-shared config raises during model + # construction; matching upstream avoids producing a detailed estimate + # for a shape the actual model code rejects. + return layer_idx >= first_shared > 0 + + +def _is_dense_mlp_layer(arch: ModelArchConfig, layer_idx: int) -> bool: + if arch.dense_layer_indices: + return layer_idx in arch.dense_layer_indices + return layer_idx < arch.num_dense_layers + + +def _per_layer_input_quantizable(arch: ModelArchConfig) -> int: + # why: Gemma4 PLE block adds per_layer_model_projection (single Linear), + # per_layer_input_gate (per layer), and per_layer_projection (per layer); + # see transformers gemma4/modular_gemma4.py:1077-1083 and :1247-1253. + pli = arch.hidden_size_per_layer_input + if pli <= 0: + return 0 + n_layers = arch.num_hidden_layers + hd = arch.hidden_size + return hd * (n_layers * pli) + (hd * pli) * n_layers + (pli * hd) * n_layers + + +def _per_layer_input_norm_elements(arch: ModelArchConfig) -> int: + pli = arch.hidden_size_per_layer_input + if pli <= 0: + return 0 + n_layers = arch.num_hidden_layers + hd = arch.hidden_size + return hd * n_layers + pli + + +def _per_layer_input_lora_params( + arch: ModelArchConfig, + r: int, + target_modules, +) -> int: + # why: Unsloth's get_peft_regex (unsloth_zoo/peft_utils.py) requires module + # names to contain a component tag (mlp/attn/...); PLE module names lack + # any tag, so all-linear training does NOT attach LoRA to them. Only count + # PLE LoRA when the user explicitly names PLE modules. + pli = arch.hidden_size_per_layer_input + if pli <= 0: + return 0 + targets = ( + {target_modules} + if isinstance(target_modules, str) + else set(target_modules or []) + ) + n_layers = arch.num_hidden_layers + hd = arch.hidden_size + total = 0 + if "per_layer_model_projection" in targets: + total += hd * r + r * (n_layers * pli) + if "per_layer_input_gate" in targets: + total += (hd * r + r * pli) * n_layers + if "per_layer_projection" in targets: + total += (pli * r + r * hd) * n_layers + return total + + +def _layer_attention_dims(arch: ModelArchConfig, layer_idx: int) -> tuple: + layer_types = _layer_types(arch) + layer_type = layer_types[layer_idx] + is_sliding = layer_type == "sliding_attention" + head_dim = ( + arch.global_head_dim + if not is_sliding and arch.global_head_dim + else _head_dim(arch) + ) + use_alt_attention = arch.attention_k_eq_v and not is_sliding + num_kv_heads = ( + arch.num_global_key_value_heads + if use_alt_attention and arch.num_global_key_value_heads + else arch.num_key_value_heads + ) + q_size = arch.num_attention_heads * head_dim + kv_size = num_kv_heads * head_dim + has_k = not _is_kv_shared_layer(arch, layer_idx) + has_v = has_k and not use_alt_attention + return q_size, kv_size, has_k, has_v + + +def _layer_mlp_size(arch: ModelArchConfig, layer_idx: int) -> int: + if arch.use_double_wide_mlp and _is_kv_shared_layer(arch, layer_idx): + return _dense_mlp_size(arch) * 2 + return _dense_mlp_size(arch) + + +def _text_linear_dims( + arch: ModelArchConfig, + layer_idx: int, +) -> Dict[str, tuple[int, int]]: + hd = arch.hidden_size + if _uses_structured_layer_shapes(arch): + q_size, kv_size, has_k, has_v = _layer_attention_dims(arch, layer_idx) + mlp_size = _layer_mlp_size(arch, layer_idx) + else: + q_size = hd + kv_size = _get_kv_size(arch) + has_k = True + has_v = True + mlp_size = _get_mlp_size(arch) + + dims = { + "q_proj": (hd, q_size), + "o_proj": (q_size, hd), + } + if has_k: + dims["k_proj"] = (hd, kv_size) + if has_v: + dims["v_proj"] = (hd, kv_size) + + dims.update( + { + "gate_proj": (hd, mlp_size), + "up_proj": (hd, mlp_size), + "down_proj": (mlp_size, hd), + } + ) + return dims + + +def _module_path_matches(skip_module: str, alias: str) -> bool: + skip_parts = [part for part in skip_module.split(".") if part] + alias_parts = [part for part in alias.split(".") if part] + if not skip_parts or not alias_parts: + return False + if alias_parts[0] == "layers": + return skip_parts == alias_parts + if len(skip_parts) <= len(alias_parts): + # why: transformers BNB quantizer suffix-matches short skip entries + # like ["q_proj"] / ["lm_head"] against full module paths, so a skip + # shorter than the alias is a tail match. + return alias_parts[-len(skip_parts) :] == skip_parts + if skip_parts[-len(alias_parts) :] != alias_parts: + return False + prefix_parts = skip_parts[: len(skip_parts) - len(alias_parts)] + if not prefix_parts: + return True + # why: bound the prefix to known text-tower roots so VLM skip names like + # vision_tower.model.layers..self_attn.q_proj do not shadow the text + # alias model.layers..self_attn.q_proj. + return ".".join(prefix_parts) in _SKIP_MODULE_TEXT_PREFIXES + + +def _add_module_aliases( + aliases: Dict[str, str], + canonical: str, + suffix: str, +) -> None: + for prefix in ( + "", + "model", + "model.model", + "language_model", + "language_model.model", + "model.language_model", + "model.language_model.model", + ): + alias = f"{prefix}.{suffix}" if prefix else suffix + aliases[alias] = canonical + + +def _build_text_module_elements( + arch: ModelArchConfig, +) -> tuple[Dict[str, int], Dict[str, str]]: + elements: Dict[str, int] = {} + aliases: Dict[str, str] = {} + + is_mla = arch.q_lora_rank is not None and not _uses_structured_layer_shapes(arch) + pli = arch.hidden_size_per_layer_input + hd_global = arch.hidden_size + + for layer_idx in range(arch.num_hidden_layers): + layer_modules: Dict[str, int] = {} + dims = _text_linear_dims(arch, layer_idx) + attn_dims = { + name: dim for name, dim in dims.items() if name in ATTENTION_TARGET_MODULES + } + mlp_dims = { + name: dim for name, dim in dims.items() if name in MLP_TARGET_MODULES + } + + if is_mla: + # why: _text_linear_dims uses (hd, hd) for q/o; MLA actually splits + # into q_a/q_b/kv_a/kv_b, so emit a single self_attn aggregate at + # the authoritative MLA per-layer total. + layer_modules["self_attn"] = _compute_attn_elements(arch) + else: + for name, (in_dim, out_dim) in attn_dims.items(): + layer_modules[f"self_attn.{name}"] = in_dim * out_dim + + if arch.num_experts and arch.num_experts > 1: + if _is_dense_mlp_layer(arch, layer_idx): + layer_modules.update( + { + f"mlp.{name}": in_dim * out_dim + for name, (in_dim, out_dim) in mlp_dims.items() + } + ) + else: + layer_modules["mlp.experts"] = _compute_routed_moe_elements(arch) + shared_moe = _compute_shared_moe_elements(arch) + if shared_moe: + # why: Qwen3.5-MoE exposes shared expert as + # mlp.shared_expert; Exaone-MoE/Laguna/GLM-style configs use + # mlp.shared_experts. Register both names so child-path + # llm_int8_skip_modules entries match the right shared block. + layer_modules["mlp.shared_expert"] = shared_moe + if arch.moe_has_dense_mlp: + # why: enable_moe_block runs the dense MLP and the MoE + # experts in parallel; register both for skip matching. + # Non-structured _text_linear_dims returns mlp_size from + # _get_mlp_size which prefers moe_intermediate_size, so + # rebuild dense dims from arch.intermediate_size directly. + if _uses_structured_layer_shapes(arch): + dense_dims = mlp_dims + else: + hd = arch.hidden_size + inter = arch.intermediate_size + dense_dims = { + "gate_proj": (hd, inter), + "up_proj": (hd, inter), + "down_proj": (inter, hd), + } + layer_modules.update( + { + f"mlp.{name}": in_dim * out_dim + for name, (in_dim, out_dim) in dense_dims.items() + } + ) + else: + layer_modules.update( + { + f"mlp.{name}": in_dim * out_dim + for name, (in_dim, out_dim) in mlp_dims.items() + } + ) + + if pli > 0: + # why: register PLE per-layer linears so llm_int8_skip_modules + # entries like model.layers.0.per_layer_input_gate match. + layer_modules["per_layer_input_gate"] = hd_global * pli + layer_modules["per_layer_projection"] = pli * hd_global + + attn_total = sum( + value + for name, value in layer_modules.items() + if name == "self_attn" or name.startswith("self_attn.") + ) + # why: gemma4 enable_moe_block puts routed experts at the sibling + # layers..experts attribute, not under self.mlp; the layer's "mlp" + # aggregate must reflect only the dense MLP path so a skip module + # `model.layers.0.mlp` does not over-skip into the experts block. + is_sibling_experts = bool(arch.moe_has_dense_mlp) + mlp_total = sum( + value + for name, value in layer_modules.items() + if ( + name == "mlp" + or ( + name.startswith("mlp.") + and not (is_sibling_experts and name == "mlp.experts") + ) + ) + ) + experts_total = layer_modules.get("mlp.experts", 0) if is_sibling_experts else 0 + layer_total = sum(layer_modules.values()) + + aggregate_modules = { + f"text.layers.{layer_idx}": layer_total, + f"text.layers.{layer_idx}.self_attn": attn_total, + f"text.layers.{layer_idx}.mlp": mlp_total, + } + if experts_total: + aggregate_modules[f"text.layers.{layer_idx}.experts"] = experts_total + elements.update(aggregate_modules) + for canonical in aggregate_modules: + suffix = canonical.removeprefix("text.") + _add_module_aliases(aliases, canonical, suffix) + + for name, value in layer_modules.items(): + canonical = f"text.layers.{layer_idx}.{name}" + elements[canonical] = value + _add_module_aliases(aliases, canonical, canonical.removeprefix("text.")) + if name == "mlp.experts" and arch.moe_has_dense_mlp: + # why: gemma4 enable_moe_block exposes routed experts at + # layers..experts (sibling of self.mlp), not under mlp. + _add_module_aliases(aliases, canonical, f"layers.{layer_idx}.experts") + elif name == "mlp.shared_expert": + # why: Exaone-MoE / Laguna / GLM-style configs use the plural + # `shared_experts` attribute name; register both spellings. + _add_module_aliases( + aliases, + canonical, + f"layers.{layer_idx}.mlp.shared_experts", + ) + + if pli > 0: + canonical = "text.per_layer_model_projection" + elements[canonical] = hd_global * (arch.num_hidden_layers * pli) + _add_module_aliases(aliases, canonical, canonical.removeprefix("text.")) + + return elements, aliases + + +def _compute_skipped_quantizable_elements(arch: ModelArchConfig) -> int: + if not arch.quantization_skip_modules: + return 0 + + module_elements, aliases = _build_text_module_elements(arch) + matched = set() + for skip_module in arch.quantization_skip_modules: + for alias, canonical in aliases.items(): + if _module_path_matches(skip_module, alias): + matched.add(canonical) + + pruned = { + canonical + for canonical in matched + if not any( + canonical != parent and canonical.startswith(f"{parent}.") + for parent in matched + ) + } + return sum(module_elements[canonical] for canonical in pruned) + + def _get_kv_size(arch: ModelArchConfig) -> int: return (arch.hidden_size // arch.num_attention_heads) * arch.num_key_value_heads @@ -226,6 +782,12 @@ def _get_mlp_size(arch: ModelArchConfig) -> int: return arch.intermediate_size +def _dense_mlp_size(arch: ModelArchConfig) -> int: + # why: Llama4 dense layers use intermediate_size_mlp; routed/shared + # experts use intermediate_size. Other configs leave the field None. + return arch.dense_intermediate_size or arch.intermediate_size + + def _get_num_experts(arch: ModelArchConfig) -> int: return arch.num_experts if arch.num_experts and arch.num_experts > 1 else 1 @@ -248,14 +810,39 @@ def _compute_attn_elements(arch: ModelArchConfig) -> int: def _compute_dense_mlp_elements(arch: ModelArchConfig) -> int: - return arch.hidden_size * arch.intermediate_size * 3 + return arch.hidden_size * _dense_mlp_size(arch) * 3 + + +def _shared_expert_size(arch: ModelArchConfig) -> int: + # why: Qwen3.5-MoE shared expert has its own intermediate_size (default 512) + # distinct from moe_intermediate_size; fall back to routed mlp_size for + # families that share it (deepseek-style configs). + return arch.shared_expert_intermediate_size or _get_mlp_size(arch) + + +def _compute_routed_moe_elements(arch: ModelArchConfig) -> int: + hd = arch.hidden_size + n_experts = _get_num_experts(arch) + return hd * _get_mlp_size(arch) * 3 * n_experts + n_experts * hd + + +def _compute_shared_moe_elements(arch: ModelArchConfig) -> int: + if not arch.n_shared_experts: + return 0 + hd = arch.hidden_size + shared_size = _shared_expert_size(arch) + total = hd * shared_size * 3 * arch.n_shared_experts + # why: only Qwen2-MoE / Qwen3.5-MoE define a shared_expert_gate Linear + # (hidden_size→1); other families (Exaone-MoE, HY-V3, GLM4-MoE-Lite, Laguna) + # have shared_experts without a gate. shared_expert_intermediate_size is the + # Qwen-style discriminator. + if arch.shared_expert_intermediate_size: + total += arch.n_shared_experts * hd + return total def _compute_moe_mlp_elements(arch: ModelArchConfig) -> int: - hd = arch.hidden_size - mlp_size = _get_mlp_size(arch) - n_experts = _get_num_experts(arch) - return hd * mlp_size * 3 * (n_experts + arch.n_shared_experts) + n_experts * hd + return _compute_routed_moe_elements(arch) + _compute_shared_moe_elements(arch) def _compute_layer_elements(arch: ModelArchConfig): @@ -267,22 +854,60 @@ def _compute_layer_elements(arch: ModelArchConfig): n_layers = arch.num_hidden_layers n_experts = _get_num_experts(arch) - attn_total = _compute_attn_elements(arch) * n_layers - - if n_experts > 1: + if _uses_structured_layer_shapes(arch): + attn_total = 0 + per_layer_dense_mlp = [] + for layer_idx in range(n_layers): + layer_dense_mlp = 0 + for name, (in_dim, out_dim) in _text_linear_dims( + arch, + layer_idx, + ).items(): + elements = in_dim * out_dim + if name in ATTENTION_TARGET_MODULES: + attn_total += elements + elif name in MLP_TARGET_MODULES: + layer_dense_mlp += elements + per_layer_dense_mlp.append(layer_dense_mlp) + if n_experts > 1: + n_dense = arch.num_dense_layers + n_moe = n_layers - n_dense + moe_mlp_total = _compute_moe_mlp_elements(arch) * n_moe + if arch.moe_has_dense_mlp: + # why: enable_moe_block runs dense MLP and MoE experts in + # parallel; count dense for every layer alongside MoE. + mlp_total = sum(per_layer_dense_mlp) + moe_mlp_total + else: + dense_only_total = sum( + value + for i, value in enumerate(per_layer_dense_mlp) + if _is_dense_mlp_layer(arch, i) + ) + mlp_total = moe_mlp_total + dense_only_total + else: + mlp_total = sum(per_layer_dense_mlp) + elif n_experts > 1: + attn_total = _compute_attn_elements(arch) * n_layers n_dense = arch.num_dense_layers n_moe = n_layers - n_dense - mlp_total = ( - _compute_moe_mlp_elements(arch) * n_moe - + _compute_dense_mlp_elements(arch) * n_dense - ) + moe_mlp_total = _compute_moe_mlp_elements(arch) * n_moe + if arch.moe_has_dense_mlp: + mlp_total = _compute_dense_mlp_elements(arch) * n_layers + moe_mlp_total + else: + mlp_total = moe_mlp_total + _compute_dense_mlp_elements(arch) * n_dense else: + attn_total = _compute_attn_elements(arch) * n_layers mlp_total = _compute_dense_mlp_elements(arch) * n_layers layernorms = 2 * hd - embed_tokens = arch.vocab_size * hd + per_layer_embed = ( + arch.vocab_size_per_layer_input * arch.hidden_size_per_layer_input * n_layers + ) + ple_text_linear = _per_layer_input_quantizable(arch) + ple_norms = _per_layer_input_norm_elements(arch) + embed_tokens = arch.vocab_size * hd + per_layer_embed + ple_norms lm_head = 0 if arch.tie_word_embeddings else arch.vocab_size * hd - return attn_total + mlp_total, layernorms, embed_tokens, lm_head + return attn_total + mlp_total + ple_text_linear, layernorms, embed_tokens, lm_head def compute_model_weights_bytes( @@ -295,7 +920,16 @@ def compute_model_weights_bytes( non_quantizable = layernorms * n_layers + embed_tokens + lm_head if training_method == "qlora" and load_in_4bit: - return int(total_quantizable * 2 / QUANT_4BIT_FACTOR + non_quantizable * 2) + skipped_quantizable = min( + _compute_skipped_quantizable_elements(arch), + total_quantizable, + ) + quantized = total_quantizable - skipped_quantizable + return int( + quantized * 2 / arch.quant_4bit_factor + + skipped_quantizable * 2 + + non_quantizable * 2 + ) return int((total_quantizable + non_quantizable) * 2) @@ -363,46 +997,130 @@ def compute_lora_params( lora_rank: int, target_modules: list, ) -> int: + all_linear = _targets_all_linear(target_modules) + selected_modules = list(DEFAULT_TARGET_MODULES) if all_linear else target_modules hd = arch.hidden_size r = lora_rank n_layers = arch.num_hidden_layers n_experts = _get_num_experts(arch) - attn_total = _lora_attn_elements(arch, r, target_modules) * n_layers - - if n_experts > 1: + use_structured_shapes = _uses_structured_layer_shapes(arch) + if use_structured_shapes: + attn_total = 0 + structured_dense_mlp = 0 + per_layer_dense_mlp = [] + for layer_idx in range(n_layers): + layer_dense = 0 + for name, (in_dim, out_dim) in _text_linear_dims( + arch, + layer_idx, + ).items(): + if name not in selected_modules: + continue + if name in ATTENTION_TARGET_MODULES: + attn_total += in_dim * r + r * out_dim + elif name in MLP_TARGET_MODULES: + layer_dense += in_dim * r + r * out_dim + per_layer_dense_mlp.append(layer_dense) + structured_dense_mlp += layer_dense + if n_experts > 1: + n_dense = arch.num_dense_layers + n_moe = n_layers - n_dense + # why: peft "all-linear" attaches LoRA to nn.Linear only; + # routed experts are nn.Parameter and need explicit + # gate_proj/up_proj/down_proj naming via Unsloth's + # get_moe_target_parameters. Shared experts are nn.Linear and + # are picked up by get_peft_regex. + routed_moe = ( + 0 + if all_linear + else _lora_mlp_elements( + hd, + _get_mlp_size(arch), + r, + selected_modules, + n_experts, + ) + ) + shared_moe = _lora_mlp_elements( + hd, + _shared_expert_size(arch), + r, + selected_modules, + arch.n_shared_experts, + ) + moe_mlp = routed_moe + shared_moe + if arch.moe_has_dense_mlp: + # why: parallel dense MLP coexists with MoE on every layer. + mlp_total = structured_dense_mlp + moe_mlp * n_moe + else: + dense_only = sum( + value + for i, value in enumerate(per_layer_dense_mlp) + if _is_dense_mlp_layer(arch, i) + ) + mlp_total = moe_mlp * n_moe + dense_only + else: + mlp_total = structured_dense_mlp + return ( + attn_total + + mlp_total + + _per_layer_input_lora_params(arch, r, target_modules) + ) + elif n_experts > 1: + attn_total = _lora_attn_elements(arch, r, selected_modules) * n_layers n_dense = arch.num_dense_layers n_moe = n_layers - n_dense - # Include shared experts alongside routed experts - moe_expert_mult = n_experts + arch.n_shared_experts - moe_mlp = _lora_mlp_elements( - hd, - _get_mlp_size(arch), - r, - target_modules, - moe_expert_mult, + # why: routed and shared experts may use different intermediate sizes + # (Qwen3.5-MoE: routed mlp_size != shared_expert_intermediate_size). + # See structured branch for the all-linear exclusion rationale; only + # routed (nn.Parameter) experts are excluded under all-linear. + routed_moe = ( + 0 + if all_linear + else _lora_mlp_elements( + hd, + _get_mlp_size(arch), + r, + selected_modules, + n_experts, + ) ) + shared_moe = _lora_mlp_elements( + hd, + _shared_expert_size(arch), + r, + selected_modules, + arch.n_shared_experts, + ) + moe_mlp = routed_moe + shared_moe dense_mlp = _lora_mlp_elements( hd, - arch.intermediate_size, + _dense_mlp_size(arch), r, - target_modules, + selected_modules, 1, ) - mlp_total = moe_mlp * n_moe + dense_mlp * n_dense + if arch.moe_has_dense_mlp: + mlp_total = moe_mlp * n_moe + dense_mlp * n_layers + else: + mlp_total = moe_mlp * n_moe + dense_mlp * n_dense else: + attn_total = _lora_attn_elements(arch, r, selected_modules) * n_layers mlp_total = ( _lora_mlp_elements( hd, - arch.intermediate_size, + _dense_mlp_size(arch), r, - target_modules, + selected_modules, 1, ) * n_layers ) - return attn_total + mlp_total + return ( + attn_total + mlp_total + _per_layer_input_lora_params(arch, r, target_modules) + ) def compute_lora_adapter_bytes(lora_params: int) -> int: @@ -419,26 +1137,88 @@ def compute_gradient_bytes(trainable_params: int) -> int: return trainable_params * 2 +def _is_linear_attention(attention_implementation: Optional[str]) -> bool: + # why: PyTorch SDPA dispatches to flash/memory-efficient O(n) backends; only + # eager (and other non-flash impls) need the quadratic correction. + return attention_implementation in LINEAR_ATTENTION_IMPLS + + +def _compute_non_flash_attention_bytes( + arch: ModelArchConfig, + batch_size: int, + seq_len: int, + effective_layers: float, +) -> int: + score_elements = batch_size * arch.num_attention_heads * seq_len * seq_len + return int(score_elements * 2 * NON_FLASH_ATTENTION_FACTOR * effective_layers) + + +def _layer_qkv_mlp_sizes(arch: ModelArchConfig, layer_idx: int) -> tuple: + n_experts = _get_num_experts(arch) + is_moe_layer = n_experts > 1 and not _is_dense_mlp_layer(arch, layer_idx) + if _uses_structured_layer_shapes(arch): + q_size, kv_size, _has_k, _has_v = _layer_attention_dims(arch, layer_idx) + # why: KV-shared layers (Gemma4/Gemma3n) drop k_proj/v_proj WEIGHTS but + # the donor layer's K/V tensors stay alive across the shared range, so + # activation memory still pays for kv_size; only the weight path uses + # has_k/has_v. + layer_type = _layer_types(arch)[layer_idx] + use_alt_attention = arch.attention_k_eq_v and layer_type != "sliding_attention" + kv_count = 1 if use_alt_attention else 2 + qkv_size = q_size + kv_size * kv_count + if is_moe_layer: + # why: each token routes through `num_experts_per_tok` experts; their + # gate/up/down intermediates are all live during MLP forward. + mlp_size = _get_mlp_size(arch) * arch.num_experts_per_tok + if arch.n_shared_experts: + mlp_size += _shared_expert_size(arch) * arch.n_shared_experts + if arch.moe_has_dense_mlp: + mlp_size += _layer_mlp_size(arch, layer_idx) + else: + mlp_size = _layer_mlp_size(arch, layer_idx) + return qkv_size, mlp_size + kv_size = _get_kv_size(arch) + if is_moe_layer: + mlp_size = _get_mlp_size(arch) * arch.num_experts_per_tok + if arch.n_shared_experts: + mlp_size += _shared_expert_size(arch) * arch.n_shared_experts + if arch.moe_has_dense_mlp: + mlp_size += arch.intermediate_size + else: + mlp_size = _get_mlp_size(arch) + return arch.hidden_size + kv_size + kv_size, mlp_size + + +def _per_layer_activation_bytes( + arch: ModelArchConfig, + layer_idx: int, + batch_size: int, + seq_len: int, +) -> int: + qkv_size, mlp_size = _layer_qkv_mlp_sizes(arch, layer_idx) + activation_qkv = seq_len * batch_size * qkv_size + residual_memory = (seq_len * batch_size) * 2 + activation_mlp = seq_len * batch_size * (mlp_size + mlp_size) + # why: per_layer_input_gate (hd-sized) and per_layer_projection (pli-sized) + # outputs materialize once per decoder layer when hidden_size_per_layer_input + # is set; see gemma4/modular_gemma4.py:1141-1145. + pli = arch.hidden_size_per_layer_input + activation_ple = seq_len * batch_size * (arch.hidden_size + pli) if pli > 0 else 0 + return int( + (activation_qkv + residual_memory + activation_mlp + activation_ple) * 2 * 1.25 + ) + + def compute_activation_bytes( arch: ModelArchConfig, batch_size: int, seq_len: int, gradient_checkpointing: str, is_lora: bool = False, + attention_implementation: Optional[str] = "flash_attention_2", ) -> int: - hd = arch.hidden_size - kv_size = _get_kv_size(arch) - mlp_size = _get_mlp_size(arch) - bsz = batch_size n_layers = arch.num_hidden_layers - activation_qkv = seq_len * bsz * (hd + kv_size + kv_size) - residual_memory = (seq_len * bsz) * 2 - activation_mlp = seq_len * bsz * (mlp_size + mlp_size) - - per_layer_bytes = (activation_qkv + residual_memory + activation_mlp) * 2 - per_layer_bytes = int(per_layer_bytes * 1.25) - gc_key = gradient_checkpointing.lower() gc_entry = GC_LAYER_MULTIPLIERS.get(gc_key, (None, None)) full_ft_mult, lora_mult = gc_entry @@ -446,10 +1226,35 @@ def compute_activation_bytes( if gc_multiplier is None: effective_layers = n_layers + linear_bytes = sum( + _per_layer_activation_bytes(arch, i, batch_size, seq_len) + for i in range(n_layers) + ) else: effective_layers = gc_multiplier + max_layer_bytes = max( + _per_layer_activation_bytes(arch, i, batch_size, seq_len) + for i in range(n_layers) + ) + linear_bytes = int(max_layer_bytes * effective_layers) - return int(per_layer_bytes * effective_layers) + # why: gemma4 per_layer_model_projection runs once outside the per-decoder + # loop and materializes a [B, S, L, PLI] tensor; see modular_gemma4.py:1247. + pli = arch.hidden_size_per_layer_input + if pli > 0: + linear_bytes += int(seq_len * batch_size * n_layers * pli * 2 * 1.25) + + if _is_linear_attention(attention_implementation): + return linear_bytes + return max( + linear_bytes, + _compute_non_flash_attention_bytes( + arch, + batch_size, + seq_len, + effective_layers, + ), + ) def estimate_training_vram( @@ -474,21 +1279,23 @@ def estimate_training_vram( trainable_params = lora_params if is_lora else compute_total_params(arch) optimizer_bytes = compute_optimizer_bytes(trainable_params, config.optimizer) - gradient_bytes = max( - compute_gradient_bytes(trainable_params), - int(model_weights * 0.15), - ) activations_computed = compute_activation_bytes( arch, config.batch_size, config.max_seq_length, config.gradient_checkpointing, is_lora = is_lora, + attention_implementation = config.attention_implementation, ) - activation_bytes = max( - activations_computed, - int(model_weights * 0.15 * (config.batch_size / 2)), - ) + raw_gradient_bytes = compute_gradient_bytes(trainable_params) + gradient_floor = int(model_weights * 0.15) + if is_lora: + gradient_floor = min( + gradient_floor, + max(activations_computed, optimizer_bytes), + ) + gradient_bytes = max(raw_gradient_bytes, gradient_floor) + activation_bytes = activations_computed return VramBreakdown( model_weights = model_weights, diff --git a/studio/backend/utils/models/model_config.py b/studio/backend/utils/models/model_config.py index a2b0c90e59..16f6d21edb 100644 --- a/studio/backend/utils/models/model_config.py +++ b/studio/backend/utils/models/model_config.py @@ -32,6 +32,7 @@ import threading import yaml +from utils.native_path_leases import child_env_without_native_path_secret from utils.subprocess_compat import ( windows_hidden_subprocess_kwargs as _windows_hidden_subprocess_kwargs, ) @@ -583,6 +584,7 @@ def _is_vision_model_subprocess( capture_output = True, text = True, timeout = 60, + env = child_env_without_native_path_secret(), **_windows_hidden_subprocess_kwargs(), ) @@ -1227,9 +1229,11 @@ def _resolve_gguf_dir(p: Path) -> Optional[Path]: return p if p.is_file() and p.suffix.lower() == ".gguf": parent = p.parent - if (parent / "config.json").exists() or ( - parent / "adapter_config.json" - ).exists(): + if ( + (parent / "config.json").exists() + or (parent / "adapter_config.json").exists() + or (parent / "export_metadata.json").exists() + ): return parent return None diff --git a/studio/backend/utils/native_path_leases.py b/studio/backend/utils/native_path_leases.py new file mode 100644 index 0000000000..a69dfab532 --- /dev/null +++ b/studio/backend/utils/native_path_leases.py @@ -0,0 +1,406 @@ +# SPDX-License-Identifier: AGPL-3.0-only +# Copyright 2026-present the Unsloth AI Inc. team. All rights reserved. See /studio/LICENSE.AGPL-3.0 + +"""Verification for Tauri native path signed grants. + +Rust signs compact ``base64url(payload_json).base64url(hmac)`` grants. The +frontend can see and forward the grant, but cannot change it without breaking +the HMAC. The backend verifies the original payload segment bytes, then +re-stats the path before any native read. +""" + +from __future__ import annotations + +import base64 +import binascii +import hashlib +import hmac +import json +import os +import stat as _stat_module +import threading +import time +from contextlib import contextmanager +from dataclasses import dataclass +from pathlib import Path +from typing import Any, Callable, Iterable, Iterator, Mapping + +LEASE_SECRET_ENV = "UNSLOTH_STUDIO_NATIVE_PATH_LEASE_SECRET" +_MAX_NATIVE_PATH_REDACTIONS = 100 +_MAX_NATIVE_PATH_LABELS = 10_000 +_MIN_LEASE_SECRET_BYTES = 32 + +_REPLAY_LOCK = threading.Lock() +_USED_NONCES: dict[str, int] = {} +_REDACTION_LOCK = threading.Lock() +_NATIVE_PATH_REDACTIONS: list[str] = [] +_NATIVE_PATH_LABELS: dict[str, str] = {} +_NATIVE_PATH_ENV_LOCK = threading.Lock() +_SECRET_INIT_LOCK = threading.Lock() +_CACHED_LEASE_SECRET: bytes | None = None +_SCRUB_REFCOUNT = 0 +_SCRUB_SAVED_SECRET: str | None = None + + +class NativePathLeaseError(ValueError): + """Raised when a native path grant is missing, invalid, or unsafe.""" + + +@dataclass(frozen = True) +class NativePathGrant: + operation: str + canonical_path: Path + path_kind: str + path_type: str + source_kind: str + token_id_hash: str + display_label: str + expires_at_ms: int + size_bytes: int | None + modified_ms: int | None + + +def native_path_leases_supported() -> bool: + try: + _decode_secret() + except NativePathLeaseError: + return False + return True + + +def child_env_without_native_path_secret( + env: Mapping[str, str] | None = None, +) -> dict[str, str]: + """Return a child-process env with the native path lease secret removed.""" + + if env is None: + with _NATIVE_PATH_ENV_LOCK: + cleaned = dict(os.environ) + else: + cleaned = dict(env) + cleaned.pop(LEASE_SECRET_ENV, None) + return cleaned + + +def run_without_native_path_secret( + target: Callable[..., Any], + *args: Any, + **kwargs: Any, +) -> Any: + """Run a multiprocessing child target without the native path lease secret.""" + + global _CACHED_LEASE_SECRET, _SCRUB_SAVED_SECRET + os.environ.pop(LEASE_SECRET_ENV, None) + _CACHED_LEASE_SECRET = None + _SCRUB_SAVED_SECRET = None + return target(*args, **kwargs) + + +@contextmanager +def native_path_secret_removed_for_child_start() -> Iterator[None]: + global _SCRUB_REFCOUNT, _SCRUB_SAVED_SECRET, _CACHED_LEASE_SECRET + with _NATIVE_PATH_ENV_LOCK: + if _SCRUB_REFCOUNT == 0: + _SCRUB_SAVED_SECRET = os.environ.pop(LEASE_SECRET_ENV, None) + _CACHED_LEASE_SECRET = None + _SCRUB_REFCOUNT += 1 + try: + yield + finally: + with _NATIVE_PATH_ENV_LOCK: + _SCRUB_REFCOUNT -= 1 + if _SCRUB_REFCOUNT == 0 and _SCRUB_SAVED_SECRET is not None: + os.environ[LEASE_SECRET_ENV] = _SCRUB_SAVED_SECRET + _SCRUB_SAVED_SECRET = None + + +def verify_native_path_lease( + lease: str | None, + *, + operation: str, + expected_kind: str | None = None, + expected_path_type: str | None = None, + allowed_suffixes: Iterable[str] | None = None, +) -> NativePathGrant: + if not lease: + raise NativePathLeaseError("Native path grant is required.") + + secret = _decode_secret() + payload_b64, signature_b64 = _split_lease(lease) + expected_signature = hmac.new( + secret, + payload_b64.encode("ascii"), + hashlib.sha256, + ).digest() + supplied_signature = _b64decode(signature_b64) + if not hmac.compare_digest(expected_signature, supplied_signature): + raise NativePathLeaseError("Native path grant signature is invalid.") + + payload = _decode_payload(payload_b64) + _validate_payload(payload, operation = operation, expected_kind = expected_kind) + + path = Path(str(payload["canonical_path"])) + _reject_network_or_device_path(path) + try: + signed_lstat = os.lstat(path) + except OSError as exc: + raise NativePathLeaseError("Native path is no longer accessible.") from exc + if _stat_module.S_ISLNK(signed_lstat.st_mode): + raise NativePathLeaseError("Native path is no longer a regular file.") + try: + resolved = path.resolve(strict = True) + except OSError as exc: + raise NativePathLeaseError("Native path is no longer accessible.") from exc + _reject_network_or_device_path(resolved) + if not _same_native_path(resolved, path): + raise NativePathLeaseError( + "Native path grant no longer resolves to the selected path." + ) + + grant = NativePathGrant( + operation = str(payload["operation"]), + canonical_path = resolved, + path_kind = str(payload["path_kind"]), + path_type = str(payload["path_type"]), + source_kind = str(payload["source_kind"]), + token_id_hash = str(payload["token_id_hash"]), + display_label = str(payload.get("display_label") or resolved.name), + expires_at_ms = _required_int(payload, "expires_at_ms"), + size_bytes = _optional_int(payload.get("size_bytes")), + modified_ms = _optional_int(payload.get("modified_ms")), + ) + + if expected_path_type and grant.path_type != expected_path_type: + raise NativePathLeaseError("Native path grant has the wrong path type.") + suffixes = tuple(s.lower() for s in (allowed_suffixes or ())) + if suffixes and resolved.suffix.lower() not in suffixes: + raise NativePathLeaseError("Native path grant has an unsupported file type.") + + _validate_current_stat(grant) + _consume_nonce(str(payload["nonce"]), grant.expires_at_ms) + _remember_native_path_for_redaction(str(resolved), grant.display_label) + return grant + + +def display_label_for_native_path(value: str | None) -> str | None: + if not value: + return value + with _REDACTION_LOCK: + return _NATIVE_PATH_LABELS.get(value, value) + + +def is_registered_native_path_label(path_value: str | None, label: str | None) -> bool: + if not path_value or not label: + return False + with _REDACTION_LOCK: + return _NATIVE_PATH_LABELS.get(path_value) == label + + +def redact_native_paths(value: str) -> str: + with _REDACTION_LOCK: + paths = sorted(_NATIVE_PATH_REDACTIONS, key = len, reverse = True) + redacted = value + for path in paths: + for variant in {path, path.replace("/", "\\"), path.replace("\\", "/")}: + if variant: + redacted = redacted.replace(variant, "") + return redacted + + +def _decode_secret() -> bytes: + global _CACHED_LEASE_SECRET + if _CACHED_LEASE_SECRET is not None: + return _CACHED_LEASE_SECRET + with _SECRET_INIT_LOCK: + if _CACHED_LEASE_SECRET is not None: + return _CACHED_LEASE_SECRET + with _NATIVE_PATH_ENV_LOCK: + encoded = os.environ.get(LEASE_SECRET_ENV) + if encoded is None and _SCRUB_SAVED_SECRET is not None: + encoded = _SCRUB_SAVED_SECRET + if not encoded: + raise NativePathLeaseError( + "Native path grants require the managed desktop backend." + ) + try: + secret = _b64decode(encoded) + except Exception as exc: + raise NativePathLeaseError("Native path grant secret is invalid.") from exc + if len(secret) < _MIN_LEASE_SECRET_BYTES: + raise NativePathLeaseError("Native path grant secret is invalid.") + _CACHED_LEASE_SECRET = secret + return secret + + +def _split_lease(lease: str) -> tuple[str, str]: + if not isinstance(lease, str): + raise NativePathLeaseError("Native path grant has an invalid format.") + try: + lease.encode("ascii") + except UnicodeEncodeError as exc: + raise NativePathLeaseError("Native path grant has an invalid format.") from exc + parts = lease.split(".") + if len(parts) != 2 or not parts[0] or not parts[1]: + raise NativePathLeaseError("Native path grant has an invalid format.") + return parts[0], parts[1] + + +def _decode_payload(payload_b64: str) -> dict[str, Any]: + try: + payload = json.loads(_b64decode(payload_b64).decode("utf-8")) + except Exception as exc: + raise NativePathLeaseError("Native path grant payload is invalid.") from exc + if not isinstance(payload, dict): + raise NativePathLeaseError("Native path grant payload is invalid.") + return payload + + +def _validate_payload( + payload: dict[str, Any], *, operation: str, expected_kind: str | None +) -> None: + required = ( + "version", + "operation", + "canonical_path", + "path_kind", + "path_type", + "source_kind", + "token_id_hash", + "issued_at_ms", + "expires_at_ms", + "nonce", + ) + missing = [key for key in required if key not in payload] + if missing: + raise NativePathLeaseError( + "Native path grant payload is missing required fields." + ) + if _required_int(payload, "version") != 1: + raise NativePathLeaseError("Native path grant version is unsupported.") + if payload["operation"] != operation: + raise NativePathLeaseError("Native path grant operation is invalid.") + if expected_kind and payload["path_kind"] != expected_kind: + raise NativePathLeaseError("Native path grant kind is invalid.") + now_ms = int(time.time() * 1000) + issued_at_ms = _required_int(payload, "issued_at_ms") + expires_at_ms = _required_int(payload, "expires_at_ms") + if issued_at_ms >= expires_at_ms: + raise NativePathLeaseError("Native path grant timestamps are inconsistent.") + if expires_at_ms <= now_ms: + raise NativePathLeaseError("Native path grant has expired.") + if issued_at_ms > now_ms + 30_000: + raise NativePathLeaseError("Native path grant issue time is invalid.") + for key in ("canonical_path", "nonce", "token_id_hash", "display_label"): + raw = payload.get(key) + if raw is None: + continue + if "\x00" in str(raw): + raise NativePathLeaseError("Native path grant contains invalid characters.") + + +def _validate_current_stat(grant: NativePathGrant) -> None: + try: + st = os.lstat(grant.canonical_path) + except OSError as exc: + raise NativePathLeaseError("Native path is no longer accessible.") from exc + if _stat_module.S_ISLNK(st.st_mode): + raise NativePathLeaseError("Native path is no longer a regular file.") + if grant.path_type == "file": + if not _stat_module.S_ISREG(st.st_mode): + raise NativePathLeaseError("Native path is no longer a regular file.") + elif grant.path_type == "directory": + if not _stat_module.S_ISDIR(st.st_mode): + raise NativePathLeaseError("Native path is no longer a directory.") + else: + raise NativePathLeaseError("Native path grant has an unsupported path type.") + + if grant.size_bytes is not None and st.st_size != grant.size_bytes: + raise NativePathLeaseError("Native path changed after it was selected.") + current_modified_ms = int(st.st_mtime_ns // 1_000_000) + if grant.modified_ms is not None and current_modified_ms != grant.modified_ms: + raise NativePathLeaseError("Native path changed after it was selected.") + + +def _consume_nonce(nonce: str, expires_at_ms: int) -> None: + now_ms = int(time.time() * 1000) + with _REPLAY_LOCK: + for key, expiry in list(_USED_NONCES.items()): + if expiry <= now_ms: + _USED_NONCES.pop(key, None) + if nonce in _USED_NONCES: + raise NativePathLeaseError("Native path grant was already used.") + _USED_NONCES[nonce] = expires_at_ms + + +def _remember_native_path_for_redaction(path: str, display_label: str) -> None: + with _REDACTION_LOCK: + _NATIVE_PATH_LABELS[path] = display_label + if len(_NATIVE_PATH_LABELS) > _MAX_NATIVE_PATH_LABELS: + excess = len(_NATIVE_PATH_LABELS) - _MAX_NATIVE_PATH_LABELS + for stale_path in list(_NATIVE_PATH_LABELS.keys())[:excess]: + _NATIVE_PATH_LABELS.pop(stale_path, None) + if path in _NATIVE_PATH_REDACTIONS: + return + _NATIVE_PATH_REDACTIONS.append(path) + del _NATIVE_PATH_REDACTIONS[:-_MAX_NATIVE_PATH_REDACTIONS] + + +def _reject_network_or_device_path(path: Path) -> None: + text = str(path) + if os.name == "nt": + normalized = text.replace("/", "\\").lower() + if normalized.startswith("\\\\?\\"): + rest = normalized[4:] + is_local_drive = len(rest) >= 3 and rest[0].isalpha() and rest[1:3] == ":\\" + if not is_local_drive: + raise NativePathLeaseError( + "Network paths are not supported for native grants." + ) + elif normalized.startswith("\\\\"): + raise NativePathLeaseError( + "Network paths are not supported for native grants." + ) + if os.name != "nt": + for root in ("/dev", "/proc", "/sys"): + if path.is_relative_to(root): + raise NativePathLeaseError( + "Device and virtual filesystem paths are not supported." + ) + if "\x00" in text: + raise NativePathLeaseError("Native path contains invalid characters.") + + +def _b64decode(value: str) -> bytes: + try: + padding = "=" * (-len(value) % 4) + return base64.urlsafe_b64decode((value + padding).encode("ascii")) + except (UnicodeEncodeError, binascii.Error, ValueError) as exc: + raise NativePathLeaseError("Native path grant has an invalid format.") from exc + + +def _same_native_path(resolved: Path, signed: Path) -> bool: + try: + return resolved.samefile(signed) + except OSError: + return os.path.normcase(str(resolved)) == os.path.normcase(str(signed)) + + +def _optional_int(value: Any) -> int | None: + if value is None: + return None + try: + return int(value) + except (TypeError, ValueError) as exc: + raise NativePathLeaseError("Native path grant payload is invalid.") from exc + + +def _required_int(payload: dict[str, Any], key: str) -> int: + raw = payload.get(key) + if raw is None: + raise NativePathLeaseError( + "Native path grant payload is missing required fields." + ) + try: + return int(raw) + except (TypeError, ValueError) as exc: + raise NativePathLeaseError("Native path grant payload is invalid.") from exc diff --git a/studio/backend/utils/transformers_version.py b/studio/backend/utils/transformers_version.py index f36bdcd6e8..17af40f663 100644 --- a/studio/backend/utils/transformers_version.py +++ b/studio/backend/utils/transformers_version.py @@ -36,6 +36,7 @@ import subprocess import sys from pathlib import Path +from utils.native_path_leases import child_env_without_native_path_secret from utils.subprocess_compat import ( windows_hidden_subprocess_kwargs as _windows_hidden_subprocess_kwargs, ) @@ -63,6 +64,7 @@ TRANSFORMERS_5_MODEL_SUBSTRINGS: tuple[str, ...] = ( TRANSFORMERS_550_MODEL_SUBSTRINGS: tuple[str, ...] = ( "gemma-4", # Gemma-4 (E2B-it, E4B-it, 31B-it, 26B-A4B-it) "gemma4", # Gemma-4 alternate naming + "qwen3.6", ) # Architecture classes / model_type values that require transformers 5.5.0. @@ -503,6 +505,7 @@ def _install_to_dir(pkg: str, target_dir: str) -> bool: stdout = subprocess.PIPE, stderr = subprocess.STDOUT, text = True, + env = child_env_without_native_path_secret(), **_windows_hidden_subprocess_kwargs(), ) if result.returncode == 0: @@ -525,6 +528,7 @@ def _install_to_dir(pkg: str, target_dir: str) -> bool: stdout = subprocess.PIPE, stderr = subprocess.STDOUT, text = True, + env = child_env_without_native_path_secret(), **_windows_hidden_subprocess_kwargs(), ) if result.returncode != 0: diff --git a/studio/backend/utils/wheel_utils.py b/studio/backend/utils/wheel_utils.py index 00240f1e69..3ed9bda827 100644 --- a/studio/backend/utils/wheel_utils.py +++ b/studio/backend/utils/wheel_utils.py @@ -13,6 +13,8 @@ import urllib.error import urllib.request from typing import Callable +from utils.native_path_leases import child_env_without_native_path_secret + _logger = logging.getLogger(__name__) FLASH_ATTN_RELEASE_BASE_URL = ( @@ -59,6 +61,7 @@ def probe_torch_wheel_env(*, timeout: int | None = None) -> dict[str, str] | Non stderr = subprocess.PIPE, text = True, timeout = timeout, + env = child_env_without_native_path_secret(), ) except subprocess.TimeoutExpired: return None @@ -142,6 +145,7 @@ def install_wheel( stdout = subprocess.PIPE, stderr = subprocess.STDOUT, text = True, + env = child_env_without_native_path_secret(), ) attempts.append(("uv", result)) if result.returncode == 0: @@ -153,6 +157,7 @@ def install_wheel( stdout = subprocess.PIPE, stderr = subprocess.STDOUT, text = True, + env = child_env_without_native_path_secret(), ) attempts.append(("pip", result)) return attempts diff --git a/studio/frontend/package.json b/studio/frontend/package.json index c5cb949ccd..6b24440964 100644 --- a/studio/frontend/package.json +++ b/studio/frontend/package.json @@ -16,6 +16,7 @@ "biome:fix": "biome check . --write" }, "dependencies": { + "@assistant-ui/core": "0.1.17", "@assistant-ui/react": "^0.12.19", "@assistant-ui/react-markdown": "^0.12.3", "@assistant-ui/react-streamdown": "^0.1.2", @@ -42,6 +43,8 @@ "@tanstack/react-router": "^1.159.10", "@tanstack/react-table": "^8.21.3", "@tauri-apps/api": "^2.10.1", + "@tauri-apps/plugin-clipboard-manager": "^2.3.2", + "@tauri-apps/plugin-notification": "^2.3.3", "@tauri-apps/plugin-opener": "^2.5.3", "@tauri-apps/plugin-process": "^2.3.1", "@tauri-apps/plugin-updater": "^2.10.1", diff --git a/studio/frontend/src/app/auth-guards.ts b/studio/frontend/src/app/auth-guards.ts index 509b8f61af..52230f0b6f 100644 --- a/studio/frontend/src/app/auth-guards.ts +++ b/studio/frontend/src/app/auth-guards.ts @@ -9,7 +9,6 @@ import { hasRefreshToken, mustChangePassword, refreshSession, - tauriAutoAuth, } from "@/features/auth"; async function hasActiveSession(): Promise { @@ -39,7 +38,7 @@ function authRedirect(to: "/login" | "/change-password"): never { export async function requireAuth(): Promise { if (isTauri) { - await tauriAutoAuth(); + // AppProvider owns backend startup + desktop auth; route guards run before it mounts. return; } @@ -59,7 +58,6 @@ export async function requireAuth(): Promise { export async function requireGuest(): Promise { if (isTauri) { - await tauriAutoAuth(); throw redirect({ to: "/chat" }); } if (!(await hasActiveSession())) return; @@ -68,7 +66,6 @@ export async function requireGuest(): Promise { export async function requirePasswordChangeFlow(): Promise { if (isTauri) { - await tauriAutoAuth(); throw redirect({ to: "/chat" }); } diff --git a/studio/frontend/src/app/provider.tsx b/studio/frontend/src/app/provider.tsx index b75998a169..62e78b809a 100644 --- a/studio/frontend/src/app/provider.tsx +++ b/studio/frontend/src/app/provider.tsx @@ -4,12 +4,19 @@ import { StartupScreen } from "@/components/tauri/startup-screen"; import { UpdateBanner } from "@/components/tauri/update-banner"; import { UpdateScreen } from "@/components/tauri/update-screen"; +import { + WindowTitlebar, + shouldUseCustomWindowTitlebar, +} from "@/components/tauri/window-titlebar"; import { Toaster } from "@/components/ui/sonner"; -import { useTauriBackend } from "@/hooks/use-tauri-backend"; +import { getTauriAuthFailure, tauriAutoAuth } from "@/features/auth"; +import { NativeIntentDrain } from "@/features/native-intents/native-intent-drain"; +import { useTauriBackend, type BackendStatus } from "@/hooks/use-tauri-backend"; import { useTauriUpdate } from "@/hooks/use-tauri-update"; import { isTauri } from "@/lib/api-base"; +import { useRouterState } from "@tanstack/react-router"; import { ThemeProvider } from "next-themes"; -import { useEffect, useRef, type ReactNode } from "react"; +import { useEffect, useRef, useState, type ReactNode } from "react"; interface AppProviderProps { children: ReactNode; @@ -19,67 +26,85 @@ interface AppProviderProps { // Tauri window helpers (only imported in Tauri mode) // --------------------------------------------------------------------------- -async function showWindow(): Promise { +type TauriWindowMode = "setup" | "app"; +type WindowLayoutGuard = () => boolean; + +async function showSetupWindow(isCurrent: WindowLayoutGuard): Promise { const { getCurrentWindow } = await import("@tauri-apps/api/window"); - await getCurrentWindow().show(); -} + if (!isCurrent()) return; -function easeOutQuart(t: number): number { - return 1 - (1 - t) ** 4; -} - -async function animateToGoldenRatio(abortRef: { current: boolean }): Promise { - const { getCurrentWindow, currentMonitor, LogicalSize } = await import("@tauri-apps/api/window"); const win = getCurrentWindow(); - - // Ensure window is visible before resizing + if (!isCurrent()) return; + await win.center(); + if (!isCurrent()) return; await win.show(); +} +async function applyAppWindowLayout(isCurrent: WindowLayoutGuard): Promise { + const { getCurrentWindow, currentMonitor, LogicalSize } = await import("@tauri-apps/api/window"); + if (!isCurrent()) return; + + const win = getCurrentWindow(); const monitor = await currentMonitor(); - if (!monitor) return; + if (!isCurrent()) return; - // Convert physical pixels to logical using scale factor - const scale = monitor.scaleFactor; - const screenW = monitor.size.width / scale; - const screenH = monitor.size.height / scale; + let finalW = 900; + let finalH = 600; - // Target: 75% of screen width, golden ratio height, capped at min 900x600 - const targetW = Math.max(900, Math.round(screenW * 0.75)); - const targetH = Math.max(600, Math.round(targetW / 1.618)); - // Don't exceed screen height - const finalH = Math.min(targetH, Math.round(screenH * 0.85)); - const finalW = targetW; + if (monitor) { + // Convert physical pixels to logical using scale factor + const scale = monitor.scaleFactor; + const screenW = monitor.size.width / scale; + const screenH = monitor.size.height / scale; - // Check reduced motion preference - const prefersReducedMotion = window.matchMedia("(prefers-reduced-motion: reduce)").matches; - - if (prefersReducedMotion) { - await win.setSize(new LogicalSize(finalW, finalH)); - } else { - // Read current size instead of hardcoding — stays correct if tauri.conf.json changes - const inner = await win.innerSize(); - const factor = await win.scaleFactor(); - const startW = Math.round(inner.width / factor); - const startH = Math.round(inner.height / factor); - const steps = 15; - const stepDuration = 23; // ~350ms total - - for (let i = 1; i <= steps; i++) { - if (abortRef.current) return; - const t = easeOutQuart(i / steps); - const w = Math.round(startW + (finalW - startW) * t); - const h = Math.round(startH + (finalH - startH) * t); - await win.setSize(new LogicalSize(w, h)); - await new Promise((r) => setTimeout(r, stepDuration)); - } + // Target: 75% of screen width, golden ratio height, capped at min 900x600 + finalW = Math.max(900, Math.round(screenW * 0.75)); + const targetH = Math.max(600, Math.round(finalW / 1.618)); + // Don't exceed screen height + finalH = Math.min(targetH, Math.round(screenH * 0.85)); } - if (abortRef.current) return; - - // Apply constraints and finalize - await win.setResizable(true); + // Apply constraints and finalize without animating through intermediate sizes + if (!isCurrent()) return; + await win.setSize(new LogicalSize(finalW, finalH)); + if (!isCurrent()) return; await win.setSizeConstraints({ minWidth: 900, minHeight: 600 }); + if (!isCurrent()) return; + await win.setResizable(true); + if (!isCurrent()) return; await win.center(); + if (!isCurrent()) return; + await win.show(); +} + +async function showWindowFallback(): Promise { + const { getCurrentWindow } = await import("@tauri-apps/api/window"); + const win = getCurrentWindow(); + await win.setResizable(true); + await win.show(); +} + +function getTauriWindowMode( + status: BackendStatus, + hasEnteredAppMode: boolean, +): TauriWindowMode | null { + switch (status) { + case "checking": + return null; + case "not-installed": + case "installing": + case "install-error": + case "needs-elevation": + case "repairing": + case "repair-error": + return "setup"; + case "starting": + case "running": + case "stopped": + return "app"; + case "error": + return hasEnteredAppMode ? "app" : "setup"; + } } // --------------------------------------------------------------------------- @@ -103,6 +128,7 @@ function TauriUpdateLayer({ isExternalServer }: { isExternalServer: boolean }) { error={update.error} onRetry={update.retryUpdate} onSkipRestart={update.skipAndRestart} + onCopyDiagnostics={update.copyDiagnostics} /> ); } @@ -112,62 +138,147 @@ function TauriUpdateLayer({ isExternalServer }: { isExternalServer: boolean }) { status={update.status} info={update.info} dismissed={update.dismissed} + lastFailure={update.lastFailure} isExternalServer={isExternalServer} onInstall={update.installUpdate} onDismiss={update.dismiss} + onCopyDiagnostics={update.copyDiagnostics} /> ); } +const HIDDEN_TITLEBAR_SIDEBAR_ROUTES = new Set([ + "/onboarding", + "/login", + "/change-password", + "/signup", +]); + function TauriWrapper({ children }: { children: ReactNode }) { + const pathname = useRouterState({ select: (s) => s.location.pathname }); const { status, logs, error, isExternalServer, currentStepIndex, progressDetail, elevationPackages, - startInstall, retry, retryInstall, approveElevation, + startInstall, retry, retryInstall, approveElevation, copyDiagnostics, } = useTauriBackend(); - const hasResized = useRef(false); - const abortRef = useRef(false); + const appliedWindowModeRef = useRef(null); + const hasEnteredAppModeRef = useRef(false); + const windowLayoutGenerationRef = useRef(0); + const [desktopAuthReady, setDesktopAuthReady] = useState(!isTauri); + const [desktopAuthRetry, setDesktopAuthRetry] = useState(0); - // Show the window once the frontend mounts (for pre-running states) useEffect(() => { - if (isTauri) void showWindow(); + if (!isTauri) return; + return () => { + windowLayoutGenerationRef.current += 1; + appliedWindowModeRef.current = null; + }; }, []); - // Animate resize when backend becomes ready + // Keep the Tauri window hidden during preflight, then show it centered in setup + // mode or apply the final app layout in one instant step. useEffect(() => { - if (status === "running" && !hasResized.current) { - hasResized.current = true; - abortRef.current = false; - animateToGoldenRatio(abortRef).catch(async () => { - // On failure, at minimum make the window resizable so user can fix manually - try { - const { getCurrentWindow } = await import("@tauri-apps/api/window"); - await getCurrentWindow().setResizable(true); - } catch { /* swallow — window may still be functional */ } - }); + if (!isTauri) return; + + const nextMode = getTauriWindowMode(status, hasEnteredAppModeRef.current); + if (!nextMode) { + appliedWindowModeRef.current = null; + windowLayoutGenerationRef.current += 1; + return; } - return () => { abortRef.current = true; }; + if (appliedWindowModeRef.current === nextMode) return; + + appliedWindowModeRef.current = nextMode; + if (nextMode === "app") hasEnteredAppModeRef.current = true; + + const layoutGeneration = windowLayoutGenerationRef.current + 1; + windowLayoutGenerationRef.current = layoutGeneration; + const isCurrent = () => windowLayoutGenerationRef.current === layoutGeneration; + const applyWindowMode = nextMode === "setup" ? showSetupWindow : applyAppWindowLayout; + applyWindowMode(isCurrent).catch(async () => { + if (!isCurrent()) return; + // On failure, at minimum make the window visible and resizable so user can fix manually. + try { + await showWindowFallback(); + } catch { /* swallow — window may still be functional */ } + }); }, [status]); - if (!isTauri) return <>{children}; - if (status === "running") return <>{children}; + useEffect(() => { + if (!isTauri) { + setDesktopAuthReady(true); + return; + } + if (status !== "running") { + setDesktopAuthReady(false); + setDesktopAuthRetry(0); + return; + } - return ( + let disposed = false; + setDesktopAuthReady(false); + tauriAutoAuth({ force: true }).then((authenticated) => { + if (disposed) return; + if (authenticated) { + setDesktopAuthReady(true); + return; + } + if (!getTauriAuthFailure()) { + window.setTimeout(() => { + if (!disposed) setDesktopAuthRetry((value) => value + 1); + }, 500); + } + }); + + return () => { disposed = true; }; + }, [status, desktopAuthRetry]); + + if (!isTauri) return <>{children}; + + const showApp = status === "running" && desktopAuthReady; + const startupStatus = status === "running" ? "starting" : status; + const startupProgressDetail = + status === "running" && !desktopAuthReady + ? "Signing in to desktop session..." + : progressDetail; + + const content = showApp ? ( + <> + + + {children} + + ) : ( ); + + if (!shouldUseCustomWindowTitlebar()) return content; + + const showSidebarSurface = + showApp && !HIDDEN_TITLEBAR_SIDEBAR_ROUTES.has(pathname); + + return ( +
+ +
+ {content} +
+
+ ); } export function AppProvider({ children }: AppProviderProps) { diff --git a/studio/frontend/src/app/routes/__root.tsx b/studio/frontend/src/app/routes/__root.tsx index 19b5763557..22d149473c 100644 --- a/studio/frontend/src/app/routes/__root.tsx +++ b/studio/frontend/src/app/routes/__root.tsx @@ -81,7 +81,7 @@ function RootLayout() { pinned={pinned} setPinned={setPinned} togglePinned={togglePinned} - className="!min-h-0 h-dvh overflow-hidden" + className="!min-h-0 h-[calc(100dvh-var(--studio-titlebar-height,0px))] overflow-hidden" > diff --git a/studio/frontend/src/components/app-sidebar.tsx b/studio/frontend/src/components/app-sidebar.tsx index 6d029a59bd..edcd5120eb 100644 --- a/studio/frontend/src/components/app-sidebar.tsx +++ b/studio/frontend/src/components/app-sidebar.tsx @@ -31,19 +31,18 @@ import { import { useAnimatedThemeToggle } from "@/components/ui/animated-theme-toggler"; import { cn } from "@/lib/utils"; import { - Book03Icon, ChefHatIcon, ColumnInsertIcon, CursorInfo02Icon, Delete02Icon, Download03Icon, GemIcon, - MessageSearch01Icon, + Globe02Icon, Search01Icon, - NewReleasesIcon, PowerIcon, PencilEdit02Icon, LayoutAlignLeftIcon, + HelpCircleIcon, Settings02Icon, ZapIcon, } from "@hugeicons/core-free-icons"; @@ -527,9 +526,9 @@ export function AppSidebar() { className="!size-8" /> -
+
{displayTitle} - Studio + Unsloth
@@ -547,6 +546,15 @@ export function AppSidebar() { Settings ⌘, + useSettingsDialogStore.getState().openDialog("api-keys")} + > + + API + + New + + } onSelect={(e) => { e.preventDefault(); toggleTheme(); }} @@ -571,47 +579,12 @@ export function AppSidebar() { - - - - - Learn More - - - - - - What's New - - - - - - Feedback - - - - + useSettingsDialogStore.getState().openDialog("about")} + > + + Help + setShutdownOpen(true)}> Shutdown diff --git a/studio/frontend/src/components/assistant-ui/model-selector.tsx b/studio/frontend/src/components/assistant-ui/model-selector.tsx index 3628252177..795bcb6d08 100644 --- a/studio/frontend/src/components/assistant-ui/model-selector.tsx +++ b/studio/frontend/src/components/assistant-ui/model-selector.tsx @@ -13,18 +13,25 @@ import { usePlatformStore } from "@/config/env"; import { cn } from "@/lib/utils"; import { ArrowDown01Icon, + FolderSearchIcon, Logout01Icon, } from "@hugeicons/core-free-icons"; import { HugeiconsIcon } from "@hugeicons/react"; import { useMemo, useState } from "react"; import type { + DeletedModelRef, LoraModelOption, ModelOption, ModelSelectorChangeMeta, } from "./model-selector/types"; import { HubModelPicker, LoraModelPicker } from "./model-selector/pickers"; -export type { LoraModelOption, ModelOption, ModelSelectorChangeMeta } from "./model-selector/types"; +export type { + DeletedModelRef, + LoraModelOption, + ModelOption, + ModelSelectorChangeMeta, +} from "./model-selector/types"; interface ModelSelectorProps { models: ModelOption[]; @@ -35,6 +42,9 @@ interface ModelSelectorProps { onValueChange?: (value: string, meta: ModelSelectorChangeMeta) => void; onEject?: () => void; onFoldersChange?: () => void; + onPickLocalModel?: () => void | Promise; + onModelsChange?: (deletedModel?: DeletedModelRef) => void; + deleteDisabled?: boolean; variant?: "outline" | "ghost" | "muted"; size?: "sm" | "default" | "lg"; className?: string; @@ -66,7 +76,7 @@ function ModelSelectorTrigger({ type="button" data-tour={dataTour} className={cn( - "flex items-center gap-2 transition-colors", + "flex min-w-0 items-center gap-2 transition-colors", variant === "outline" && "rounded-[8px] border border-border/60 hover:bg-[#ececec] dark:hover:bg-[#2e3035]", variant === "ghost" && "rounded-[8px] hover:bg-[#ececec] dark:hover:bg-[#2e3035]", @@ -80,17 +90,23 @@ function ModelSelectorTrigger({ {isLoaded && ( )} - - {currentModel?.name ?? "Select model"} + + + {currentModel?.name ?? "Select model"} + + {currentModel?.description && ( + + {currentModel.description} + + )} + + + - {currentModel?.description && ( - {currentModel.description} - )} - ); @@ -103,6 +119,9 @@ function ModelSelectorContent({ onSelect, onEject, onFoldersChange, + onPickLocalModel, + onModelsChange, + deleteDisabled, className, dataTour, }: { @@ -112,6 +131,9 @@ function ModelSelectorContent({ onSelect: (id: string, meta: ModelSelectorChangeMeta) => void; onEject?: () => void; onFoldersChange?: () => void; + onPickLocalModel?: () => void; + onModelsChange?: (deletedModel?: DeletedModelRef) => void; + deleteDisabled?: boolean; className?: string; dataTour?: string; }) { @@ -145,11 +167,26 @@ function ModelSelectorContent({ loraModels={loraModels} value={value} onSelect={onSelect} + onModelsChange={onModelsChange} + deleteDisabled={deleteDisabled} /> )} + {onPickLocalModel ? ( +
+ +
+ ) : null} {hasSelection && onEject ? (
+ + { + if (!nextOpen && deleting) return; + setOpen(nextOpen); + }} + > + + + {title} + {description} + + + No + { + e.preventDefault(); + handleConfirm(); + }} + > + {deleting ? loadingLabel : "Yes"} + + + + + + ); +} diff --git a/studio/frontend/src/components/assistant-ui/model-selector/pickers.tsx b/studio/frontend/src/components/assistant-ui/model-selector/pickers.tsx index 30b8fdecb8..fae3c22caf 100644 --- a/studio/frontend/src/components/assistant-ui/model-selector/pickers.tsx +++ b/studio/frontend/src/components/assistant-ui/model-selector/pickers.tsx @@ -1,16 +1,6 @@ // SPDX-License-Identifier: AGPL-3.0-only // Copyright 2026-present the Unsloth AI Inc. team. All rights reserved. See /studio/LICENSE.AGPL-3.0 -import { - AlertDialog, - AlertDialogAction, - AlertDialogCancel, - AlertDialogContent, - AlertDialogDescription, - AlertDialogFooter, - AlertDialogHeader, - AlertDialogTitle, -} from "@/components/ui/alert-dialog"; import { Input } from "@/components/ui/input"; import { Spinner } from "@/components/ui/spinner"; import { @@ -23,6 +13,7 @@ import { type ScanFolderInfo, addScanFolder, deleteCachedModel, + deleteFineTunedModel, listCachedGguf, listCachedModels, listGgufVariants, @@ -50,7 +41,8 @@ import { checkVramFit, estimateLoadingVram } from "@/lib/vram"; import { Add01Icon, Cancel01Icon, Folder02Icon, Search01Icon } from "@hugeicons/core-free-icons"; import { HugeiconsIcon } from "@hugeicons/react"; import { FolderBrowser } from "./folder-browser"; -import { ChevronDownIcon, ChevronRightIcon, DownloadIcon, StarIcon, Trash2Icon } from "lucide-react"; +import { ModelDeleteAction } from "./model-delete-action"; +import { ChevronDownIcon, ChevronRightIcon, DownloadIcon, StarIcon } from "lucide-react"; import { type ReactNode, useCallback, @@ -60,6 +52,7 @@ import { } from "react"; import { toast } from "sonner"; import type { + DeletedModelRef, LoraModelOption, ModelOption, ModelSelectorChangeMeta, @@ -211,12 +204,22 @@ function GgufVariantExpander({ gpuGb, systemRamGb, onDeleteVariant, + sourceOverride, + deleteVariantTitle = "Delete cached model?", + renderDeleteVariantDescription, + getDeleteVariantSuccessMessage, + deleteDisabled = false, }: { repoId: string; onSelect: (id: string, meta: ModelSelectorChangeMeta) => void; gpuGb?: number; systemRamGb?: number; - onDeleteVariant?: (quant: string) => void; + onDeleteVariant?: (quant: string) => Promise | void; + sourceOverride?: ModelSelectorChangeMeta["source"]; + deleteVariantTitle?: string; + renderDeleteVariantDescription?: (quant: string) => ReactNode; + getDeleteVariantSuccessMessage?: (quant: string) => string; + deleteDisabled?: boolean; }) { const [variants, setVariants] = useState(null); const [defaultVariant, setDefaultVariant] = useState(null); @@ -259,14 +262,14 @@ function GgufVariantExpander({ const handleVariantClick = useCallback( (quant: string, downloaded?: boolean, sizeBytes?: number) => { onSelect(repoId, { - source: isLocalPath ? "local" : "hub", + source: sourceOverride ?? (isLocalPath ? "local" : "hub"), isLora: false, ggufVariant: quant, isDownloaded: isLocalPath ? true : downloaded, expectedBytes: sizeBytes, }); }, - [repoId, isLocalPath, onSelect], + [repoId, isLocalPath, onSelect, sourceOverride], ); // GGUF fit classification matching llama-server's _select_gpus logic: @@ -408,16 +411,29 @@ function GgufVariantExpander({ {v.downloaded && onDeleteVariant && ( - + + This will remove{" "} + + {repoId} ({v.quant}) + {" "} + from disk. You can re-download it later. + + ) + } + successMessage={ + getDeleteVariantSuccessMessage?.(v.quant) ?? + `Deleted ${repoId} ${v.quant}` + } + buttonClassName="p-1" + iconClassName="size-3" + disabled={deleteDisabled} + onConfirm={() => onDeleteVariant(v.quant)} + /> )}
); @@ -512,9 +528,6 @@ export function HubModelPicker({ // Track which GGUF repo is expanded for variant selection const [expandedGguf, setExpandedGguf] = useState(null); - // Delete confirmation dialog state - const [deleteTarget, setDeleteTarget] = useState(null); - const [deleting, setDeleting] = useState(false); const [downloadedCollapsed, setDownloadedCollapsed] = useState(false); const [customFoldersCollapsed, setCustomFoldersCollapsed] = useState(false); const [recommendedCollapsed, setRecommendedCollapsed] = useState(false); @@ -675,27 +688,6 @@ export function HubModelPicker({ .finally(check); }, [refreshLocalModelsList, refreshScanFolders]); - const handleDeleteConfirm = useCallback(async () => { - if (!deleteTarget) return; - setDeleting(true); - try { - // deleteTarget is "repo_id" or "repo_id::variant" - const sepIdx = deleteTarget.indexOf("::"); - const repoId = sepIdx >= 0 ? deleteTarget.slice(0, sepIdx) : deleteTarget; - const variant = sepIdx >= 0 ? deleteTarget.slice(sepIdx + 2) : undefined; - await deleteCachedModel(repoId, variant); - toast.success(`Deleted ${variant ? `${repoId} ${variant}` : repoId}`); - refreshCachedLists(); - } catch (err) { - toast.error( - err instanceof Error ? err.message : "Failed to delete model", - ); - } finally { - setDeleting(false); - setDeleteTarget(null); - } - }, [deleteTarget, refreshCachedLists]); - // Deduplicate: don't show downloaded models in the recommended list. // Compare case-insensitively since HF cache lowercases repo IDs. const downloadedSet = useMemo(() => { @@ -952,9 +944,10 @@ export function HubModelPicker({ systemRamGb={ gpu.available ? gpu.systemRamAvailableGb : undefined } - onDeleteVariant={(quant) => - setDeleteTarget(`${c.repo_id}::${quant}`) - } + onDeleteVariant={async (quant) => { + await deleteCachedModel(c.repo_id, quant); + refreshCachedLists(); + }} /> )}
@@ -977,16 +970,22 @@ export function HubModelPicker({ vramStatus={null} /> - + + This will remove{" "} + + {c.repo_id} + {" "} + from disk. You can re-download it later. + + } + successMessage={`Deleted ${c.repo_id}`} + onConfirm={() => deleteCachedModel(c.repo_id)} + onDeleted={refreshCachedLists} + /> ))} @@ -1409,40 +1408,6 @@ export function HubModelPicker({ - { - if (!open && !deleting) setDeleteTarget(null); - }} - > - - - Delete cached model? - - This will remove{" "} - - {deleteTarget?.includes("::") - ? `${deleteTarget.split("::")[0]} (${deleteTarget.split("::")[1]})` - : deleteTarget} - {" "} - from disk. You can re-download it later. - - - - No - { - e.preventDefault(); - handleDeleteConfirm(); - }} - > - {deleting ? "Deleting..." : "Yes"} - - - - ); } @@ -1451,10 +1416,14 @@ export function LoraModelPicker({ loraModels, value, onSelect, + onModelsChange, + deleteDisabled = false, }: { loraModels: LoraModelOption[]; value?: string; onSelect: (id: string, meta: ModelSelectorChangeMeta) => void; + onModelsChange?: (deletedModel?: DeletedModelRef) => void; + deleteDisabled?: boolean; }) { const [query, setQuery] = useState(""); const [expandedGguf, setExpandedGguf] = useState(null); @@ -1541,6 +1510,8 @@ export function LoraModelPicker({ const isExported = adapter.source === "exported"; const isMerged = adapter.exportType === "merged"; const isGguf = adapter.exportType === "gguf"; + const isExportedGguf = isExported && isGguf; + const canDelete = (isTraining || isExported) && !isExportedGguf; const isTrainingFull = isTraining && isMerged; const isLocalGgufDir = isLocal && @@ -1569,38 +1540,69 @@ export function LoraModelPicker({ : tag; return (
- { - if (isLocalGgufDir) { - setExpandedGguf((prev) => - prev === adapter.id ? null : adapter.id, - ); - } else { - onSelect(adapter.id, { - source: isLocal - ? "local" - : isExported - ? "exported" - : "lora", - isLora: !isLocal && !isMerged && !isGguf, - isDownloaded: true, - }); - } - }} - tooltipText={ - <> - - {adapter.name} - - - {adapter.id} - - - } - /> +
+
+ { + if (isLocalGgufDir || isExportedGguf) { + setExpandedGguf((prev) => + prev === adapter.id ? null : adapter.id, + ); + } else { + onSelect(adapter.id, { + source: isLocal + ? "local" + : isExported + ? "exported" + : "lora", + isLora: !isLocal && !isMerged && !isGguf, + isDownloaded: true, + }); + } + }} + tooltipText={ + <> + + {adapter.name} + + + {adapter.id} + + + } + /> +
+ {canDelete && ( + + This will remove{" "} + + {adapter.name} + {" "} + from disk. This cannot be undone. + + } + successMessage={`Deleted ${adapter.name}`} + disabled={deleteDisabled} + onConfirm={() => + deleteFineTunedModel({ + modelPath: adapter.id, + source: isExported ? "exported" : "training", + exportType: adapter.exportType, + }) + } + onDeleted={() => + onModelsChange?.({ id: adapter.id }) + } + /> + )} +
{expandedGguf === adapter.id && ( ( + <> + This will remove{" "} + + {adapter.name} ({quant}) + {" "} + from disk. This cannot be undone. + + )} + getDeleteVariantSuccessMessage={(quant) => + `Deleted ${adapter.name} ${quant}` + } + deleteDisabled={deleteDisabled} + onDeleteVariant={ + isExportedGguf + ? async (quant) => { + await deleteFineTunedModel({ + modelPath: adapter.id, + source: "exported", + exportType: "gguf", + ggufVariant: quant, + }); + onModelsChange?.({ + id: adapter.id, + ggufVariant: quant, + }); + } + : undefined + } /> )}
@@ -1619,6 +1652,7 @@ export function LoraModelPicker({ )} + ); } diff --git a/studio/frontend/src/components/assistant-ui/model-selector/types.ts b/studio/frontend/src/components/assistant-ui/model-selector/types.ts index 215dd2b38e..3da75b4d4e 100644 --- a/studio/frontend/src/components/assistant-ui/model-selector/types.ts +++ b/studio/frontend/src/components/assistant-ui/model-selector/types.ts @@ -25,3 +25,8 @@ export interface ModelSelectorChangeMeta { isDownloaded?: boolean; expectedBytes?: number; } + +export interface DeletedModelRef { + id: string; + ggufVariant?: string; +} diff --git a/studio/frontend/src/components/assistant-ui/thread.tsx b/studio/frontend/src/components/assistant-ui/thread.tsx index 3f67b11b51..0d6cd3bbf9 100644 --- a/studio/frontend/src/components/assistant-ui/thread.tsx +++ b/studio/frontend/src/components/assistant-ui/thread.tsx @@ -33,6 +33,7 @@ import { import { sentAudioNames } from "@/features/chat/api/chat-adapter"; import { useChatRuntimeStore } from "@/features/chat/stores/chat-runtime-store"; import { applyQwenThinkingParams } from "@/features/chat/utils/qwen-params"; +import { isTauri } from "@/lib/api-base"; import { deleteThreadMessage } from "@/features/chat/utils/delete-thread-message"; import { AUDIO_ACCEPT, MAX_AUDIO_SIZE, fileToBase64 } from "@/lib/audio-utils"; import { copyToClipboard } from "@/lib/copy-to-clipboard"; @@ -298,27 +299,41 @@ const Composer: FC<{ disabled?: boolean }> = ({ disabled }) => { [disabled], ); + const composerContent = ( + <> + + + + + + + ); + return ( - - - - - - - + {isTauri ? ( + // Phase 1 native model drops own Tauri local-path drops. Restore browser + // attachment drops in Tauri when Phase 1d adds attachment-token bridging. +
+ {composerContent} +
+ ) : ( + + {composerContent} + + )}
); }; @@ -485,7 +500,7 @@ const PreserveThinkingToggle: FC = () => { : "bg-muted text-muted-foreground hover:bg-muted-foreground/15", )} aria-label={ - preserveThinking ? "Disable preserve thinking" : "Enable preserve thinking" + preserveThinking ? "Disable preserve think" : "Enable preserve think" } > {preserveThinking && !disabled ? ( @@ -493,7 +508,7 @@ const PreserveThinkingToggle: FC = () => { ) : ( )} - Preserve Thinking + Preserve Think ); }; diff --git a/studio/frontend/src/components/tauri/startup-screen.tsx b/studio/frontend/src/components/tauri/startup-screen.tsx index af5f241416..1eb3f81d10 100644 --- a/studio/frontend/src/components/tauri/startup-screen.tsx +++ b/studio/frontend/src/components/tauri/startup-screen.tsx @@ -3,7 +3,9 @@ import { ShimmerButton } from "@/components/ui/shimmer-button"; import type { BackendStatus } from "@/hooks/use-tauri-backend"; +import type { CopySupportDiagnosticsResult } from "@/lib/tauri-diagnostics"; import { AnimatePresence, motion } from "motion/react"; +import { useState } from "react"; interface StartupScreenProps { status: BackendStatus; @@ -17,6 +19,63 @@ interface StartupScreenProps { onRetryInstall: () => void; onApproveElevation: () => void; onStartServer: () => void; + onCopyDiagnostics: () => Promise; +} + +function DiagnosticsCopyActions({ + onCopyDiagnostics, + children, +}: { + onCopyDiagnostics: () => Promise; + children: React.ReactNode; +}) { + const [copying, setCopying] = useState(false); + const [manualReport, setManualReport] = useState(null); + const [manualMessage, setManualMessage] = useState(null); + + async function handleCopyDiagnostics() { + setCopying(true); + try { + const result = await onCopyDiagnostics(); + if (result.ok) { + setManualReport(null); + setManualMessage(null); + } else { + setManualReport(result.report); + setManualMessage(result.error ?? "Clipboard copy failed. Select and copy the diagnostics below."); + } + } catch (error) { + setManualReport(null); + setManualMessage(`Diagnostics copy failed: ${String(error)}`); + } finally { + setCopying(false); + } + } + + return ( +
+
+ void handleCopyDiagnostics()} + > + {copying ? "Copying..." : "Copy Diagnostics"} + + {children} +
+ {manualMessage && ( +

{manualMessage}

+ )} + {manualReport && ( +