Merge branch 'main' into feature/chat-api
This commit is contained in:
commit
882de198b6
48 changed files with 21909 additions and 666 deletions
1
.gitignore
vendored
1
.gitignore
vendored
|
|
@ -228,3 +228,4 @@ setup_leo.sh
|
|||
server.pid
|
||||
*.log
|
||||
package-lock.json
|
||||
llama.cpp/
|
||||
|
|
|
|||
400
install.ps1
400
install.ps1
|
|
@ -3,6 +3,11 @@
|
|||
# Local: Set-ExecutionPolicy -Scope Process -ExecutionPolicy Bypass; .\install.ps1 --local
|
||||
# NoTorch: .\install.ps1 --no-torch (skip PyTorch, GGUF-only mode)
|
||||
# Test: .\install.ps1 --package roland-sloth
|
||||
#
|
||||
# Env vars (priority: UNSLOTH_STUDIO_HOME > STUDIO_HOME > USERPROFILE-redirect > default):
|
||||
# UNSLOTH_STUDIO_HOME / STUDIO_HOME = path -> install under that path
|
||||
# (DataDir nests inside; user PATH not modified persistently).
|
||||
# Default ($USERPROFILE\.unsloth\studio) is preserved when no env var is set.
|
||||
|
||||
function Install-UnslothStudio {
|
||||
$ErrorActionPreference = "Stop"
|
||||
|
|
@ -126,7 +131,94 @@ function Install-UnslothStudio {
|
|||
}
|
||||
|
||||
$PythonVersion = "3.13"
|
||||
$StudioHome = Join-Path $env:USERPROFILE ".unsloth\studio"
|
||||
|
||||
# Resolve install destinations. Priority: UNSLOTH_STUDIO_HOME, then
|
||||
# STUDIO_HOME alias, then USERPROFILE-redirect, then default.
|
||||
# Reject whitespace-only values so " " is treated as unset (matches the
|
||||
# Python resolvers' .strip()), preventing install/runtime layout drift.
|
||||
$envOverrideVar = $null
|
||||
$envOverride = $null
|
||||
if (-not [string]::IsNullOrWhiteSpace($env:UNSLOTH_STUDIO_HOME)) {
|
||||
$envOverrideVar = "UNSLOTH_STUDIO_HOME"
|
||||
$envOverride = $env:UNSLOTH_STUDIO_HOME.Trim()
|
||||
} elseif (-not [string]::IsNullOrWhiteSpace($env:STUDIO_HOME)) {
|
||||
$envOverrideVar = "STUDIO_HOME"
|
||||
$envOverride = $env:STUDIO_HOME.Trim()
|
||||
}
|
||||
|
||||
# Custom Studio roots are not supported with --tauri (desktop app still
|
||||
# resolves %USERPROFILE%\.unsloth\studio). Pass through if override == legacy.
|
||||
if ($TauriMode -and $envOverride) {
|
||||
$_tauriOverride = $envOverride
|
||||
if ($_tauriOverride -eq "~" -or $_tauriOverride -like "~/*" -or $_tauriOverride -like "~\*") {
|
||||
$_tauriOverride = (Join-Path $env:USERPROFILE $_tauriOverride.Substring(1).TrimStart('/','\'))
|
||||
}
|
||||
try {
|
||||
$_tauriOverride = [System.IO.Path]::GetFullPath($_tauriOverride)
|
||||
} catch {}
|
||||
$_legacyTauriRoot = Join-Path $env:USERPROFILE ".unsloth\studio"
|
||||
try {
|
||||
$_legacyTauriRoot = [System.IO.Path]::GetFullPath($_legacyTauriRoot)
|
||||
} catch {}
|
||||
# Strip trailing separators so ".../studio\" matches ".../studio".
|
||||
$_trimSeps = @(
|
||||
[System.IO.Path]::DirectorySeparatorChar,
|
||||
[System.IO.Path]::AltDirectorySeparatorChar
|
||||
)
|
||||
$_tauriOverride = $_tauriOverride.TrimEnd($_trimSeps)
|
||||
$_legacyTauriRoot = $_legacyTauriRoot.TrimEnd($_trimSeps)
|
||||
if ($_tauriOverride -ne $_legacyTauriRoot) {
|
||||
Write-Host "ERROR: $envOverrideVar is not supported with --tauri." -ForegroundColor Red
|
||||
Write-Host " The desktop app still uses the legacy %USERPROFILE%\.unsloth\studio root." -ForegroundColor Red
|
||||
Write-Host " Run install.ps1 without --tauri for custom-root shell installs," -ForegroundColor Yellow
|
||||
Write-Host " or unset the env var for default desktop installs." -ForegroundColor Yellow
|
||||
throw "$envOverrideVar is not supported with --tauri."
|
||||
}
|
||||
}
|
||||
|
||||
$defaultProfile = $null
|
||||
try { $defaultProfile = [Environment]::GetFolderPath("UserProfile") } catch {}
|
||||
|
||||
# LOCALAPPDATA may be unset in service / CI contexts; Join-Path would abort
|
||||
# under ErrorActionPreference=Stop without this guard.
|
||||
$defaultDataDir = if ($env:LOCALAPPDATA -and -not [string]::IsNullOrWhiteSpace($env:LOCALAPPDATA)) {
|
||||
Join-Path $env:LOCALAPPDATA "Unsloth Studio"
|
||||
} else { $null }
|
||||
|
||||
if ($envOverride) {
|
||||
# Tilde expansion: env vars aren't subject to it when quoted on assignment.
|
||||
if ($envOverride -eq "~" -or $envOverride -like "~/*" -or $envOverride -like "~\*") {
|
||||
$envOverride = (Join-Path $env:USERPROFILE $envOverride.Substring(1).TrimStart('/','\'))
|
||||
}
|
||||
try {
|
||||
# .NET API: New-Item -Path treats brackets as wildcards and has no
|
||||
# -LiteralPath in PS 5.1, so a root like C:\studio[abc] would fail.
|
||||
[System.IO.Directory]::CreateDirectory($envOverride) | Out-Null
|
||||
$StudioHome = (Resolve-Path -LiteralPath $envOverride).Path
|
||||
} catch {
|
||||
Write-Host "ERROR: $envOverrideVar=$envOverride cannot be created or accessed." -ForegroundColor Red
|
||||
throw "$envOverrideVar=$envOverride cannot be created or accessed."
|
||||
}
|
||||
$probe = Join-Path $StudioHome (".unsloth-write-probe-" + [guid]::NewGuid())
|
||||
try {
|
||||
# WriteAllText: literal-path safe + closes handle so Remove-Item works.
|
||||
[System.IO.File]::WriteAllText($probe, "")
|
||||
Remove-Item -LiteralPath $probe -Force -ErrorAction SilentlyContinue
|
||||
} catch {
|
||||
Write-Host "ERROR: $envOverrideVar=$StudioHome is not writable." -ForegroundColor Red
|
||||
throw "$envOverrideVar=$StudioHome is not writable."
|
||||
}
|
||||
$StudioDataDir = Join-Path $StudioHome "share"
|
||||
$StudioRedirectMode = 'env'
|
||||
} elseif ($defaultProfile -and $env:USERPROFILE -and ($env:USERPROFILE -ne $defaultProfile)) {
|
||||
$StudioHome = Join-Path $env:USERPROFILE ".unsloth\studio"
|
||||
$StudioDataDir = $defaultDataDir
|
||||
$StudioRedirectMode = 'profile'
|
||||
} else {
|
||||
$StudioHome = Join-Path $env:USERPROFILE ".unsloth\studio"
|
||||
$StudioDataDir = $defaultDataDir
|
||||
$StudioRedirectMode = 'default'
|
||||
}
|
||||
$VenvDir = Join-Path $StudioHome "unsloth_studio"
|
||||
|
||||
$Rule = [string]::new([char]0x2500, 52)
|
||||
|
|
@ -378,24 +470,24 @@ function Install-UnslothStudio {
|
|||
[Parameter(Mandatory = $true)][string]$UnslothExePath
|
||||
)
|
||||
|
||||
if (-not (Test-Path $UnslothExePath)) {
|
||||
if (-not (Test-Path -LiteralPath $UnslothExePath)) {
|
||||
substep "cannot create shortcuts, unsloth.exe not found at $UnslothExePath" "Yellow"
|
||||
return
|
||||
}
|
||||
try {
|
||||
# Persist an absolute path in launcher scripts so shortcut working
|
||||
# directory changes do not break process startup.
|
||||
$UnslothExePath = (Resolve-Path $UnslothExePath).Path
|
||||
$UnslothExePath = (Resolve-Path -LiteralPath $UnslothExePath).Path
|
||||
# Escape for single-quoted embedding in generated launcher script.
|
||||
# This prevents runtime variable expansion for paths containing '$'.
|
||||
$SingleQuotedExePath = $UnslothExePath -replace "'", "''"
|
||||
|
||||
$localAppDataDir = $env:LOCALAPPDATA
|
||||
if (-not $localAppDataDir -or [string]::IsNullOrWhiteSpace($localAppDataDir)) {
|
||||
substep "LOCALAPPDATA path unavailable; skipped shortcut creation" "Yellow"
|
||||
# $StudioDataDir = LOCALAPPDATA\Unsloth Studio, or $StudioHome\share in env-mode.
|
||||
if (-not $StudioDataDir -or [string]::IsNullOrWhiteSpace($StudioDataDir)) {
|
||||
substep "DataDir path unavailable; skipped shortcut creation" "Yellow"
|
||||
return
|
||||
}
|
||||
$appDir = Join-Path $localAppDataDir "Unsloth Studio"
|
||||
$appDir = $StudioDataDir
|
||||
$launcherPs1 = Join-Path $appDir "launch-studio.ps1"
|
||||
$launcherVbs = Join-Path $appDir "launch-studio.vbs"
|
||||
$desktopDir = [Environment]::GetFolderPath("Desktop")
|
||||
|
|
@ -427,23 +519,89 @@ function Install-UnslothStudio {
|
|||
}
|
||||
$iconUrl = "https://raw.githubusercontent.com/unslothai/unsloth/main/studio/frontend/public/unsloth.ico"
|
||||
|
||||
if (-not (Test-Path $appDir)) {
|
||||
New-Item -ItemType Directory -Path $appDir -Force | Out-Null
|
||||
if (-not (Test-Path -LiteralPath $appDir)) {
|
||||
[System.IO.Directory]::CreateDirectory($appDir) | Out-Null
|
||||
}
|
||||
|
||||
# Same-install discriminator: per-install opaque id written once at
|
||||
# install time and read by both this launcher and the backend
|
||||
# (/api/health). Replaces the older sha256(resolved $StudioHome)
|
||||
# scheme to (a) avoid leaking the install path on -H 0.0.0.0
|
||||
# deployments and (b) sidestep launcher/backend canonicalization
|
||||
# drift (Resolve-Path vs Path.resolve() junction handling). Lives
|
||||
# at $StudioHome\share\ (not $appDir) so the backend can find it
|
||||
# via _STUDIO_ROOT_RESOLVED / "share" / "studio_install_id"
|
||||
# regardless of mode. 32 bytes of crypto random -> 64 hex chars.
|
||||
$_studioIdDir = Join-Path $StudioHome "share"
|
||||
if (-not (Test-Path -LiteralPath $_studioIdDir)) {
|
||||
[System.IO.Directory]::CreateDirectory($_studioIdDir) | Out-Null
|
||||
}
|
||||
$_studioIdFile = Join-Path $_studioIdDir "studio_install_id"
|
||||
$_studioRootId = ""
|
||||
if ((Test-Path -LiteralPath $_studioIdFile) -and `
|
||||
((Get-Item -LiteralPath $_studioIdFile).Length -gt 0)) {
|
||||
$_studioRootId = ([System.IO.File]::ReadAllText($_studioIdFile)).Trim()
|
||||
}
|
||||
if (-not $_studioRootId) {
|
||||
$_idBytes = New-Object byte[] 32
|
||||
[Security.Cryptography.RandomNumberGenerator]::Create().GetBytes($_idBytes)
|
||||
$_studioRootId = -join ($_idBytes | ForEach-Object { $_.ToString('x2') })
|
||||
# Atomic write: write to a temp sibling then rename, so a partial
|
||||
# install cannot leave a half-written id.
|
||||
$_idTmp = $_studioIdFile + ".$PID.tmp"
|
||||
[System.IO.File]::WriteAllText($_idTmp, $_studioRootId)
|
||||
Move-Item -LiteralPath $_idTmp -Destination $_studioIdFile -Force
|
||||
}
|
||||
|
||||
# Env-mode: persist UNSLOTH_STUDIO_HOME (and llama path) so fresh
|
||||
# shells don't need to re-export, and bake per-install $portFile /
|
||||
# $mutexName so concurrent custom-root launchers cannot serialize
|
||||
# through one global mutex on 8888..8908. Default installs get an
|
||||
# empty prefix to match pre-PR behavior.
|
||||
$studioHomeExport = if ($StudioRedirectMode -eq 'env') {
|
||||
# When override == legacy default, llama.cpp stays at
|
||||
# ~/.unsloth/llama.cpp (one shared build). Canonicalize the
|
||||
# legacy side so the comparison survives path normalization.
|
||||
$_legacyStudio = Join-Path $env:USERPROFILE ".unsloth\studio"
|
||||
if (Test-Path -LiteralPath $_legacyStudio -PathType Container) {
|
||||
$_legacyStudio = (Resolve-Path -LiteralPath $_legacyStudio).Path
|
||||
}
|
||||
$_llamaPath = if ($StudioHome -eq $_legacyStudio) {
|
||||
Join-Path $env:USERPROFILE ".unsloth\llama.cpp"
|
||||
} else {
|
||||
Join-Path $StudioHome "llama.cpp"
|
||||
}
|
||||
$_sq = $StudioHome -replace "'", "''"
|
||||
$_llama = $_llamaPath -replace "'", "''"
|
||||
$_appDirSq = $appDir -replace "'", "''"
|
||||
$_appBytes = [Text.Encoding]::UTF8.GetBytes($appDir)
|
||||
$_appHash = ([BitConverter]::ToString(
|
||||
[Security.Cryptography.SHA256]::Create().ComputeHash($_appBytes)
|
||||
) -replace '-', '').Substring(0, 16)
|
||||
# UNSLOTH_LLAMA_CPP_PATH is a pre-existing user override; only default if unset.
|
||||
"`$env:UNSLOTH_STUDIO_HOME = '$_sq'`nif (-not `$env:UNSLOTH_LLAMA_CPP_PATH) {`n `$env:UNSLOTH_LLAMA_CPP_PATH = '$_llama'`n}`n`$portFile = '$_appDirSq\studio.port'`n`$mutexName = 'Local\UnslothStudioLauncher-$_appHash'`n"
|
||||
} else {
|
||||
"`$portFile = `$null`n`$mutexName = 'Local\UnslothStudioLauncher'`n"
|
||||
}
|
||||
|
||||
$launcherContent = @"
|
||||
`$ErrorActionPreference = 'Stop'
|
||||
$studioHomeExport`$ErrorActionPreference = 'Stop'
|
||||
`$basePort = 8888
|
||||
`$maxPortOffset = 20
|
||||
`$timeoutSec = 60
|
||||
`$pollIntervalMs = 1000
|
||||
`$_ExpectedStudioRootId = '$_studioRootId'
|
||||
|
||||
function Test-StudioHealth {
|
||||
param([Parameter(Mandatory = `$true)][int]`$Port)
|
||||
try {
|
||||
`$url = "http://127.0.0.1:`$Port/api/health"
|
||||
`$resp = Invoke-RestMethod -Uri `$url -TimeoutSec 1 -Method Get
|
||||
return (`$resp -and `$resp.status -eq 'healthy' -and `$resp.service -eq 'Unsloth UI Backend')
|
||||
if (-not (`$resp -and `$resp.status -eq 'healthy' -and `$resp.service -eq 'Unsloth UI Backend')) { return `$false }
|
||||
# why: verify the backend belongs to THIS install via the install-time
|
||||
# hex digest; raw path is not leaked over /api/health.
|
||||
if (`$_ExpectedStudioRootId -and `$resp.studio_root_id -ne `$_ExpectedStudioRootId) { return `$false }
|
||||
return `$true
|
||||
} catch {
|
||||
return `$false
|
||||
}
|
||||
|
|
@ -469,6 +627,17 @@ function Get-CandidatePorts {
|
|||
}
|
||||
|
||||
function Find-HealthyStudioPort {
|
||||
if (`$portFile) {
|
||||
if (Test-Path -LiteralPath `$portFile) {
|
||||
`$cached = Get-Content -LiteralPath `$portFile -ErrorAction SilentlyContinue | Select-Object -First 1
|
||||
if (`$cached -match '^\d+`$') {
|
||||
`$cachedPort = [int]`$cached
|
||||
if (Test-StudioHealth -Port `$cachedPort) { return `$cachedPort }
|
||||
Remove-Item -LiteralPath `$portFile -Force -ErrorAction SilentlyContinue
|
||||
}
|
||||
}
|
||||
return `$null
|
||||
}
|
||||
foreach (`$candidate in (Get-CandidatePorts)) {
|
||||
if (Test-StudioHealth -Port `$candidate) {
|
||||
return `$candidate
|
||||
|
|
@ -522,7 +691,7 @@ if (`$existingPort) {
|
|||
exit 0
|
||||
}
|
||||
|
||||
`$launchMutex = [System.Threading.Mutex]::new(`$false, 'Local\UnslothStudioLauncher')
|
||||
`$launchMutex = [System.Threading.Mutex]::new(`$false, `$mutexName)
|
||||
`$haveMutex = `$false
|
||||
try {
|
||||
try {
|
||||
|
|
@ -552,7 +721,9 @@ try {
|
|||
} catch {}
|
||||
exit 1
|
||||
}
|
||||
`$studioCommand = '& "' + `$studioExe + '" studio -p ' + `$launchPort
|
||||
# Single-quote the path in the child -Command so `$` / backtick in custom
|
||||
# roots don't get reparsed; double any apostrophes so 'O''Brien' survives.
|
||||
`$studioCommand = "& '" + (`$studioExe -replace "'", "''") + "' studio -p " + `$launchPort
|
||||
`$launchArgs = @(
|
||||
'-NoExit',
|
||||
'-NoProfile',
|
||||
|
|
@ -576,9 +747,13 @@ try {
|
|||
`$browserOpened = `$false
|
||||
`$deadline = (Get-Date).AddSeconds(`$timeoutSec)
|
||||
while ((Get-Date) -lt `$deadline) {
|
||||
`$healthyPort = Find-HealthyStudioPort
|
||||
if (`$healthyPort) {
|
||||
Start-Process "http://localhost:`$healthyPort"
|
||||
if (Test-StudioHealth -Port `$launchPort) {
|
||||
if (`$portFile) {
|
||||
try {
|
||||
[System.IO.File]::WriteAllText(`$portFile, "`$launchPort`n")
|
||||
} catch {}
|
||||
}
|
||||
Start-Process "http://localhost:`$launchPort"
|
||||
`$browserOpened = `$true
|
||||
break
|
||||
}
|
||||
|
|
@ -613,19 +788,19 @@ cmd = "powershell -NoProfile -ExecutionPolicy Bypass -WindowStyle Hidden -File "
|
|||
shell.Run cmd, 0, False
|
||||
"@
|
||||
# WSH handles UTF-16LE reliably for .vbs files with non-ASCII paths.
|
||||
Set-Content -Path $launcherVbs -Value $vbsContent -Encoding Unicode -Force
|
||||
Set-Content -LiteralPath $launcherVbs -Value $vbsContent -Encoding Unicode -Force
|
||||
|
||||
# Prefer bundled icon from local clone/dev installs.
|
||||
# If not available, best-effort download from raw GitHub.
|
||||
# We only attach the icon if the resulting file has a valid ICO header.
|
||||
$hasValidIcon = $false
|
||||
if ($bundledIcon -and (Test-Path $bundledIcon)) {
|
||||
if ($bundledIcon -and (Test-Path -LiteralPath $bundledIcon)) {
|
||||
try {
|
||||
Copy-Item -Path $bundledIcon -Destination $iconPath -Force
|
||||
Copy-Item -LiteralPath $bundledIcon -Destination $iconPath -Force
|
||||
} catch {
|
||||
Write-Host "[DEBUG] Error copying bundled icon: $($_.Exception.Message)" -ForegroundColor DarkGray
|
||||
}
|
||||
} elseif (-not (Test-Path $iconPath)) {
|
||||
} elseif (-not (Test-Path -LiteralPath $iconPath)) {
|
||||
try {
|
||||
Invoke-WebRequest -Uri $iconUrl -OutFile $iconPath -UseBasicParsing
|
||||
} catch {
|
||||
|
|
@ -633,7 +808,7 @@ shell.Run cmd, 0, False
|
|||
}
|
||||
}
|
||||
|
||||
if (Test-Path $iconPath) {
|
||||
if (Test-Path -LiteralPath $iconPath) {
|
||||
try {
|
||||
$bytes = [System.IO.File]::ReadAllBytes($iconPath)
|
||||
if (
|
||||
|
|
@ -645,14 +820,21 @@ shell.Run cmd, 0, False
|
|||
) {
|
||||
$hasValidIcon = $true
|
||||
} else {
|
||||
Remove-Item $iconPath -Force -ErrorAction SilentlyContinue
|
||||
Remove-Item -LiteralPath $iconPath -Force -ErrorAction SilentlyContinue
|
||||
}
|
||||
} catch {
|
||||
Write-Host "[DEBUG] Error validating or removing icon: $($_.Exception.Message)" -ForegroundColor DarkGray
|
||||
Remove-Item $iconPath -Force -ErrorAction SilentlyContinue
|
||||
Remove-Item -LiteralPath $iconPath -Force -ErrorAction SilentlyContinue
|
||||
}
|
||||
}
|
||||
|
||||
# Env-mode: skip persistent Desktop / Start Menu .lnk shortcuts
|
||||
# that may point at a deleted workspace; launcher + icon stay.
|
||||
if ($StudioRedirectMode -eq 'env') {
|
||||
substep "wrote launcher at $launcherPs1 (persistent shortcuts skipped in env-override mode)"
|
||||
return
|
||||
}
|
||||
|
||||
$wscriptExe = Join-Path $env:SystemRoot "System32\wscript.exe"
|
||||
$shortcutArgs = "//B //Nologo `"$launcherVbs`""
|
||||
|
||||
|
|
@ -850,8 +1032,9 @@ shell.Run cmd, 0, False
|
|||
# Pass the resolved executable path to uv so it does not re-resolve
|
||||
# a version string back to a conda interpreter.
|
||||
Write-TauriLog "STEP" "Creating virtual environment"
|
||||
if (-not (Test-Path $StudioHome)) {
|
||||
New-Item -ItemType Directory -Path $StudioHome -Force | Out-Null
|
||||
if (-not (Test-Path -LiteralPath $StudioHome)) {
|
||||
# .NET API: New-Item -Path treats brackets as wildcards.
|
||||
[System.IO.Directory]::CreateDirectory($StudioHome) | Out-Null
|
||||
}
|
||||
|
||||
$VenvPython = Join-Path $VenvDir "Scripts\python.exe"
|
||||
|
|
@ -865,11 +1048,13 @@ shell.Run cmd, 0, False
|
|||
$stamp = Get-Date -Format "yyyyMMddHHmmss"
|
||||
$candidate = Join-Path $StudioHome "unsloth_studio.rollback.$stamp.$PID"
|
||||
$suffix = 0
|
||||
while (Test-Path $candidate) {
|
||||
# -LiteralPath: a custom $StudioHome may contain [ ] * ? which
|
||||
# plain Test-Path / Move-Item would interpret as wildcards.
|
||||
while (Test-Path -LiteralPath $candidate) {
|
||||
$suffix++
|
||||
$candidate = Join-Path $StudioHome "unsloth_studio.rollback.$stamp.$PID.$suffix"
|
||||
}
|
||||
Move-Item -Path $ExistingDir -Destination $candidate -ErrorAction Stop
|
||||
Move-Item -LiteralPath $ExistingDir -Destination $candidate -ErrorAction Stop
|
||||
$script:StudioVenvRollbackDir = $candidate
|
||||
$script:StudioVenvRollbackTarget = $ExistingDir
|
||||
$script:StudioVenvRollbackActive = $true
|
||||
|
|
@ -880,16 +1065,16 @@ shell.Run cmd, 0, False
|
|||
if (-not $script:StudioVenvRollbackActive) { return }
|
||||
$backup = $script:StudioVenvRollbackDir
|
||||
$target = $script:StudioVenvRollbackTarget
|
||||
if (-not $backup -or -not (Test-Path $backup)) {
|
||||
if (-not $backup -or -not (Test-Path -LiteralPath $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
|
||||
if (Test-Path -LiteralPath $target) {
|
||||
Remove-Item -LiteralPath $target -Recurse -Force -ErrorAction SilentlyContinue
|
||||
}
|
||||
Move-Item -Path $backup -Destination $target -Force -ErrorAction Stop
|
||||
Move-Item -LiteralPath $backup -Destination $target -Force -ErrorAction Stop
|
||||
substep "restored previous environment"
|
||||
$script:StudioVenvRollbackActive = $false
|
||||
$script:StudioVenvRollbackDir = $null
|
||||
|
|
@ -902,14 +1087,29 @@ shell.Run cmd, 0, False
|
|||
function Complete-StudioVenvRollback {
|
||||
if (-not $script:StudioVenvRollbackActive) { return }
|
||||
$backup = $script:StudioVenvRollbackDir
|
||||
if ($backup -and (Test-Path $backup)) {
|
||||
Remove-Item -Recurse -Force $backup -ErrorAction SilentlyContinue
|
||||
if ($backup -and (Test-Path -LiteralPath $backup)) {
|
||||
Remove-Item -LiteralPath $backup -Recurse -Force -ErrorAction SilentlyContinue
|
||||
}
|
||||
$script:StudioVenvRollbackActive = $false
|
||||
$script:StudioVenvRollbackDir = $null
|
||||
}
|
||||
|
||||
if (Test-Path $VenvPython) {
|
||||
if (Test-Path -LiteralPath $VenvPython) {
|
||||
# why: matching guard to the .venv branch below -- in env-mode
|
||||
# $StudioHome is a user-chosen workspace, so refuse to nuke an
|
||||
# existing $StudioHome\unsloth_studio that lacks Studio sentinels.
|
||||
# -PathType Leaf rejects a directory at the sentinel path. Accept the
|
||||
# in-VENV ownership marker so partial-install retries are not blocked.
|
||||
if (
|
||||
$StudioRedirectMode -eq 'env' -and
|
||||
-not (Test-Path -LiteralPath (Join-Path $VenvDir ".unsloth-studio-owned") -PathType Leaf) -and
|
||||
-not (Test-Path -LiteralPath (Join-Path $StudioHome "share\studio.conf") -PathType Leaf) -and
|
||||
-not (Test-Path -LiteralPath (Join-Path $StudioHome "bin\unsloth.exe") -PathType Leaf)
|
||||
) {
|
||||
Write-Host "[ERROR] $VenvDir already exists but does not look like an Unsloth Studio install." -ForegroundColor Red
|
||||
Write-Host " Move it aside or choose an empty UNSLOTH_STUDIO_HOME." -ForegroundColor Yellow
|
||||
throw "Refusing to delete non-Studio venv at $VenvDir"
|
||||
}
|
||||
# New layout already exists -- replace only after preserving rollback copy.
|
||||
substep "preserving existing environment for rollback..."
|
||||
try {
|
||||
|
|
@ -918,8 +1118,13 @@ shell.Run cmd, 0, False
|
|||
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
|
||||
} elseif (
|
||||
$StudioRedirectMode -ne 'env' `
|
||||
-and (Test-Path -LiteralPath (Join-Path $StudioHome ".venv\Scripts\python.exe"))
|
||||
) {
|
||||
# Old layout (~/.unsloth/studio/.venv) exists -- validate before migrating.
|
||||
# Skip in env-mode so we don't blow away an unrelated .venv at the
|
||||
# workspace root (e.g. user's existing project Python venv).
|
||||
$OldVenv = Join-Path $StudioHome ".venv"
|
||||
$OldPy = Join-Path $OldVenv "Scripts\python.exe"
|
||||
substep "found legacy Studio environment, validating..."
|
||||
|
|
@ -936,24 +1141,29 @@ shell.Run cmd, 0, False
|
|||
$ErrorActionPreference = $prevEAP2
|
||||
if ($legacyOk) {
|
||||
substep "legacy environment is healthy -- migrating..."
|
||||
Move-Item -Path $OldVenv -Destination $VenvDir -Force
|
||||
Move-Item -LiteralPath $OldVenv -Destination $VenvDir -Force
|
||||
substep "moved .venv -> unsloth_studio"
|
||||
$_Migrated = $true
|
||||
} else {
|
||||
substep "legacy environment failed validation -- creating fresh environment" "Yellow"
|
||||
$invalidVenv = Join-Path $StudioHome (".venv.invalid.{0}.{1}" -f (Get-Date -Format "yyyyMMddHHmmss"), $PID)
|
||||
Move-Item -Path $OldVenv -Destination $invalidVenv -Force -ErrorAction SilentlyContinue
|
||||
Move-Item -LiteralPath $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
|
||||
} elseif (
|
||||
$StudioRedirectMode -ne 'env' `
|
||||
-and (Test-Path -LiteralPath (Join-Path $env:USERPROFILE "unsloth_studio\Scripts\python.exe"))
|
||||
) {
|
||||
# CWD-relative venv from old install.ps1 -> migrate to absolute path.
|
||||
# Skip in env-mode so we don't relocate the default-install venv into
|
||||
# the workspace root.
|
||||
$CwdVenv = Join-Path $env:USERPROFILE "unsloth_studio"
|
||||
substep "found CWD-relative Studio environment, migrating to $VenvDir..."
|
||||
Move-Item -Path $CwdVenv -Destination $VenvDir -Force
|
||||
Move-Item -LiteralPath $CwdVenv -Destination $VenvDir -Force
|
||||
substep "moved ~/unsloth_studio -> ~/.unsloth/studio/unsloth_studio"
|
||||
$_Migrated = $true
|
||||
}
|
||||
|
||||
if (-not (Test-Path $VenvPython)) {
|
||||
if (-not (Test-Path -LiteralPath $VenvPython)) {
|
||||
step "venv" "creating Python $($DetectedPython.Version) virtual environment"
|
||||
substep "$VenvDir"
|
||||
$venvExit = Invoke-InstallCommand { uv venv $VenvDir --python "$($DetectedPython.Path)" }
|
||||
|
|
@ -966,6 +1176,13 @@ shell.Run cmd, 0, False
|
|||
substep "$VenvDir"
|
||||
}
|
||||
|
||||
# Mark the freshly-created venv as Studio-owned so a partial install can be
|
||||
# repaired by re-running install.ps1; the env-mode deletion guard above
|
||||
# accepts this marker as the primary sentinel.
|
||||
if (Test-Path -LiteralPath $VenvDir -PathType Container) {
|
||||
try { [System.IO.File]::WriteAllText((Join-Path $VenvDir ".unsloth-studio-owned"), "") } catch {}
|
||||
}
|
||||
|
||||
# ── Detect GPU (robust: PATH + hardcoded fallback paths, mirrors setup.ps1) ──
|
||||
$HasNvidiaSmi = $false
|
||||
$NvidiaSmiExe = $null
|
||||
|
|
@ -1054,7 +1271,7 @@ shell.Run cmd, 0, False
|
|||
if ($StudioLocalInstall -and (Test-Path (Join-Path $RepoRoot "studio\backend\requirements\no-torch-runtime.txt"))) {
|
||||
return Join-Path $RepoRoot "studio\backend\requirements\no-torch-runtime.txt"
|
||||
}
|
||||
$installed = Get-ChildItem -Path $VenvDir -Recurse -Filter "no-torch-runtime.txt" -ErrorAction SilentlyContinue |
|
||||
$installed = Get-ChildItem -LiteralPath $VenvDir -Recurse -Filter "no-torch-runtime.txt" -ErrorAction SilentlyContinue |
|
||||
Where-Object { $_.FullName -like "*studio*backend*requirements*no-torch-runtime.txt" } |
|
||||
Select-Object -ExpandProperty FullName -First 1
|
||||
return $installed
|
||||
|
|
@ -1068,7 +1285,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.5.1" unsloth-zoo }
|
||||
$baseInstallExit = Invoke-InstallCommand { uv pip install --python $VenvPython --no-deps --reinstall-package unsloth --reinstall-package unsloth-zoo "unsloth>=2026.5.2" unsloth-zoo }
|
||||
if ($baseInstallExit -eq 0) {
|
||||
$NoTorchReq = Find-NoTorchRuntimeFile
|
||||
if ($NoTorchReq) {
|
||||
|
|
@ -1076,7 +1293,7 @@ shell.Run cmd, 0, False
|
|||
}
|
||||
}
|
||||
} else {
|
||||
$baseInstallExit = Invoke-InstallCommand { uv pip install --python $VenvPython --reinstall-package unsloth --reinstall-package unsloth-zoo "unsloth>=2026.5.1" unsloth-zoo }
|
||||
$baseInstallExit = Invoke-InstallCommand { uv pip install --python $VenvPython --reinstall-package unsloth --reinstall-package unsloth-zoo "unsloth>=2026.5.2" unsloth-zoo }
|
||||
}
|
||||
if ($baseInstallExit -ne 0) {
|
||||
Write-Host "[ERROR] Failed to install unsloth (exit code $baseInstallExit)" -ForegroundColor Red
|
||||
|
|
@ -1114,7 +1331,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.5.1" unsloth-zoo }
|
||||
$baseInstallExit = Invoke-InstallCommand { uv pip install --python $VenvPython --no-deps --upgrade-package unsloth --upgrade-package unsloth-zoo "unsloth>=2026.5.2" unsloth-zoo }
|
||||
if ($baseInstallExit -eq 0) {
|
||||
$NoTorchReq = Find-NoTorchRuntimeFile
|
||||
if ($NoTorchReq) {
|
||||
|
|
@ -1122,7 +1339,7 @@ shell.Run cmd, 0, False
|
|||
}
|
||||
}
|
||||
} elseif ($StudioLocalInstall) {
|
||||
$baseInstallExit = Invoke-InstallCommand { uv pip install --python $VenvPython --upgrade-package unsloth "unsloth>=2026.5.1" unsloth-zoo }
|
||||
$baseInstallExit = Invoke-InstallCommand { uv pip install --python $VenvPython --upgrade-package unsloth "unsloth>=2026.5.2" unsloth-zoo }
|
||||
} else {
|
||||
$baseInstallExit = Invoke-InstallCommand { uv pip install --python $VenvPython --upgrade-package unsloth -- "$PackageName" }
|
||||
}
|
||||
|
|
@ -1150,7 +1367,7 @@ 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.5.1" --torch-backend=auto }
|
||||
$baseInstallExit = Invoke-InstallCommand { uv pip install --python $VenvPython unsloth-zoo "unsloth>=2026.5.2" --torch-backend=auto }
|
||||
if ($baseInstallExit -ne 0) {
|
||||
Write-Host "[ERROR] Failed to install unsloth (exit code $baseInstallExit)" -ForegroundColor Red
|
||||
return (Exit-InstallFailure "Failed to install unsloth (exit code $baseInstallExit)" $baseInstallExit)
|
||||
|
|
@ -1192,23 +1409,25 @@ shell.Run cmd, 0, False
|
|||
foreach ($rel in $overlayMap.Keys) {
|
||||
$src = Join-Path $scriptDir $rel
|
||||
$dst = Join-Path $VenvDir $overlayMap[$rel]
|
||||
if (-not (Test-Path $src)) { continue }
|
||||
# -LiteralPath: $VenvDir derives from $StudioHome which may
|
||||
# contain [ ] * ? when the user overrode UNSLOTH_STUDIO_HOME.
|
||||
if (-not (Test-Path -LiteralPath $src)) { continue }
|
||||
$dstParent = Split-Path -Parent $dst
|
||||
if (-not (Test-Path $dstParent)) {
|
||||
if (-not (Test-Path -LiteralPath $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)) {
|
||||
if (-not (Test-Path -LiteralPath $dst)) {
|
||||
# Backfill: target file missing but parent dir exists.
|
||||
Copy-Item $src $dst -Force
|
||||
Copy-Item -LiteralPath $src -Destination $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
|
||||
$srcHash = (Get-FileHash -LiteralPath $src -Algorithm SHA256).Hash
|
||||
$dstHash = (Get-FileHash -LiteralPath $dst -Algorithm SHA256).Hash
|
||||
if ($srcHash -ne $dstHash) {
|
||||
Copy-Item $src $dst -Force
|
||||
Copy-Item -LiteralPath $src -Destination $dst -Force
|
||||
substep ("applied bundled " + (Split-Path -Leaf $rel))
|
||||
}
|
||||
}
|
||||
|
|
@ -1225,7 +1444,8 @@ shell.Run cmd, 0, False
|
|||
Write-TauriLog "STEP" "Running studio setup"
|
||||
step "setup" "running unsloth studio setup..."
|
||||
$UnslothExe = Join-Path $VenvDir "Scripts\unsloth.exe"
|
||||
if (-not (Test-Path $UnslothExe)) {
|
||||
if (-not (Test-Path -LiteralPath $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
|
||||
|
|
@ -1250,6 +1470,15 @@ shell.Run cmd, 0, False
|
|||
# Use 'studio setup' (not 'studio update') because 'update' pops
|
||||
# SKIP_STUDIO_BASE, which would cause redundant package reinstallation
|
||||
# and bypass the fast-path version check from PR #4667.
|
||||
# Propagate UNSLOTH_STUDIO_HOME only for env-override installs; otherwise
|
||||
# an inherited value would put llama.cpp in the wrong place.
|
||||
$previousUnslothStudioHome = $env:UNSLOTH_STUDIO_HOME
|
||||
$hadPreviousUnslothStudioHome = ($null -ne $previousUnslothStudioHome)
|
||||
if ($StudioRedirectMode -eq 'env') {
|
||||
$env:UNSLOTH_STUDIO_HOME = $StudioHome
|
||||
} else {
|
||||
Remove-Item Env:UNSLOTH_STUDIO_HOME -ErrorAction SilentlyContinue
|
||||
}
|
||||
$studioArgs = @('studio', 'setup')
|
||||
if ($script:UnslothVerbose) { $studioArgs += '--verbose' }
|
||||
$env:UNSLOTH_INSTALL_ROLLBACK_MANAGED = "1"
|
||||
|
|
@ -1257,6 +1486,11 @@ shell.Run cmd, 0, False
|
|||
& $UnslothExe @studioArgs
|
||||
$setupExit = $LASTEXITCODE
|
||||
} finally {
|
||||
if ($hadPreviousUnslothStudioHome) {
|
||||
$env:UNSLOTH_STUDIO_HOME = $previousUnslothStudioHome
|
||||
} else {
|
||||
Remove-Item Env:UNSLOTH_STUDIO_HOME -ErrorAction SilentlyContinue
|
||||
}
|
||||
Remove-Item Env:UNSLOTH_INSTALL_ROLLBACK_MANAGED -ErrorAction SilentlyContinue
|
||||
}
|
||||
if ($setupExit -ne 0) {
|
||||
|
|
@ -1301,20 +1535,32 @@ shell.Run cmd, 0, False
|
|||
}
|
||||
} catch { }
|
||||
$ShimDir = Join-Path $StudioHome "bin"
|
||||
New-Item -ItemType Directory -Force -Path $ShimDir | Out-Null
|
||||
[System.IO.Directory]::CreateDirectory($ShimDir) | Out-Null
|
||||
$ShimExe = Join-Path $ShimDir "unsloth.exe"
|
||||
# Fatal preflight outside the lock-handling try/catch -- a directory at
|
||||
# the shim path must not be downgraded to "Continuing with the existing
|
||||
# launcher", or the install finishes with no usable shim.
|
||||
if (Test-Path -LiteralPath $ShimExe -PathType Container) {
|
||||
Write-Host "[ERROR] Cannot create unsloth launcher: $ShimExe is a directory." -ForegroundColor Red
|
||||
Write-Host " Move or remove it manually, then re-run the installer." -ForegroundColor Yellow
|
||||
throw "Cannot create unsloth launcher: $ShimExe is a directory."
|
||||
}
|
||||
# try/catch: if unsloth.exe is locked (Studio running), keep the old shim.
|
||||
$shimUpdated = $false
|
||||
try {
|
||||
if (Test-Path $ShimExe) { Remove-Item $ShimExe -Force -ErrorAction Stop }
|
||||
if (Test-Path -LiteralPath $ShimExe) { Remove-Item -LiteralPath $ShimExe -Force -ErrorAction Stop }
|
||||
try {
|
||||
# New-Item -ItemType HardLink does NOT accept -LiteralPath in any
|
||||
# PowerShell version, so use -Path. Wildcards in $ShimExe (e.g.
|
||||
# brackets in custom roots) glob-expand here and fall through to
|
||||
# the Copy-Item -LiteralPath fallback below.
|
||||
New-Item -ItemType HardLink -Path $ShimExe -Target $UnslothExe -ErrorAction Stop | Out-Null
|
||||
} catch {
|
||||
Copy-Item -Path $UnslothExe -Destination $ShimExe -Force -ErrorAction Stop # fallback: copy
|
||||
Copy-Item -LiteralPath $UnslothExe -Destination $ShimExe -Force -ErrorAction Stop # fallback: copy
|
||||
}
|
||||
$shimUpdated = $true
|
||||
} catch {
|
||||
if (Test-Path $ShimExe) {
|
||||
if (Test-Path -LiteralPath $ShimExe) {
|
||||
Write-Host "[WARN] Could not refresh unsloth launcher at $ShimExe." -ForegroundColor Yellow
|
||||
Write-Host " This usually means a running 'unsloth studio' process still holds the file open." -ForegroundColor Yellow
|
||||
Write-Host " Close Studio and re-run the installer to pick up the latest launcher." -ForegroundColor Yellow
|
||||
|
|
@ -1325,10 +1571,13 @@ shell.Run cmd, 0, False
|
|||
Write-Host " Launch unsloth studio directly via '$UnslothExe' until the next successful install." -ForegroundColor Yellow
|
||||
}
|
||||
}
|
||||
# Only add to PATH when the launcher actually exists on disk.
|
||||
# Add to PATH only when launcher exists. Env-mode: session-only export,
|
||||
# no registry change (workspace path may be deleted later).
|
||||
$pathAdded = $false
|
||||
if (Test-Path $ShimExe) {
|
||||
$pathAdded = Add-ToUserPath -Directory $ShimDir -Position 'Prepend'
|
||||
if (Test-Path -LiteralPath $ShimExe) {
|
||||
if ($StudioRedirectMode -ne 'env') {
|
||||
$pathAdded = Add-ToUserPath -Directory $ShimDir -Position 'Prepend'
|
||||
}
|
||||
}
|
||||
if ($shimUpdated -and $pathAdded) {
|
||||
step "path" "added unsloth launcher to PATH"
|
||||
|
|
@ -1336,12 +1585,20 @@ shell.Run cmd, 0, False
|
|||
Refresh-SessionPath # sync current session with registry
|
||||
Complete-StudioVenvRollback
|
||||
|
||||
# Env-mode session export AFTER Refresh-SessionPath; otherwise a legacy
|
||||
# User PATH entry (Machine > User > current $env:Path) would win.
|
||||
if ($StudioRedirectMode -eq 'env' -and (Test-Path -LiteralPath $ShimExe)) {
|
||||
$env:Path = "$ShimDir;$env:Path"
|
||||
step "path" "exported $ShimDir for this session (no registry PATH change in env-override mode)"
|
||||
}
|
||||
|
||||
# ── Tauri mode: done, skip shortcuts and auto-launch ──
|
||||
if ($TauriMode) {
|
||||
Write-TauriLog "DONE" ""
|
||||
return
|
||||
}
|
||||
|
||||
# New-StudioShortcuts gates the .lnk shortcuts on env-mode internally.
|
||||
New-StudioShortcuts -UnslothExePath $UnslothExe
|
||||
|
||||
# In interactive terminals, ask the user before starting Studio.
|
||||
|
|
@ -1360,8 +1617,21 @@ shell.Run cmd, 0, False
|
|||
}
|
||||
} else {
|
||||
step "launch" "manual commands:"
|
||||
substep "& `"$VenvDir\Scripts\Activate.ps1`""
|
||||
substep "unsloth studio -p 8888"
|
||||
# Single-quote the printed paths so $-vars / backticks in custom roots
|
||||
# do not reparse when the user pastes the command.
|
||||
$_actLiteral = "'" + ((Join-Path $VenvDir "Scripts\Activate.ps1") -replace "'", "''") + "'"
|
||||
if ($StudioRedirectMode -eq 'env') {
|
||||
# Env-mode skips registry PATH; print the absolute shim path.
|
||||
$_shim = Join-Path $StudioHome "bin\unsloth.exe"
|
||||
$_shimLiteral = "'" + ($_shim -replace "'", "''") + "'"
|
||||
substep "& $_shimLiteral studio -p 8888"
|
||||
substep "or activate env first:"
|
||||
substep "& $_actLiteral"
|
||||
substep "unsloth studio -p 8888"
|
||||
} else {
|
||||
substep "& $_actLiteral"
|
||||
substep "unsloth studio -p 8888"
|
||||
}
|
||||
substep "(add -H 0.0.0.0 to allow network / cloud access)"
|
||||
Write-Host ""
|
||||
}
|
||||
|
|
|
|||
431
install.sh
431
install.sh
|
|
@ -6,6 +6,12 @@
|
|||
# Usage (no-torch): ./install.sh --no-torch (skip PyTorch, GGUF-only mode)
|
||||
# Usage (test): ./install.sh --package roland-sloth (install a different package name)
|
||||
# Usage (py): ./install.sh --python 3.12 (override auto-detected Python version)
|
||||
#
|
||||
# Env vars (priority: UNSLOTH_STUDIO_HOME > STUDIO_HOME > HOME-redirect > default):
|
||||
# UNSLOTH_STUDIO_HOME=/abs/path -> install under that path
|
||||
# STUDIO_HOME=/abs/path -> alias, same effect (UNSLOTH_STUDIO_HOME wins)
|
||||
# (DATA_DIR + unsloth CLI shim nest inside; no shell rc-file append.)
|
||||
# Default ($HOME/.unsloth/studio) is preserved when no env var is set.
|
||||
set -e
|
||||
|
||||
# ── Output style (aligned with studio/setup.sh) ──
|
||||
|
|
@ -66,6 +72,56 @@ if [ "$_VERBOSE" = true ]; then
|
|||
export UNSLOTH_VERBOSE=1
|
||||
fi
|
||||
|
||||
# Custom Studio roots are not supported with --tauri (desktop app still
|
||||
# resolves ~/.unsloth/studio). Pass through if the override == legacy default.
|
||||
if [ "$TAURI_MODE" = true ]; then
|
||||
_tauri_override_var=""
|
||||
_tauri_override="${UNSLOTH_STUDIO_HOME:-}"
|
||||
if [ -n "$_tauri_override" ]; then
|
||||
_tauri_override_var="UNSLOTH_STUDIO_HOME"
|
||||
else
|
||||
_tauri_override="${STUDIO_HOME:-}"
|
||||
[ -n "$_tauri_override" ] && _tauri_override_var="STUDIO_HOME"
|
||||
fi
|
||||
# Strip whitespace so " " is treated as unset (matches Python .strip()).
|
||||
_tauri_override=$(printf '%s' "$_tauri_override" | sed -e 's/^[[:space:]]*//' -e 's/[[:space:]]*$//')
|
||||
if [ -n "$_tauri_override" ]; then
|
||||
case "$_tauri_override" in
|
||||
"~") _tauri_override="$HOME" ;;
|
||||
"~/"*) _tauri_override="$HOME/${_tauri_override#'~/'}" ;;
|
||||
esac
|
||||
# Canonicalize both sides (CDPATH=, -P) so a CDPATH-set env or
|
||||
# symlinked $HOME doesn't break the legacy-equality comparison.
|
||||
if [ -d "$_tauri_override" ]; then
|
||||
_tauri_override_abs=$(CDPATH= cd -P -- "$_tauri_override" 2>/dev/null && pwd -P) \
|
||||
|| _tauri_override_abs="$_tauri_override"
|
||||
else
|
||||
_tauri_override_abs="$_tauri_override"
|
||||
fi
|
||||
# Strip trailing separators so ".../studio/" matches ".../studio".
|
||||
while [ "$_tauri_override_abs" != "/" ] \
|
||||
&& [ "${_tauri_override_abs%/}" != "$_tauri_override_abs" ]; do
|
||||
_tauri_override_abs=${_tauri_override_abs%/}
|
||||
done
|
||||
_tauri_legacy_root="$HOME/.unsloth/studio"
|
||||
if [ -d "$_tauri_legacy_root" ]; then
|
||||
_tauri_legacy_root=$(CDPATH= cd -P -- "$_tauri_legacy_root" 2>/dev/null && pwd -P) \
|
||||
|| _tauri_legacy_root="$HOME/.unsloth/studio"
|
||||
fi
|
||||
while [ "$_tauri_legacy_root" != "/" ] \
|
||||
&& [ "${_tauri_legacy_root%/}" != "$_tauri_legacy_root" ]; do
|
||||
_tauri_legacy_root=${_tauri_legacy_root%/}
|
||||
done
|
||||
if [ "$_tauri_override_abs" != "$_tauri_legacy_root" ]; then
|
||||
echo "ERROR: $_tauri_override_var is not supported with --tauri." >&2
|
||||
echo " The desktop app still uses the legacy ~/.unsloth/studio root." >&2
|
||||
echo " Run install.sh without --tauri for custom-root shell installs," >&2
|
||||
echo " or unset the env var for default desktop installs." >&2
|
||||
exit 1
|
||||
fi
|
||||
fi
|
||||
fi
|
||||
|
||||
_is_verbose() {
|
||||
[ "${UNSLOTH_VERBOSE:-0}" = "1" ]
|
||||
}
|
||||
|
|
@ -219,7 +275,67 @@ _tauri_gpu_branch() {
|
|||
}
|
||||
|
||||
PYTHON_VERSION="" # resolved after platform detection
|
||||
STUDIO_HOME="$HOME/.unsloth/studio"
|
||||
|
||||
# Resolve install destinations: env override, HOME-redirect (best-effort
|
||||
# via getent/dscl), or default. Env-var priority: UNSLOTH_STUDIO_HOME wins
|
||||
# over STUDIO_HOME (the more specific signal beats the generic alias).
|
||||
_resolve_studio_destinations() {
|
||||
_override_var=""
|
||||
_override="${UNSLOTH_STUDIO_HOME:-}"
|
||||
if [ -n "$_override" ]; then
|
||||
_override_var="UNSLOTH_STUDIO_HOME"
|
||||
else
|
||||
_override="${STUDIO_HOME:-}"
|
||||
[ -n "$_override" ] && _override_var="STUDIO_HOME"
|
||||
fi
|
||||
# Strip surrounding whitespace so " " is treated as unset (matches the
|
||||
# Python resolvers' .strip()), preventing install/runtime layout drift.
|
||||
_override=$(printf '%s' "$_override" | sed -e 's/^[[:space:]]*//' -e 's/[[:space:]]*$//')
|
||||
# Tilde expansion: env vars are not subject to it when quoted on assignment.
|
||||
case "$_override" in
|
||||
"~") _override="$HOME" ;;
|
||||
"~/"*) _override="$HOME/${_override#'~/'}" ;;
|
||||
esac
|
||||
if [ -n "$_override" ]; then
|
||||
mkdir -p -- "$_override" 2>/dev/null || { echo "ERROR: $_override_var=$_override cannot be created." >&2; exit 1; }
|
||||
[ -w "$_override" ] || { echo "ERROR: $_override_var=$_override is not writable." >&2; exit 1; }
|
||||
STUDIO_HOME="$(CDPATH= cd -P -- "$_override" && pwd -P)" || exit 1
|
||||
DATA_DIR="$STUDIO_HOME/share"
|
||||
_LOCAL_BIN="$STUDIO_HOME/bin"
|
||||
_STUDIO_HOME_REDIRECT=env
|
||||
substep "custom $_override_var=$STUDIO_HOME"
|
||||
return 0
|
||||
fi
|
||||
_default_home=""
|
||||
if command -v getent >/dev/null 2>&1; then
|
||||
_default_home=$(getent passwd "${USER:-$(whoami)}" 2>/dev/null | cut -d: -f6)
|
||||
elif [ "$(uname)" = "Darwin" ] && command -v dscl >/dev/null 2>&1; then
|
||||
_default_home=$(dscl . -read "/Users/${USER:-$(whoami)}" NFSHomeDirectory 2>/dev/null | awk '{print $2}')
|
||||
fi
|
||||
# Canonicalize both sides so a trailing slash on $HOME (or symlink mismatch
|
||||
# with passwd-DB output) doesn't misfire the redirection branch.
|
||||
_home_canon="$HOME"
|
||||
if [ -d "$_home_canon" ]; then
|
||||
_home_canon=$(CDPATH= cd -P -- "$_home_canon" 2>/dev/null && pwd -P) || _home_canon="$HOME"
|
||||
fi
|
||||
_default_home_canon="$_default_home"
|
||||
if [ -n "$_default_home_canon" ] && [ -d "$_default_home_canon" ]; then
|
||||
_default_home_canon=$(CDPATH= cd -P -- "$_default_home_canon" 2>/dev/null && pwd -P) || _default_home_canon="$_default_home"
|
||||
fi
|
||||
if [ -n "$_default_home_canon" ] && [ "$_home_canon" != "$_default_home_canon" ]; then
|
||||
STUDIO_HOME="$HOME/.unsloth/studio"
|
||||
DATA_DIR="$HOME/.local/share/unsloth"
|
||||
_LOCAL_BIN="$HOME/.local/bin"
|
||||
_STUDIO_HOME_REDIRECT=home
|
||||
substep "HOME redirected ($HOME); install follows \$HOME"
|
||||
return 0
|
||||
fi
|
||||
STUDIO_HOME="$HOME/.unsloth/studio"
|
||||
DATA_DIR="$HOME/.local/share/unsloth"
|
||||
_LOCAL_BIN="$HOME/.local/bin"
|
||||
_STUDIO_HOME_REDIRECT=default
|
||||
}
|
||||
_resolve_studio_destinations
|
||||
VENV_DIR="$STUDIO_HOME/unsloth_studio"
|
||||
_VENV_ROLLBACK_DIR=""
|
||||
_VENV_ROLLBACK_TARGET="$VENV_DIR"
|
||||
|
|
@ -383,23 +499,65 @@ create_studio_shortcuts() {
|
|||
_css_exe_dir=$(cd "$(dirname "$_css_exe")" && pwd)
|
||||
_css_exe="$_css_exe_dir/$(basename "$_css_exe")"
|
||||
|
||||
_css_data_dir="$HOME/.local/share/unsloth"
|
||||
_css_data_dir="$DATA_DIR"
|
||||
_css_launcher="$_css_data_dir/launch-studio.sh"
|
||||
_css_icon_png="$_css_data_dir/unsloth-studio.png"
|
||||
_css_gem_png="$_css_data_dir/unsloth-gem.png"
|
||||
|
||||
mkdir -p "$_css_data_dir"
|
||||
|
||||
# Same-install discriminator: per-install opaque id written once at install
|
||||
# time and read by both this launcher and the backend (/api/health). Replaces
|
||||
# the older sha256(canonical $STUDIO_HOME) scheme to (a) avoid leaking the
|
||||
# install path on -H 0.0.0.0 deployments and (b) sidestep launcher/backend
|
||||
# canonicalization drift (cd -P vs Path.resolve() symlink/junction handling).
|
||||
# Lives at $STUDIO_HOME/share/ (not $DATA_DIR) so the backend can find it
|
||||
# via _STUDIO_ROOT_RESOLVED / "share" / "studio_install_id" regardless of
|
||||
# mode (in env-mode $STUDIO_HOME/share == $DATA_DIR; in default mode they
|
||||
# diverge but the backend only knows the studio_root). 32 bytes of urandom
|
||||
# -> 64 hex chars, byte-compatible with the prior digest so launcher
|
||||
# placeholder, _check_health, and tests stay length-agnostic.
|
||||
_css_id_dir="$STUDIO_HOME/share"
|
||||
mkdir -p "$_css_id_dir"
|
||||
_css_id_file="$_css_id_dir/studio_install_id"
|
||||
if [ ! -s "$_css_id_file" ]; then
|
||||
if [ -r /dev/urandom ]; then
|
||||
_css_new_id=$(od -An -N32 -tx1 /dev/urandom 2>/dev/null | tr -d ' \n')
|
||||
fi
|
||||
if [ -z "${_css_new_id:-}" ] && command -v python3 >/dev/null 2>&1; then
|
||||
_css_new_id=$(python3 -c 'import secrets; print(secrets.token_hex(32))' 2>/dev/null)
|
||||
fi
|
||||
if [ -z "${_css_new_id:-}" ]; then
|
||||
echo "[WARN] Cannot create launcher: no entropy source for studio_install_id" >&2
|
||||
return 1
|
||||
fi
|
||||
# Atomic write so a partial install can't leave a half-written id.
|
||||
_css_id_tmp="$_css_id_file.$$.tmp"
|
||||
printf '%s' "$_css_new_id" > "$_css_id_tmp" \
|
||||
&& mv "$_css_id_tmp" "$_css_id_file"
|
||||
chmod 600 "$_css_id_file" 2>/dev/null || true
|
||||
unset _css_new_id _css_id_tmp
|
||||
fi
|
||||
_css_studio_root_id=$(cat "$_css_id_file" 2>/dev/null)
|
||||
if [ -z "$_css_studio_root_id" ]; then
|
||||
echo "[WARN] Cannot create launcher: failed to read $_css_id_file" >&2
|
||||
return 1
|
||||
fi
|
||||
_css_is_env_mode=false
|
||||
[ "$_STUDIO_HOME_REDIRECT" = "env" ] && _css_is_env_mode=true
|
||||
|
||||
# ── Write launcher script ──
|
||||
# The launcher is Bash (not POSIX sh).
|
||||
# We write it with a placeholder and substitute the exe path via sed.
|
||||
# Single-quoted heredoc; @@DATA_DIR@@, @@STUDIO_ROOT_ID@@, and
|
||||
# @@INSTALLED_IS_ENV_MODE@@ are substituted via sed below.
|
||||
cat > "$_css_launcher" << 'LAUNCHER_EOF'
|
||||
#!/usr/bin/env bash
|
||||
# Unsloth Studio Launcher
|
||||
# Auto-generated by install.sh -- do not edit manually.
|
||||
set -euo pipefail
|
||||
|
||||
DATA_DIR="$HOME/.local/share/unsloth"
|
||||
DATA_DIR='@@DATA_DIR@@'
|
||||
_EXPECTED_STUDIO_ROOT_ID='@@STUDIO_ROOT_ID@@'
|
||||
_INSTALLED_IS_ENV_MODE='@@INSTALLED_IS_ENV_MODE@@'
|
||||
|
||||
# Read exe path from config written at install time.
|
||||
# Sourcing is safe: the config file is written by install.sh, not user input.
|
||||
|
|
@ -416,7 +574,23 @@ MAX_PORT_OFFSET=20
|
|||
TIMEOUT_SEC=60
|
||||
POLL_INTERVAL_SEC=1
|
||||
LOG_FILE="$DATA_DIR/studio.log"
|
||||
# why: in env-override mode multiple installs share an OS user; namespace the
|
||||
# lock and remember our own healthy port so we never attach to an unrelated
|
||||
# Studio listening on the global 8888..8908 range.
|
||||
LOCK_DIR="${XDG_RUNTIME_DIR:-/tmp}/unsloth-studio-launcher-$(id -u).lock"
|
||||
PORT_FILE=""
|
||||
# why: gate on the install-time mode (baked above) instead of the runtime env
|
||||
# var; sourcing a custom-root studio.conf in shell must not flip a default-mode
|
||||
# launcher into env-mode behavior with stale state.
|
||||
if [ "$_INSTALLED_IS_ENV_MODE" = "true" ]; then
|
||||
if command -v cksum >/dev/null 2>&1; then
|
||||
_LOCK_KEY=$(printf '%s' "$DATA_DIR" | cksum | awk '{print $1}')
|
||||
else
|
||||
_LOCK_KEY=""
|
||||
fi
|
||||
[ -n "$_LOCK_KEY" ] && LOCK_DIR="${XDG_RUNTIME_DIR:-/tmp}/unsloth-studio-launcher-$(id -u)-${_LOCK_KEY}.lock"
|
||||
PORT_FILE="$DATA_DIR/studio.port"
|
||||
fi
|
||||
|
||||
# ── HTTP GET helper (supports curl and wget) ──
|
||||
_http_get() {
|
||||
|
|
@ -435,10 +609,20 @@ _check_health() {
|
|||
_port=$1
|
||||
_resp=$(_http_get "http://127.0.0.1:$_port/api/health") || return 1
|
||||
case "$_resp" in
|
||||
*'"status"'*'"healthy"'*'"service"'*'"Unsloth UI Backend"'*) return 0 ;;
|
||||
*'"service"'*'"Unsloth UI Backend"'*'"status"'*'"healthy"'*) return 0 ;;
|
||||
*'"status"'*'"healthy"'*'"service"'*'"Unsloth UI Backend"'*) ;;
|
||||
*'"service"'*'"Unsloth UI Backend"'*'"status"'*'"healthy"'*) ;;
|
||||
*) return 1 ;;
|
||||
esac
|
||||
return 1
|
||||
# why: verify the backend belongs to THIS install. Baked hex digest avoids
|
||||
# JSON-escape mismatches on paths with `\`/`"` and avoids leaking the raw
|
||||
# install path to unauthenticated callers.
|
||||
if [ -n "$_EXPECTED_STUDIO_ROOT_ID" ]; then
|
||||
case "$_resp" in
|
||||
*"\"studio_root_id\":\"$_EXPECTED_STUDIO_ROOT_ID\""*|*"\"studio_root_id\": \"$_EXPECTED_STUDIO_ROOT_ID\""*) return 0 ;;
|
||||
*) return 1 ;;
|
||||
esac
|
||||
fi
|
||||
return 0
|
||||
}
|
||||
|
||||
# ── Port scanning ──
|
||||
|
|
@ -461,6 +645,25 @@ _candidate_ports() {
|
|||
}
|
||||
|
||||
_find_healthy_port() {
|
||||
if [ -n "$PORT_FILE" ] && [ -f "$PORT_FILE" ]; then
|
||||
# why: env-mode installs only attach to a port we previously launched
|
||||
# ourselves; never to a sibling Studio that happens to be healthy.
|
||||
_p=$(cat "$PORT_FILE" 2>/dev/null || true)
|
||||
case "$_p" in
|
||||
''|*[!0-9]*) ;;
|
||||
*)
|
||||
if _check_health "$_p"; then
|
||||
echo "$_p"
|
||||
return 0
|
||||
fi
|
||||
rm -f "$PORT_FILE"
|
||||
;;
|
||||
esac
|
||||
return 1
|
||||
fi
|
||||
if [ -n "$PORT_FILE" ]; then
|
||||
return 1
|
||||
fi
|
||||
for _p in $(_candidate_ports | sort -un); do
|
||||
if _check_health "$_p"; then
|
||||
echo "$_p"
|
||||
|
|
@ -611,6 +814,7 @@ if [ -t 1 ]; then
|
|||
_obwr_deadline=$(($(date +%s) + TIMEOUT_SEC))
|
||||
while [ "$(date +%s)" -lt "$_obwr_deadline" ]; do
|
||||
if _check_health "$_launch_port"; then
|
||||
[ -n "$PORT_FILE" ] && printf '%s\n' "$_launch_port" > "$PORT_FILE" 2>/dev/null || true
|
||||
_release_lock
|
||||
_open_browser "http://localhost:$_launch_port"
|
||||
exit 0
|
||||
|
|
@ -634,6 +838,7 @@ else
|
|||
_deadline=$(($(date +%s) + TIMEOUT_SEC))
|
||||
while [ "$(date +%s)" -lt "$_deadline" ]; do
|
||||
if _check_health "$_launch_port"; then
|
||||
[ -n "$PORT_FILE" ] && printf '%s\n' "$_launch_port" > "$PORT_FILE" 2>/dev/null || true
|
||||
_open_browser "http://localhost:$_launch_port"
|
||||
exit 0
|
||||
fi
|
||||
|
|
@ -646,13 +851,62 @@ else
|
|||
fi
|
||||
LAUNCHER_EOF
|
||||
|
||||
# why: bake non-user-controlled placeholders FIRST so a literal
|
||||
# `@@STUDIO_ROOT_ID@@` inside $DATA_DIR cannot be rewritten below.
|
||||
sed -e "s|@@STUDIO_ROOT_ID@@|$_css_studio_root_id|g" \
|
||||
-e "s|@@INSTALLED_IS_ENV_MODE@@|$_css_is_env_mode|g" \
|
||||
"$_css_launcher" > "$_css_launcher.tmp" \
|
||||
&& mv "$_css_launcher.tmp" "$_css_launcher"
|
||||
|
||||
# Env-mode bakes an absolute DATA_DIR (root fixed at install time);
|
||||
# default / HOME-redirect keeps the literal $HOME/.local/share/unsloth
|
||||
# so behavior is byte-identical to pre-override.
|
||||
if [ "$_STUDIO_HOME_REDIRECT" = "env" ]; then
|
||||
# Two-stage escape: (1) `'` -> `'\''` for shell single-quote embedding,
|
||||
# (2) backslash/&/| escape so the value survives the s|...|VALUE| sed
|
||||
# below. Verified end-to-end with apostrophes, spaces, &, |, $.
|
||||
_sq_escaped=$(printf '%s' "$DATA_DIR" | sed "s/'/'\\\\''/g")
|
||||
_sed_safe=$(printf '%s' "$_sq_escaped" | sed 's/[\\&|]/\\&/g')
|
||||
sed "s|@@DATA_DIR@@|$_sed_safe|g" "$_css_launcher" > "$_css_launcher.tmp" \
|
||||
&& mv "$_css_launcher.tmp" "$_css_launcher"
|
||||
else
|
||||
sed "s|DATA_DIR='@@DATA_DIR@@'|DATA_DIR=\"\$HOME/.local/share/unsloth\"|" \
|
||||
"$_css_launcher" > "$_css_launcher.tmp" \
|
||||
&& mv "$_css_launcher.tmp" "$_css_launcher"
|
||||
fi
|
||||
|
||||
chmod +x "$_css_launcher"
|
||||
|
||||
# Write the exe path to a separate conf file sourced by the launcher.
|
||||
# Using single-quote wrapping with the standard '\'' escape for any
|
||||
# embedded apostrophes. This avoids all sed metacharacter issues.
|
||||
# studio.conf: exe path + (env-mode only) persisted env vars so fresh
|
||||
# shells launch the right install without re-exporting.
|
||||
_css_quoted_exe=$(printf '%s' "$_css_exe" | sed "s/'/'\\\\''/g")
|
||||
printf '%s\n' "UNSLOTH_EXE='$_css_quoted_exe'" > "$_css_data_dir/studio.conf"
|
||||
{
|
||||
printf '%s\n' "UNSLOTH_EXE='$_css_quoted_exe'"
|
||||
if [ "$_STUDIO_HOME_REDIRECT" = "env" ]; then
|
||||
# When an override resolves to the legacy default, llama.cpp
|
||||
# still lives at ~/.unsloth/llama.cpp (one shared build).
|
||||
# Canonicalize the legacy side so a symlinked $HOME doesn't
|
||||
# break the comparison.
|
||||
_css_legacy_studio="$HOME/.unsloth/studio"
|
||||
if [ -d "$_css_legacy_studio" ]; then
|
||||
_css_legacy_studio=$(CDPATH= cd -P -- "$_css_legacy_studio" 2>/dev/null && pwd -P) \
|
||||
|| _css_legacy_studio="$HOME/.unsloth/studio"
|
||||
fi
|
||||
if [ "$STUDIO_HOME" = "$_css_legacy_studio" ]; then
|
||||
_css_llama_path="$HOME/.unsloth/llama.cpp"
|
||||
else
|
||||
_css_llama_path="$STUDIO_HOME/llama.cpp"
|
||||
fi
|
||||
_css_quoted_home=$(printf '%s' "$STUDIO_HOME" | sed "s/'/'\\\\''/g")
|
||||
_css_quoted_llama=$(printf '%s' "$_css_llama_path" | sed "s/'/'\\\\''/g")
|
||||
printf '%s\n' "export UNSLOTH_STUDIO_HOME='$_css_quoted_home'"
|
||||
# UNSLOTH_LLAMA_CPP_PATH is a pre-existing user-controlled
|
||||
# llama.cpp dir override; only default it if unset.
|
||||
printf '%s\n' 'if [ -z "${UNSLOTH_LLAMA_CPP_PATH:-}" ]; then'
|
||||
printf '%s\n' " export UNSLOTH_LLAMA_CPP_PATH='$_css_quoted_llama'"
|
||||
printf '%s\n' 'fi'
|
||||
fi
|
||||
} > "$_css_data_dir/studio.conf"
|
||||
|
||||
# ── Icon: try bundled, then download ──
|
||||
# rounded-512.png used for both Linux and macOS icons
|
||||
|
|
@ -698,6 +952,14 @@ LAUNCHER_EOF
|
|||
fi
|
||||
|
||||
# ── Platform-specific shortcuts ──
|
||||
# Env-mode installs are workspace-scoped: skip persistent desktop /
|
||||
# Start-Menu / dock launchers that may point at a deleted workspace.
|
||||
# Runtime launcher + studio.conf + icon are still written above.
|
||||
if [ "$_STUDIO_HOME_REDIRECT" = "env" ]; then
|
||||
substep "wrote launcher at $_css_launcher (persistent shortcuts skipped in env-override mode)"
|
||||
return 0
|
||||
fi
|
||||
|
||||
_css_created=0
|
||||
|
||||
if [ "$_css_os" = "linux" ]; then
|
||||
|
|
@ -775,11 +1037,18 @@ DESKTOP_EOF
|
|||
</plist>
|
||||
PLIST_EOF
|
||||
|
||||
# Executable stub
|
||||
cat > "$_css_macos_dir/launch-studio" << STUB_EOF
|
||||
# Executable stub: same single-quoted-heredoc + sed-substitute
|
||||
# pattern as launch-studio.sh so $-vars in $_css_data_dir don't
|
||||
# expand at .app launch time.
|
||||
_css_sq_dir=$(printf '%s' "$_css_data_dir" | sed "s/'/'\\\\''/g")
|
||||
_css_sed_dir=$(printf '%s' "$_css_sq_dir" | sed 's/[\\&|]/\\&/g')
|
||||
cat > "$_css_macos_dir/launch-studio" << 'STUB_EOF'
|
||||
#!/bin/sh
|
||||
exec "$HOME/.local/share/unsloth/launch-studio.sh" "\$@"
|
||||
exec '@@DATA_DIR@@/launch-studio.sh' "$@"
|
||||
STUB_EOF
|
||||
sed "s|@@DATA_DIR@@|$_css_sed_dir|g" "$_css_macos_dir/launch-studio" \
|
||||
> "$_css_macos_dir/launch-studio.tmp" \
|
||||
&& mv "$_css_macos_dir/launch-studio.tmp" "$_css_macos_dir/launch-studio"
|
||||
chmod +x "$_css_macos_dir/launch-studio"
|
||||
|
||||
# Build AppIcon.icns from unsloth-gem.png (2240x2240)
|
||||
|
|
@ -1079,11 +1348,28 @@ mkdir -p "$STUDIO_HOME"
|
|||
_MIGRATED=false
|
||||
|
||||
if [ -x "$VENV_DIR/bin/python" ]; then
|
||||
# why: matching guard to the .venv branch below -- in env-mode
|
||||
# $STUDIO_HOME is a user-chosen workspace, so refuse to nuke an
|
||||
# existing $STUDIO_HOME/unsloth_studio that lacks Studio sentinels.
|
||||
# Accept the in-VENV ownership marker so partial-install retries are
|
||||
# not blocked. Sentinels must be regular files: -f follows symlinks
|
||||
# to files (the legitimate ln -s shim shape) but rejects directories
|
||||
# and broken/dir-targeted symlinks.
|
||||
if [ "$_STUDIO_HOME_REDIRECT" = "env" ] \
|
||||
&& [ ! -f "$VENV_DIR/.unsloth-studio-owned" ] \
|
||||
&& [ ! -f "$STUDIO_HOME/share/studio.conf" ] \
|
||||
&& [ ! -f "$STUDIO_HOME/bin/unsloth" ]; then
|
||||
echo "ERROR: $VENV_DIR already exists but does not look like an Unsloth Studio install." >&2
|
||||
echo " Move it aside or choose an empty UNSLOTH_STUDIO_HOME." >&2
|
||||
exit 1
|
||||
fi
|
||||
# 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
|
||||
elif [ "$_STUDIO_HOME_REDIRECT" != "env" ] && [ -x "$STUDIO_HOME/.venv/bin/python" ]; then
|
||||
# Old layout exists — validate before migrating.
|
||||
# Skip in env-mode so we don't rm -rf an unrelated .venv at the
|
||||
# workspace root (e.g. user's existing project Python venv).
|
||||
# In no-torch mode, a missing torch package is expected; validate Python only.
|
||||
substep "found legacy Studio environment, validating..."
|
||||
_legacy_ok=false
|
||||
|
|
@ -1132,6 +1418,13 @@ if [ ! -x "$VENV_DIR/bin/python" ]; then
|
|||
run_install_cmd "create venv" uv venv "$VENV_DIR" --python "$PYTHON_VERSION"
|
||||
fi
|
||||
|
||||
# Mark the freshly-created venv as Studio-owned so a partial install can be
|
||||
# repaired by re-running install.sh; the env-mode deletion guard above accepts
|
||||
# this marker as the primary sentinel.
|
||||
if [ -x "$VENV_DIR/bin/python" ]; then
|
||||
: > "$VENV_DIR/.unsloth-studio-owned" 2>/dev/null || true
|
||||
fi
|
||||
|
||||
# Guard against Python 3.13.8 torch import bug on Apple Silicon
|
||||
# (skip when the user explicitly chose a version via --python)
|
||||
if [ -z "$_USER_PYTHON" ] && [ "$OS" = "macos" ] && [ "$_ARCH" = "arm64" ]; then
|
||||
|
|
@ -1143,6 +1436,9 @@ if [ -z "$_USER_PYTHON" ] && [ "$OS" = "macos" ] && [ "$_ARCH" = "arm64" ]; then
|
|||
rm -rf "$VENV_DIR"
|
||||
PYTHON_VERSION="3.12"
|
||||
run_install_cmd "recreate venv" uv venv "$VENV_DIR" --python "$PYTHON_VERSION"
|
||||
if [ -x "$VENV_DIR/bin/python" ]; then
|
||||
: > "$VENV_DIR/.unsloth-studio-owned" 2>/dev/null || true
|
||||
fi
|
||||
fi
|
||||
fi
|
||||
|
||||
|
|
@ -1486,7 +1782,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.5.1" unsloth-zoo
|
||||
"unsloth>=2026.5.2" 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"
|
||||
|
|
@ -1494,7 +1790,7 @@ 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.5.1" unsloth-zoo
|
||||
"unsloth>=2026.5.2" unsloth-zoo
|
||||
fi
|
||||
if [ "$STUDIO_LOCAL_INSTALL" = true ]; then
|
||||
substep "overlaying local repo (editable)..."
|
||||
|
|
@ -1662,7 +1958,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.5.1" unsloth-zoo
|
||||
"unsloth>=2026.5.2" 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"
|
||||
|
|
@ -1677,7 +1973,7 @@ elif [ -n "$TORCH_INDEX_URL" ]; then
|
|||
fi
|
||||
elif [ "$STUDIO_LOCAL_INSTALL" = true ]; then
|
||||
run_install_cmd "install unsloth (local)" uv pip install --python "$_VENV_PY" \
|
||||
--upgrade-package unsloth "unsloth>=2026.5.1" unsloth-zoo
|
||||
--upgrade-package unsloth "unsloth>=2026.5.2" 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..."
|
||||
|
|
@ -1709,7 +2005,7 @@ 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.5.1" --torch-backend=auto
|
||||
run_install_cmd "install unsloth (auto torch backend)" uv pip install --python "$_VENV_PY" unsloth-zoo "unsloth>=2026.5.2" --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..."
|
||||
|
|
@ -1721,6 +2017,12 @@ else
|
|||
fi
|
||||
fi
|
||||
|
||||
# ── Install mlx-vlm on Apple Silicon (optional, for VLM training) ──
|
||||
if [ "$OS" = "macos" ] && [ "$_ARCH" = "arm64" ]; then
|
||||
substep "installing mlx-vlm (VLM training support)..."
|
||||
run_install_cmd "install mlx-vlm" uv pip install --python "$_VENV_PY" mlx-vlm
|
||||
fi
|
||||
|
||||
# ── Run studio setup ──
|
||||
tauri_log "STEP" "Running Studio setup"
|
||||
# When --local, use the repo's own setup.sh directly.
|
||||
|
|
@ -1768,7 +2070,17 @@ _SKIP_FRONTEND=0
|
|||
if [ "$TAURI_MODE" = true ]; then
|
||||
_SKIP_FRONTEND=1
|
||||
fi
|
||||
# Prepend UNSLOTH_STUDIO_HOME=$STUDIO_HOME to "$@" for env-override installs
|
||||
# without word-splitting on whitespace paths.
|
||||
_run_setup_with_studio_home() {
|
||||
if [ "$_STUDIO_HOME_REDIRECT" = "env" ]; then
|
||||
UNSLOTH_STUDIO_HOME="$STUDIO_HOME" "$@"
|
||||
else
|
||||
"$@"
|
||||
fi
|
||||
}
|
||||
if [ "$STUDIO_LOCAL_INSTALL" = true ]; then
|
||||
_run_setup_with_studio_home env \
|
||||
SKIP_STUDIO_BASE="$_SKIP_BASE" \
|
||||
SKIP_STUDIO_FRONTEND="$_SKIP_FRONTEND" \
|
||||
STUDIO_PACKAGE_NAME="$PACKAGE_NAME" \
|
||||
|
|
@ -1782,6 +2094,7 @@ else
|
|||
# the same session) does not silently flip a normal install onto the
|
||||
# local-dev path in setup.sh and install_python_stack.py. Mirrors the
|
||||
# reset already done in install.ps1 for PowerShell.
|
||||
_run_setup_with_studio_home env \
|
||||
SKIP_STUDIO_BASE="$_SKIP_BASE" \
|
||||
SKIP_STUDIO_FRONTEND="$_SKIP_FRONTEND" \
|
||||
STUDIO_PACKAGE_NAME="$PACKAGE_NAME" \
|
||||
|
|
@ -1791,36 +2104,53 @@ else
|
|||
bash "$SETUP_SH" </dev/null || _SETUP_EXIT=$?
|
||||
fi
|
||||
|
||||
# ── Make 'unsloth' available globally via ~/.local/bin ──
|
||||
mkdir -p "$HOME/.local/bin"
|
||||
ln -sf "$VENV_DIR/bin/unsloth" "$HOME/.local/bin/unsloth"
|
||||
# ── Make 'unsloth' available via $_LOCAL_BIN (resolved earlier) ──
|
||||
# Env-mode: $_LOCAL_BIN is $STUDIO_HOME/bin; skip shell-rc PATH append so we
|
||||
# don't pollute the user's profile with a workspace-scoped path.
|
||||
mkdir -p "$_LOCAL_BIN"
|
||||
# ln -sf into an existing dir creates link inside it. Refuse to delete a
|
||||
# real directory at the shim path -- that could destroy unrelated user data.
|
||||
_shim_path="$_LOCAL_BIN/unsloth"
|
||||
if [ -d "$_shim_path" ] && [ ! -L "$_shim_path" ]; then
|
||||
echo "ERROR: $_shim_path is a directory; refusing to delete it." >&2
|
||||
echo " Move or remove it manually, then re-run the installer." >&2
|
||||
exit 1
|
||||
fi
|
||||
# why: -sfn is atomic and -n prevents descent into a symlink-to-directory at
|
||||
# the shim path (the directory guard above already rejects a real directory).
|
||||
ln -sfn "$VENV_DIR/bin/unsloth" "$_shim_path"
|
||||
|
||||
_LOCAL_BIN="$HOME/.local/bin"
|
||||
case ":$PATH:" in
|
||||
*":$_LOCAL_BIN:"*) ;; # already on PATH
|
||||
*)
|
||||
_SHELL_PROFILE=""
|
||||
if [ -n "${ZSH_VERSION:-}" ] || [ "$(basename "${SHELL:-}")" = "zsh" ]; then
|
||||
_SHELL_PROFILE="$HOME/.zshrc"
|
||||
elif [ -f "$HOME/.bashrc" ]; then
|
||||
_SHELL_PROFILE="$HOME/.bashrc"
|
||||
elif [ -f "$HOME/.profile" ]; then
|
||||
_SHELL_PROFILE="$HOME/.profile"
|
||||
fi
|
||||
|
||||
if [ -n "$_SHELL_PROFILE" ]; then
|
||||
if ! grep -q '\.local/bin' "$_SHELL_PROFILE" 2>/dev/null; then
|
||||
echo '' >> "$_SHELL_PROFILE"
|
||||
echo '# Added by Unsloth installer' >> "$_SHELL_PROFILE"
|
||||
echo 'export PATH="$HOME/.local/bin:$PATH"' >> "$_SHELL_PROFILE"
|
||||
step "path" "added ~/.local/bin to PATH in $_SHELL_PROFILE"
|
||||
if [ "$_STUDIO_HOME_REDIRECT" = "env" ]; then
|
||||
export PATH="$_LOCAL_BIN:$PATH"
|
||||
step "path" "exported $_LOCAL_BIN for this session (no rc-file append in env-override mode)"
|
||||
else
|
||||
_SHELL_PROFILE=""
|
||||
if [ -n "${ZSH_VERSION:-}" ] || [ "$(basename "${SHELL:-}")" = "zsh" ]; then
|
||||
_SHELL_PROFILE="$HOME/.zshrc"
|
||||
elif [ -f "$HOME/.bashrc" ]; then
|
||||
_SHELL_PROFILE="$HOME/.bashrc"
|
||||
elif [ -f "$HOME/.profile" ]; then
|
||||
_SHELL_PROFILE="$HOME/.profile"
|
||||
fi
|
||||
if [ -n "$_SHELL_PROFILE" ]; then
|
||||
if ! grep -q '\.local/bin' "$_SHELL_PROFILE" 2>/dev/null; then
|
||||
echo '' >> "$_SHELL_PROFILE"
|
||||
echo '# Added by Unsloth installer' >> "$_SHELL_PROFILE"
|
||||
echo 'export PATH="$HOME/.local/bin:$PATH"' >> "$_SHELL_PROFILE"
|
||||
step "path" "added ~/.local/bin to PATH in $_SHELL_PROFILE"
|
||||
fi
|
||||
fi
|
||||
export PATH="$_LOCAL_BIN:$PATH"
|
||||
fi
|
||||
export PATH="$_LOCAL_BIN:$PATH"
|
||||
;;
|
||||
esac
|
||||
|
||||
# Non-Tauri installs keep shortcuts even if setup reports failure.
|
||||
# create_studio_shortcuts gates persistent menu shortcuts on env-mode;
|
||||
# launcher + studio.conf + icon are always written.
|
||||
if [ "$TAURI_MODE" != true ]; then
|
||||
create_studio_shortcuts "$VENV_ABS_BIN/unsloth" "$OS"
|
||||
fi
|
||||
|
|
@ -1883,10 +2213,21 @@ if [ -t 1 ]; then
|
|||
esac
|
||||
else
|
||||
step "launch" "manual commands:"
|
||||
substep "unsloth studio -p 8888"
|
||||
substep "or activate env first:"
|
||||
substep "source ${VENV_DIR}/bin/activate"
|
||||
substep "unsloth studio -p 8888"
|
||||
# Single-quote-escape so paths with spaces / apostrophes copy-paste cleanly.
|
||||
_li_shim_q="'$(printf '%s' "${_LOCAL_BIN}/unsloth" | sed "s/'/'\\\\''/g")'"
|
||||
_li_act_q="'$(printf '%s' "${VENV_DIR}/bin/activate" | sed "s/'/'\\\\''/g")'"
|
||||
if [ "$_STUDIO_HOME_REDIRECT" = "env" ]; then
|
||||
# Env-mode skips the rc PATH append, so print the absolute shim path.
|
||||
substep "$_li_shim_q studio -p 8888"
|
||||
substep "or activate env first:"
|
||||
substep "source $_li_act_q"
|
||||
substep "unsloth studio -p 8888"
|
||||
else
|
||||
substep "unsloth studio -p 8888"
|
||||
substep "or activate env first:"
|
||||
substep "source $_li_act_q"
|
||||
substep "unsloth studio -p 8888"
|
||||
fi
|
||||
substep "(add -H 0.0.0.0 to allow network / cloud access)"
|
||||
echo ""
|
||||
fi
|
||||
|
|
|
|||
|
|
@ -89,7 +89,7 @@ huggingfacenotorch = [
|
|||
]
|
||||
huggingface = [
|
||||
"unsloth[huggingfacenotorch]",
|
||||
"unsloth_zoo>=2026.5.1",
|
||||
"unsloth_zoo>=2026.4.8",
|
||||
"torchvision",
|
||||
"unsloth[triton]",
|
||||
]
|
||||
|
|
@ -579,7 +579,7 @@ colab-ampere-torch220 = [
|
|||
"flash-attn>=2.6.3 ; ('linux' in sys_platform)",
|
||||
]
|
||||
colab-new = [
|
||||
"unsloth_zoo>=2026.5.1",
|
||||
"unsloth_zoo>=2026.4.8",
|
||||
"packaging",
|
||||
"tyro",
|
||||
"transformers>=4.51.3,!=4.52.0,!=4.52.1,!=4.52.2,!=4.52.3,!=4.53.0,!=4.54.0,!=4.55.0,!=4.55.1,!=4.57.0,!=4.57.4,!=4.57.5,!=5.0.0,!=5.1.0,<=5.5.0",
|
||||
|
|
|
|||
|
|
@ -9,16 +9,14 @@ Export backend - handles model exporting in various formats
|
|||
import glob
|
||||
import json
|
||||
import structlog
|
||||
import tempfile
|
||||
from loggers import get_logger
|
||||
import os
|
||||
import shutil
|
||||
from pathlib import Path
|
||||
from typing import Optional, Tuple, List
|
||||
from peft import PeftModel, PeftModelForCausalLM
|
||||
from unsloth import FastLanguageModel, FastVisionModel
|
||||
from unsloth import FastLanguageModel, FastVisionModel, _IS_MLX
|
||||
from huggingface_hub import HfApi, ModelCard
|
||||
from transformers.modeling_utils import PushToHubMixin
|
||||
import torch
|
||||
from utils.hardware import clear_gpu_cache
|
||||
|
||||
from utils.models import is_vision_model, get_base_model_from_lora
|
||||
|
|
@ -26,6 +24,12 @@ from utils.models.model_config import detect_audio_type
|
|||
from utils.paths import ensure_dir, outputs_root, resolve_export_dir, resolve_output_dir
|
||||
from core.inference import get_inference_backend
|
||||
|
||||
# GPU-only imports — guarded for Apple Silicon where these aren't needed
|
||||
if not _IS_MLX:
|
||||
from peft import PeftModel, PeftModelForCausalLM
|
||||
from transformers.modeling_utils import PushToHubMixin
|
||||
import torch
|
||||
|
||||
logger = get_logger(__name__)
|
||||
|
||||
_LLAMA_CPP_SCRIPTS_WARNING_EMITTED = False
|
||||
|
|
@ -225,7 +229,7 @@ class ExportBackend:
|
|||
model, tokenizer = FastModel.from_pretrained(
|
||||
model_name = checkpoint_path,
|
||||
max_seq_length = max_seq_length,
|
||||
dtype = torch.float32,
|
||||
dtype = None if _IS_MLX else torch.float32,
|
||||
load_in_4bit = False,
|
||||
trust_remote_code = trust_remote_code,
|
||||
)
|
||||
|
|
@ -262,8 +266,12 @@ class ExportBackend:
|
|||
trust_remote_code = trust_remote_code,
|
||||
)
|
||||
|
||||
# Check if PEFT model
|
||||
self.is_peft = isinstance(model, (PeftModel, PeftModelForCausalLM))
|
||||
# Check if PEFT / LoRA model
|
||||
if _IS_MLX:
|
||||
# MLX doesn't use PeftModel — detect LoRA via adapter_config.json
|
||||
self.is_peft = adapter_config.exists()
|
||||
else:
|
||||
self.is_peft = isinstance(model, (PeftModel, PeftModelForCausalLM))
|
||||
|
||||
# Store loaded model
|
||||
self.current_model = model
|
||||
|
|
@ -325,9 +333,7 @@ class ExportBackend:
|
|||
private: Whether to make the repo private
|
||||
|
||||
Returns:
|
||||
Tuple of (success, message, output_path). output_path is the
|
||||
resolved absolute on-disk directory of the saved model when
|
||||
``save_directory`` was set, else None.
|
||||
Tuple of (success: bool, message: str, output_path: Optional[str])
|
||||
"""
|
||||
if not self.current_model or not self.current_tokenizer:
|
||||
return False, "No model loaded. Please select a checkpoint first.", None
|
||||
|
|
@ -341,14 +347,17 @@ class ExportBackend:
|
|||
|
||||
output_path: Optional[str] = None
|
||||
try:
|
||||
# Determine save method
|
||||
if format_type == "4-bit (FP4)":
|
||||
save_method = "merged_4bit_forced"
|
||||
elif self._audio_type == "whisper":
|
||||
# Whisper uses save_method=None for local 16-bit merged save
|
||||
save_method = None
|
||||
else: # 16-bit (FP16)
|
||||
save_method = "merged_16bit"
|
||||
if _IS_MLX:
|
||||
mlx_save_method = (
|
||||
"merged_4bit" if format_type == "4-bit (FP4)" else "merged_16bit"
|
||||
)
|
||||
else:
|
||||
if format_type == "4-bit (FP4)":
|
||||
save_method = "merged_4bit_forced"
|
||||
elif self._audio_type == "whisper":
|
||||
save_method = None
|
||||
else:
|
||||
save_method = "merged_16bit"
|
||||
|
||||
# Save locally if requested
|
||||
if save_directory:
|
||||
|
|
@ -356,11 +365,17 @@ class ExportBackend:
|
|||
logger.info(f"Saving merged model locally to: {save_directory}")
|
||||
ensure_dir(Path(save_directory))
|
||||
|
||||
self.current_model.save_pretrained_merged(
|
||||
save_directory, self.current_tokenizer, save_method = save_method
|
||||
)
|
||||
if _IS_MLX:
|
||||
self.current_model.save_pretrained_merged(
|
||||
save_directory,
|
||||
self.current_tokenizer,
|
||||
save_method = mlx_save_method,
|
||||
)
|
||||
else:
|
||||
self.current_model.save_pretrained_merged(
|
||||
save_directory, self.current_tokenizer, save_method = save_method
|
||||
)
|
||||
|
||||
# Write export metadata so the Chat page can identify the base model
|
||||
self._write_export_metadata(save_directory)
|
||||
logger.info(f"Model saved successfully to {save_directory}")
|
||||
output_path = str(Path(save_directory).resolve())
|
||||
|
|
@ -376,17 +391,40 @@ class ExportBackend:
|
|||
|
||||
logger.info(f"Pushing merged model to Hub: {repo_id}")
|
||||
|
||||
# Whisper uses save_method=None for local but "merged_16bit" for hub push
|
||||
hub_save_method = (
|
||||
save_method if save_method is not None else "merged_16bit"
|
||||
)
|
||||
self.current_model.push_to_hub_merged(
|
||||
repo_id,
|
||||
self.current_tokenizer,
|
||||
save_method = hub_save_method,
|
||||
token = hf_token,
|
||||
private = private,
|
||||
)
|
||||
if _IS_MLX:
|
||||
if save_directory:
|
||||
self.current_model.push_to_hub_merged(
|
||||
repo_id,
|
||||
self.current_tokenizer,
|
||||
save_directory = save_directory,
|
||||
token = hf_token,
|
||||
private = private,
|
||||
)
|
||||
else:
|
||||
with tempfile.TemporaryDirectory() as tmp_dir:
|
||||
self.current_model.save_pretrained_merged(
|
||||
tmp_dir,
|
||||
self.current_tokenizer,
|
||||
save_method = mlx_save_method,
|
||||
)
|
||||
self.current_model.push_to_hub_merged(
|
||||
repo_id,
|
||||
self.current_tokenizer,
|
||||
save_directory = tmp_dir,
|
||||
token = hf_token,
|
||||
private = private,
|
||||
)
|
||||
else:
|
||||
hub_save_method = (
|
||||
save_method if save_method is not None else "merged_16bit"
|
||||
)
|
||||
self.current_model.push_to_hub_merged(
|
||||
repo_id,
|
||||
self.current_tokenizer,
|
||||
save_method = hub_save_method,
|
||||
token = hf_token,
|
||||
private = private,
|
||||
)
|
||||
logger.info(f"Model pushed successfully to {repo_id}")
|
||||
|
||||
return True, "Model exported successfully", output_path
|
||||
|
|
@ -411,9 +449,7 @@ class ExportBackend:
|
|||
Export base model (for non-PEFT models).
|
||||
|
||||
Returns:
|
||||
Tuple of (success, message, output_path). output_path is the
|
||||
resolved absolute on-disk directory of the saved model when
|
||||
``save_directory`` was set, else None.
|
||||
Tuple of (success: bool, message: str, output_path: Optional[str])
|
||||
"""
|
||||
if not self.current_model or not self.current_tokenizer:
|
||||
return False, "No model loaded. Please select a checkpoint first.", None
|
||||
|
|
@ -433,8 +469,16 @@ class ExportBackend:
|
|||
logger.info(f"Saving base model locally to: {save_directory}")
|
||||
ensure_dir(Path(save_directory))
|
||||
|
||||
self.current_model.save_pretrained(save_directory)
|
||||
self.current_tokenizer.save_pretrained(save_directory)
|
||||
if _IS_MLX:
|
||||
# MLX: save_pretrained_merged handles non-LoRA models too
|
||||
# (fuse() is a no-op when there are no LoRA layers)
|
||||
self.current_model.save_pretrained_merged(
|
||||
save_directory,
|
||||
self.current_tokenizer,
|
||||
)
|
||||
else:
|
||||
self.current_model.save_pretrained(save_directory)
|
||||
self.current_tokenizer.save_pretrained(save_directory)
|
||||
|
||||
# Write export metadata so the Chat page can identify the base model
|
||||
self._write_export_metadata(save_directory)
|
||||
|
|
@ -452,44 +496,73 @@ class ExportBackend:
|
|||
|
||||
logger.info(f"Pushing base model to Hub: {repo_id}")
|
||||
|
||||
# Get base model name from request or model config
|
||||
base_model = (
|
||||
base_model_id
|
||||
or self.current_model.config._name_or_path
|
||||
or "unknown"
|
||||
)
|
||||
|
||||
# Create repo
|
||||
hf_api = HfApi(token = hf_token)
|
||||
repo_id = PushToHubMixin._create_repo(
|
||||
PushToHubMixin,
|
||||
repo_id = repo_id,
|
||||
private = private,
|
||||
token = hf_token,
|
||||
)
|
||||
username = repo_id.split("/")[0]
|
||||
|
||||
# Create and push model card
|
||||
content = MODEL_CARD.format(
|
||||
username = username,
|
||||
base_model = base_model,
|
||||
model_type = self.current_model.config.model_type,
|
||||
method = "",
|
||||
extra = "unsloth",
|
||||
)
|
||||
card = ModelCard(content)
|
||||
card.push_to_hub(
|
||||
repo_id, token = hf_token, commit_message = "Unsloth Model Card"
|
||||
)
|
||||
|
||||
# Upload model files
|
||||
if save_directory:
|
||||
hf_api.upload_folder(
|
||||
folder_path = save_directory, repo_id = repo_id, repo_type = "model"
|
||||
)
|
||||
logger.info(f"Model pushed successfully to {repo_id}")
|
||||
if _IS_MLX:
|
||||
if save_directory:
|
||||
self.current_model.push_to_hub_merged(
|
||||
repo_id,
|
||||
self.current_tokenizer,
|
||||
save_directory = save_directory,
|
||||
token = hf_token,
|
||||
private = private,
|
||||
)
|
||||
else:
|
||||
with tempfile.TemporaryDirectory() as tmp_dir:
|
||||
self.current_model.save_pretrained_merged(
|
||||
tmp_dir,
|
||||
self.current_tokenizer,
|
||||
)
|
||||
self.current_model.push_to_hub_merged(
|
||||
repo_id,
|
||||
self.current_tokenizer,
|
||||
save_directory = tmp_dir,
|
||||
token = hf_token,
|
||||
private = private,
|
||||
)
|
||||
else:
|
||||
return False, "Local save directory required for Hub upload", None
|
||||
# Get base model name from request or model config
|
||||
base_model = (
|
||||
base_model_id
|
||||
or self.current_model.config._name_or_path
|
||||
or "unknown"
|
||||
)
|
||||
|
||||
# Create repo
|
||||
hf_api = HfApi(token = hf_token)
|
||||
repo_id = PushToHubMixin._create_repo(
|
||||
PushToHubMixin,
|
||||
repo_id = repo_id,
|
||||
private = private,
|
||||
token = hf_token,
|
||||
)
|
||||
username = repo_id.split("/")[0]
|
||||
|
||||
# Create and push model card
|
||||
content = MODEL_CARD.format(
|
||||
username = username,
|
||||
base_model = base_model,
|
||||
model_type = self.current_model.config.model_type,
|
||||
method = "",
|
||||
extra = "unsloth",
|
||||
)
|
||||
card = ModelCard(content)
|
||||
card.push_to_hub(
|
||||
repo_id, token = hf_token, commit_message = "Unsloth Model Card"
|
||||
)
|
||||
|
||||
# Upload model files
|
||||
if save_directory:
|
||||
hf_api.upload_folder(
|
||||
folder_path = save_directory,
|
||||
repo_id = repo_id,
|
||||
repo_type = "model",
|
||||
)
|
||||
logger.info(f"Model pushed successfully to {repo_id}")
|
||||
else:
|
||||
return (
|
||||
False,
|
||||
"Local save directory required for Hub upload",
|
||||
None,
|
||||
)
|
||||
|
||||
return True, "Model exported successfully", output_path
|
||||
|
||||
|
|
@ -519,9 +592,7 @@ class ExportBackend:
|
|||
hf_token: Hugging Face token
|
||||
|
||||
Returns:
|
||||
Tuple of (success, message, output_path). output_path is the
|
||||
resolved absolute on-disk directory containing the .gguf
|
||||
files when ``save_directory`` was set, else None.
|
||||
Tuple of (success: bool, message: str, output_path: Optional[str])
|
||||
"""
|
||||
if not self.current_model or not self.current_tokenizer:
|
||||
return False, "No model loaded. Please select a checkpoint first.", None
|
||||
|
|
@ -692,9 +763,7 @@ class ExportBackend:
|
|||
Export LoRA adapter only (not merged).
|
||||
|
||||
Returns:
|
||||
Tuple of (success, message, output_path). output_path is the
|
||||
resolved absolute on-disk directory of the saved adapter
|
||||
when ``save_directory`` was set, else None.
|
||||
Tuple of (success: bool, message: str, output_path: Optional[str])
|
||||
"""
|
||||
if not self.current_model or not self.current_tokenizer:
|
||||
return False, "No model loaded. Please select a checkpoint first.", None
|
||||
|
|
@ -710,8 +779,13 @@ class ExportBackend:
|
|||
logger.info(f"Saving LoRA adapter locally to: {save_directory}")
|
||||
ensure_dir(Path(save_directory))
|
||||
|
||||
self.current_model.save_pretrained(save_directory)
|
||||
self.current_tokenizer.save_pretrained(save_directory)
|
||||
if _IS_MLX:
|
||||
# MLX: save adapters.safetensors + tokenizer files
|
||||
self.current_model.save_lora_adapters(save_directory)
|
||||
self.current_tokenizer.save_pretrained(save_directory)
|
||||
else:
|
||||
self.current_model.save_pretrained(save_directory)
|
||||
self.current_tokenizer.save_pretrained(save_directory)
|
||||
logger.info(f"Adapter saved successfully to {save_directory}")
|
||||
output_path = str(Path(save_directory).resolve())
|
||||
|
||||
|
|
@ -726,10 +800,24 @@ class ExportBackend:
|
|||
|
||||
logger.info(f"Pushing LoRA adapter to Hub: {repo_id}")
|
||||
|
||||
self.current_model.push_to_hub(repo_id, token = hf_token, private = private)
|
||||
self.current_tokenizer.push_to_hub(
|
||||
repo_id, token = hf_token, private = private
|
||||
)
|
||||
if _IS_MLX:
|
||||
with tempfile.TemporaryDirectory() as tmp_dir:
|
||||
self.current_model.save_lora_adapters(tmp_dir)
|
||||
self.current_tokenizer.save_pretrained(tmp_dir)
|
||||
hf_api = HfApi(token = hf_token)
|
||||
hf_api.create_repo(repo_id, private = private, exist_ok = True)
|
||||
hf_api.upload_folder(
|
||||
folder_path = tmp_dir,
|
||||
repo_id = repo_id,
|
||||
repo_type = "model",
|
||||
)
|
||||
else:
|
||||
self.current_model.push_to_hub(
|
||||
repo_id, token = hf_token, private = private
|
||||
)
|
||||
self.current_tokenizer.push_to_hub(
|
||||
repo_id, token = hf_token, private = private
|
||||
)
|
||||
logger.info(f"Adapter pushed successfully to {repo_id}")
|
||||
|
||||
return True, "LoRA adapter exported successfully", output_path
|
||||
|
|
|
|||
|
|
@ -732,22 +732,46 @@ class LlamaCppBackend:
|
|||
if win_bin.is_file():
|
||||
return str(win_bin)
|
||||
|
||||
# 2–4. ~/.unsloth/llama.cpp (primary — setup.sh / setup.ps1 build here)
|
||||
unsloth_home = Path.home() / ".unsloth" / "llama.cpp"
|
||||
# Root dir (make builds copy binaries here)
|
||||
home_root = unsloth_home / binary_name
|
||||
if home_root.is_file():
|
||||
return str(home_root)
|
||||
# build/bin/ (cmake builds on Linux)
|
||||
home_linux = unsloth_home / "build" / "bin" / binary_name
|
||||
if home_linux.is_file():
|
||||
return str(home_linux)
|
||||
# 2-4. Match installer layout: env-mode -> $STUDIO_HOME/llama.cpp;
|
||||
# default/HOME-redirect -> ~/.unsloth/llama.cpp (sibling of studio).
|
||||
legacy_llama = Path.home() / ".unsloth" / "llama.cpp"
|
||||
try:
|
||||
from utils.paths.storage_roots import studio_root as _sr # noqa: WPS433
|
||||
|
||||
# 3. Windows MSVC build has Release subdir
|
||||
if sys.platform == "win32":
|
||||
home_win = unsloth_home / "build" / "bin" / "Release" / binary_name
|
||||
if home_win.is_file():
|
||||
return str(home_win)
|
||||
_resolved_sr = _sr()
|
||||
_legacy_studio = Path.home() / ".unsloth" / "studio"
|
||||
try:
|
||||
_is_legacy = _resolved_sr.resolve() == _legacy_studio.resolve()
|
||||
except (OSError, ValueError):
|
||||
_is_legacy = _resolved_sr == _legacy_studio
|
||||
if _is_legacy:
|
||||
search_roots = [legacy_llama]
|
||||
else:
|
||||
# why: _kill_orphaned_servers excludes the legacy root in custom
|
||||
# mode; discovery must match so we never spawn a server we then
|
||||
# refuse to clean up. UNSLOTH_LLAMA_CPP_PATH (handled earlier)
|
||||
# is the explicit way to share a build across roots.
|
||||
search_roots = [_resolved_sr / "llama.cpp"]
|
||||
except (ImportError, OSError, ValueError):
|
||||
search_roots = [legacy_llama]
|
||||
_seen_roots: set[str] = set()
|
||||
_unique_roots: list[Path] = []
|
||||
for r in search_roots:
|
||||
k = str(r)
|
||||
if k not in _seen_roots:
|
||||
_seen_roots.add(k)
|
||||
_unique_roots.append(r)
|
||||
for unsloth_home in _unique_roots:
|
||||
home_root = unsloth_home / binary_name
|
||||
if home_root.is_file():
|
||||
return str(home_root)
|
||||
home_linux = unsloth_home / "build" / "bin" / binary_name
|
||||
if home_linux.is_file():
|
||||
return str(home_linux)
|
||||
if sys.platform == "win32":
|
||||
home_win = unsloth_home / "build" / "bin" / "Release" / binary_name
|
||||
if home_win.is_file():
|
||||
return str(home_win)
|
||||
|
||||
# 5–6. Legacy: in-tree build (older setup.sh / setup.ps1 versions)
|
||||
project_root = Path(__file__).resolve().parents[4]
|
||||
|
|
@ -2592,8 +2616,27 @@ class LlamaCppBackend:
|
|||
# (binary must be *under* one of these)
|
||||
install_roots: list[Path] = []
|
||||
|
||||
# Primary install dir (setup.sh / prebuilt installer)
|
||||
install_roots.append(Path.home() / ".unsloth" / "llama.cpp")
|
||||
# Env-mode custom root (mirrors _find_llama_server_binary).
|
||||
_is_custom_root = False
|
||||
try:
|
||||
from utils.paths.storage_roots import studio_root as _sr # noqa: WPS433
|
||||
|
||||
_resolved_sr = _sr()
|
||||
_legacy_studio = Path.home() / ".unsloth" / "studio"
|
||||
try:
|
||||
_is_custom_root = _resolved_sr.resolve() != _legacy_studio.resolve()
|
||||
except (OSError, ValueError):
|
||||
_is_custom_root = _resolved_sr != _legacy_studio
|
||||
if _is_custom_root:
|
||||
install_roots.append(_resolved_sr / "llama.cpp")
|
||||
except (ImportError, OSError, ValueError):
|
||||
pass
|
||||
|
||||
# Primary install dir (default mode only). Env-mode skips this so
|
||||
# a custom-root Studio cannot kill a concurrent default-install
|
||||
# Studio's llama-server (same OS user, different install).
|
||||
if not _is_custom_root:
|
||||
install_roots.append(Path.home() / ".unsloth" / "llama.cpp")
|
||||
|
||||
# Legacy in-tree build dirs (older setup.sh versions)
|
||||
project_root = Path(__file__).resolve().parents[4]
|
||||
|
|
|
|||
395
studio/backend/core/inference/mlx_inference.py
Normal file
395
studio/backend/core/inference/mlx_inference.py
Normal file
|
|
@ -0,0 +1,395 @@
|
|||
# SPDX-License-Identifier: AGPL-3.0-only
|
||||
"""MLX inference backend for Apple Silicon.
|
||||
|
||||
Drop-in replacement for InferenceBackend — same interface, uses mlx-lm/mlx-vlm
|
||||
instead of torch/transformers for model loading and generation.
|
||||
"""
|
||||
|
||||
import threading
|
||||
from typing import Optional, Generator
|
||||
from loggers import get_logger
|
||||
|
||||
logger = get_logger(__name__)
|
||||
|
||||
|
||||
class MLXInferenceBackend:
|
||||
def __init__(self):
|
||||
self.models = {}
|
||||
self.active_model_name = None
|
||||
self.loading_models = set()
|
||||
self.loaded_local_models = []
|
||||
self.device = "mlx"
|
||||
self._generation_lock = threading.Lock()
|
||||
|
||||
# MLX state
|
||||
self._model = None
|
||||
self._tokenizer = None
|
||||
self._processor = None
|
||||
self._is_vlm = False
|
||||
self._config = {}
|
||||
|
||||
# Recorded for unload to release pinned memory back to the OS.
|
||||
self._memory_limits_applied = {}
|
||||
|
||||
def _configure_memory_limits(self):
|
||||
"""Apply Metal memory caps before loading a model.
|
||||
|
||||
Mirrors MLXTrainer._configure_memory_limits's defaults:
|
||||
memory_limit = 85% of recommended working-set,
|
||||
wired_limit = min(recommended, memory_limit). Recorded so unload
|
||||
can lower wired_limit back to release pinned RAM.
|
||||
"""
|
||||
import mlx.core as mx
|
||||
|
||||
if not mx.metal.is_available():
|
||||
return
|
||||
info = mx.device_info()
|
||||
rec_bytes = info.get("max_recommended_working_set_size")
|
||||
if not rec_bytes or rec_bytes <= 0:
|
||||
return
|
||||
rec_gb = rec_bytes / 1e9
|
||||
memory_limit_gb = rec_gb * 0.85
|
||||
wired_limit_gb = min(rec_gb, memory_limit_gb)
|
||||
mx.set_memory_limit(int(memory_limit_gb * 1e9))
|
||||
mx.set_wired_limit(int(wired_limit_gb * 1e9))
|
||||
self._memory_limits_applied = {
|
||||
"memory_limit_gb": memory_limit_gb,
|
||||
"wired_limit_gb": wired_limit_gb,
|
||||
"recommended_gb": rec_gb,
|
||||
}
|
||||
logger.info(
|
||||
"MLX memory caps: memory_limit=%.2f GB, wired_limit=%.2f GB",
|
||||
memory_limit_gb,
|
||||
wired_limit_gb,
|
||||
)
|
||||
|
||||
def load_model(
|
||||
self,
|
||||
config,
|
||||
max_seq_length = 2048,
|
||||
load_in_4bit = True,
|
||||
hf_token = None,
|
||||
trust_remote_code = False,
|
||||
gpu_ids = None,
|
||||
dtype = None,
|
||||
) -> bool:
|
||||
import mlx.core as mx
|
||||
|
||||
model_name = config.identifier if hasattr(config, "identifier") else str(config)
|
||||
is_vision = getattr(config, "is_vision", False)
|
||||
|
||||
if hf_token:
|
||||
import os
|
||||
|
||||
os.environ["HF_TOKEN"] = hf_token
|
||||
self._configure_memory_limits()
|
||||
|
||||
is_lora = getattr(config, "is_lora", False)
|
||||
|
||||
logger.info(
|
||||
"Loading %s via %s (is_lora=%s)",
|
||||
model_name,
|
||||
"mlx-vlm" if is_vision else "mlx-lm",
|
||||
is_lora,
|
||||
)
|
||||
|
||||
try:
|
||||
from unsloth_zoo.mlx_loader import FastMLXModel
|
||||
except ImportError as e:
|
||||
raise ImportError(
|
||||
"Unsloth: MLX inference requires unsloth-zoo with the MLX modules "
|
||||
"(unsloth_zoo.mlx_loader). Reinstall via install.sh on Apple Silicon."
|
||||
) from e
|
||||
|
||||
model, tokenizer_or_processor = FastMLXModel.from_pretrained(
|
||||
model_name,
|
||||
max_seq_length = max_seq_length,
|
||||
dtype = dtype,
|
||||
load_in_4bit = load_in_4bit,
|
||||
token = hf_token,
|
||||
trust_remote_code = trust_remote_code,
|
||||
text_only = False if is_vision else True,
|
||||
)
|
||||
|
||||
if is_vision:
|
||||
processor = tokenizer_or_processor
|
||||
self._model = model
|
||||
self._processor = processor
|
||||
self._tokenizer = getattr(processor, "tokenizer", processor)
|
||||
self._is_vlm = True
|
||||
else:
|
||||
tokenizer = tokenizer_or_processor
|
||||
self._model = model
|
||||
self._tokenizer = tokenizer
|
||||
self._processor = None
|
||||
self._is_vlm = False
|
||||
|
||||
self.active_model_name = model_name
|
||||
self.models[model_name] = {
|
||||
"model": self._model,
|
||||
"tokenizer": self._tokenizer,
|
||||
"processor": self._processor,
|
||||
"is_vision": is_vision,
|
||||
"is_lora": getattr(config, "is_lora", False),
|
||||
"is_audio": False,
|
||||
"audio_type": None,
|
||||
"has_audio_input": False,
|
||||
}
|
||||
|
||||
logger.info("Model %s loaded successfully", model_name)
|
||||
return True
|
||||
|
||||
def unload_model(self, model_name: str) -> bool:
|
||||
import mlx.core as mx
|
||||
import gc
|
||||
|
||||
if model_name in self.models:
|
||||
del self.models[model_name]
|
||||
self._model = None
|
||||
self._tokenizer = None
|
||||
self._processor = None
|
||||
if self.active_model_name == model_name:
|
||||
self.active_model_name = None
|
||||
gc.collect()
|
||||
mx.clear_cache()
|
||||
|
||||
if mx.metal.is_available() and self._memory_limits_applied and not self.models:
|
||||
try:
|
||||
mx.set_wired_limit(0)
|
||||
logger.info("MLX wired_limit released back to OS on unload")
|
||||
except Exception as e:
|
||||
logger.warning("Failed to release wired_limit: %s", e)
|
||||
self._memory_limits_applied = {}
|
||||
logger.info("Model %s unloaded", model_name)
|
||||
return True
|
||||
|
||||
def generate_chat_response(
|
||||
self,
|
||||
messages,
|
||||
system_prompt = "",
|
||||
image = None,
|
||||
temperature = 0.7,
|
||||
top_p = 0.9,
|
||||
top_k = 40,
|
||||
min_p = 0.0,
|
||||
max_new_tokens = 256,
|
||||
repetition_penalty = 1.0,
|
||||
cancel_event = None,
|
||||
) -> Generator[str, None, None]:
|
||||
if self._model is None:
|
||||
raise RuntimeError("No model loaded")
|
||||
|
||||
# Build messages with system prompt
|
||||
full_messages = []
|
||||
if system_prompt:
|
||||
full_messages.append({"role": "system", "content": system_prompt})
|
||||
full_messages.extend(messages)
|
||||
|
||||
# Inject image into the last user message for VLM
|
||||
if self._is_vlm and image is not None:
|
||||
for msg in reversed(full_messages):
|
||||
if msg.get("role") == "user":
|
||||
content = msg.get("content", "")
|
||||
if isinstance(content, str):
|
||||
msg["content"] = [
|
||||
{"type": "image"},
|
||||
{"type": "text", "text": content},
|
||||
]
|
||||
elif isinstance(content, list):
|
||||
# Prepend image if not already there
|
||||
has_image = any(
|
||||
p.get("type") == "image"
|
||||
for p in content
|
||||
if isinstance(p, dict)
|
||||
)
|
||||
if not has_image:
|
||||
content.insert(0, {"type": "image"})
|
||||
break
|
||||
|
||||
if self._is_vlm:
|
||||
yield from self._generate_vlm(
|
||||
full_messages,
|
||||
image,
|
||||
temperature,
|
||||
top_p,
|
||||
top_k,
|
||||
min_p,
|
||||
max_new_tokens,
|
||||
repetition_penalty,
|
||||
cancel_event,
|
||||
)
|
||||
else:
|
||||
yield from self._generate_text(
|
||||
full_messages,
|
||||
temperature,
|
||||
top_p,
|
||||
top_k,
|
||||
min_p,
|
||||
max_new_tokens,
|
||||
repetition_penalty,
|
||||
cancel_event,
|
||||
)
|
||||
|
||||
def _generate_text(
|
||||
self,
|
||||
messages,
|
||||
temperature,
|
||||
top_p,
|
||||
top_k,
|
||||
min_p,
|
||||
max_new_tokens,
|
||||
repetition_penalty,
|
||||
cancel_event,
|
||||
):
|
||||
from mlx_lm import stream_generate
|
||||
from mlx_lm.sample_utils import make_sampler, make_logits_processors
|
||||
|
||||
prompt = self._tokenizer.apply_chat_template(
|
||||
messages,
|
||||
tokenize = False,
|
||||
add_generation_prompt = True,
|
||||
)
|
||||
if prompt is None:
|
||||
raise RuntimeError(
|
||||
"apply_chat_template returned None — tokenizer may be incompatible"
|
||||
)
|
||||
|
||||
sampler = make_sampler(
|
||||
temp = temperature,
|
||||
top_p = top_p,
|
||||
top_k = int(top_k or 0),
|
||||
min_p = float(min_p or 0.0),
|
||||
min_tokens_to_keep = 1,
|
||||
)
|
||||
# Only build a logits processor when we actually have a non-trivial
|
||||
# repetition penalty (1.0 is the no-op value).
|
||||
logits_processors = None
|
||||
if repetition_penalty is not None and float(repetition_penalty) not in (
|
||||
0.0,
|
||||
1.0,
|
||||
):
|
||||
logits_processors = make_logits_processors(
|
||||
repetition_penalty = float(repetition_penalty),
|
||||
)
|
||||
|
||||
token_ids = []
|
||||
logger.info(
|
||||
"Generating: prompt_len=%d, max_tokens=%d, model=%s, tokenizer=%s",
|
||||
len(prompt),
|
||||
max_new_tokens,
|
||||
type(self._model).__name__,
|
||||
type(self._tokenizer).__name__,
|
||||
)
|
||||
with self._generation_lock:
|
||||
try:
|
||||
gen_kwargs = dict(
|
||||
prompt = prompt,
|
||||
max_tokens = max_new_tokens,
|
||||
sampler = sampler,
|
||||
)
|
||||
if logits_processors is not None:
|
||||
gen_kwargs["logits_processors"] = logits_processors
|
||||
for response in stream_generate(
|
||||
self._model,
|
||||
self._tokenizer,
|
||||
**gen_kwargs,
|
||||
):
|
||||
token_ids.append(response.token)
|
||||
# Decode full sequence with skip_special_tokens — same as GPU
|
||||
cumulative = self._tokenizer.decode(
|
||||
token_ids,
|
||||
skip_special_tokens = True,
|
||||
)
|
||||
yield cumulative
|
||||
|
||||
if cancel_event and cancel_event.is_set():
|
||||
break
|
||||
except Exception as e:
|
||||
import traceback
|
||||
|
||||
logger.error("stream_generate failed:\n%s", traceback.format_exc())
|
||||
raise
|
||||
|
||||
def _generate_vlm(
|
||||
self,
|
||||
messages,
|
||||
image,
|
||||
temperature,
|
||||
top_p,
|
||||
top_k,
|
||||
min_p,
|
||||
max_new_tokens,
|
||||
repetition_penalty,
|
||||
cancel_event,
|
||||
):
|
||||
from mlx_vlm import stream_generate as vlm_stream
|
||||
|
||||
# Apply chat template
|
||||
chat_fn = getattr(self._processor, "apply_chat_template", None)
|
||||
if (
|
||||
chat_fn is None
|
||||
or not hasattr(self._processor, "chat_template")
|
||||
or self._processor.chat_template is None
|
||||
):
|
||||
tok = getattr(self._processor, "tokenizer", self._processor)
|
||||
chat_fn = tok.apply_chat_template
|
||||
|
||||
prompt = chat_fn(messages, tokenize = False, add_generation_prompt = True)
|
||||
|
||||
# For VLM: always use mlx_vlm's stream_generate which handles
|
||||
# pixel_values properly (passes None for text-only, image for VLM)
|
||||
images = [image] if image is not None else None
|
||||
|
||||
cumulative = ""
|
||||
logger.info(
|
||||
"VLM generating: prompt_len=%d, has_image=%s",
|
||||
len(prompt),
|
||||
image is not None,
|
||||
)
|
||||
# mlx_vlm.stream_generate forwards **kwargs into generate_step, which
|
||||
# accepts temp/top_p/top_k/repetition_penalty (and builds the sampler
|
||||
# + logits_processors internally). Pass them through.
|
||||
# NOTE: mlx_vlm.generate_step expects ``temperature=`` (long form) —
|
||||
# passing ``temp=`` silently falls into **kwargs and is ignored,
|
||||
# leaving generation stuck at the default 0.0 (greedy).
|
||||
vlm_kwargs = dict(
|
||||
max_tokens = max_new_tokens,
|
||||
temperature = temperature,
|
||||
top_p = top_p,
|
||||
top_k = int(top_k or 0),
|
||||
min_p = float(min_p or 0.0),
|
||||
)
|
||||
if repetition_penalty is not None and float(repetition_penalty) not in (
|
||||
0.0,
|
||||
1.0,
|
||||
):
|
||||
vlm_kwargs["repetition_penalty"] = float(repetition_penalty)
|
||||
|
||||
with self._generation_lock:
|
||||
for response in vlm_stream(
|
||||
self._model,
|
||||
self._processor,
|
||||
prompt,
|
||||
images,
|
||||
**vlm_kwargs,
|
||||
):
|
||||
token_text = (
|
||||
response.text if hasattr(response, "text") else str(response)
|
||||
)
|
||||
cumulative += token_text
|
||||
yield cumulative
|
||||
if cancel_event and cancel_event.is_set():
|
||||
break
|
||||
|
||||
def generate_with_adapter_control(
|
||||
self, use_adapter = None, cancel_event = None, **gen_kwargs
|
||||
) -> Generator[str, None, None]:
|
||||
# MLX LoRA adapter toggling not yet supported — generate normally
|
||||
yield from self.generate_chat_response(cancel_event = cancel_event, **gen_kwargs)
|
||||
|
||||
def reset_generation_state(self):
|
||||
import mlx.core as mx
|
||||
import gc
|
||||
|
||||
gc.collect()
|
||||
mx.clear_cache()
|
||||
|
|
@ -663,6 +663,98 @@ def run_inference_process(
|
|||
|
||||
model_name = config["model_name"]
|
||||
|
||||
# ── 0. MLX fast-path — skip torch/transformers entirely ──
|
||||
backend_path = str(Path(__file__).resolve().parent.parent.parent)
|
||||
if backend_path not in sys.path:
|
||||
sys.path.insert(0, backend_path)
|
||||
|
||||
from utils.hardware import hardware as _hw
|
||||
|
||||
_hw.detect_hardware()
|
||||
if _hw.DEVICE == _hw.DeviceType.MLX:
|
||||
try:
|
||||
_activate_transformers_version(model_name)
|
||||
except Exception:
|
||||
pass
|
||||
try:
|
||||
from core.inference.mlx_inference import MLXInferenceBackend
|
||||
|
||||
backend = MLXInferenceBackend()
|
||||
_send_response(
|
||||
resp_queue,
|
||||
{"type": "status", "message": "Loading model...", "ts": time.time()},
|
||||
)
|
||||
_handle_load(backend, config, resp_queue)
|
||||
except Exception as exc:
|
||||
_send_response(
|
||||
resp_queue,
|
||||
{
|
||||
"type": "error",
|
||||
"error": f"MLX inference init failed: {exc}",
|
||||
"stack": traceback.format_exc(limit = 20),
|
||||
"ts": time.time(),
|
||||
},
|
||||
)
|
||||
return
|
||||
|
||||
# Enter same command loop as GPU path
|
||||
logger.info("MLX inference subprocess ready, entering command loop")
|
||||
while True:
|
||||
try:
|
||||
cmd = cmd_queue.get(timeout = 1.0)
|
||||
except _queue.Empty:
|
||||
continue
|
||||
except (EOFError, OSError):
|
||||
return
|
||||
if cmd is None:
|
||||
continue
|
||||
cmd_type = cmd.get("type", "")
|
||||
try:
|
||||
if cmd_type == "generate":
|
||||
cancel_event.clear()
|
||||
_handle_generate(backend, cmd, resp_queue, cancel_event)
|
||||
elif cmd_type == "load":
|
||||
if backend.active_model_name:
|
||||
backend.unload_model(backend.active_model_name)
|
||||
_handle_load(backend, cmd, resp_queue)
|
||||
elif cmd_type == "unload":
|
||||
_handle_unload(backend, cmd, resp_queue)
|
||||
elif cmd_type == "cancel":
|
||||
cancel_event.set()
|
||||
elif cmd_type == "reset":
|
||||
cancel_event.set()
|
||||
backend.reset_generation_state()
|
||||
_send_response(resp_queue, {"type": "reset_ack", "ts": time.time()})
|
||||
elif cmd_type == "status":
|
||||
_send_response(
|
||||
resp_queue,
|
||||
{
|
||||
"type": "status_response",
|
||||
"active_model": backend.active_model_name,
|
||||
"models": {
|
||||
k: {kk: vv for kk, vv in v.items() if kk != "model"}
|
||||
for k, v in backend.models.items()
|
||||
},
|
||||
"loading": list(backend.loading_models),
|
||||
"ts": time.time(),
|
||||
},
|
||||
)
|
||||
elif cmd_type == "shutdown":
|
||||
return
|
||||
except Exception as exc:
|
||||
logger.error("MLX command error (%s): %s", cmd_type, exc)
|
||||
_send_response(
|
||||
resp_queue,
|
||||
{
|
||||
"type": "gen_error" if cmd_type == "generate" else "error",
|
||||
"request_id": cmd.get("request_id"),
|
||||
"error": str(exc),
|
||||
"stack": traceback.format_exc(limit = 20),
|
||||
"ts": time.time(),
|
||||
},
|
||||
)
|
||||
return
|
||||
|
||||
# ── 1. Activate correct transformers version BEFORE any ML imports ──
|
||||
try:
|
||||
_activate_transformers_version(model_name)
|
||||
|
|
|
|||
|
|
@ -62,6 +62,7 @@ class TrainingProgress:
|
|||
grad_norm: Optional[float] = None
|
||||
num_tokens: Optional[int] = None
|
||||
eval_loss: Optional[float] = None
|
||||
peak_memory_gb: Optional[float] = None
|
||||
|
||||
|
||||
class TrainingBackend:
|
||||
|
|
@ -199,21 +200,27 @@ class TrainingBackend:
|
|||
config["load_in_4bit"] = False
|
||||
|
||||
# Spawn subprocess — use locals so state is untouched on failure
|
||||
resolved_gpu_ids, gpu_selection = prepare_gpu_selection(
|
||||
kwargs.get("gpu_ids"),
|
||||
model_name = config["model_name"],
|
||||
hf_token = config["hf_token"] or None,
|
||||
training_type = config["training_type"],
|
||||
load_in_4bit = config["load_in_4bit"],
|
||||
batch_size = config.get("batch_size", 4),
|
||||
max_seq_length = config.get("max_seq_length", 2048),
|
||||
lora_rank = config.get("lora_r", 16),
|
||||
target_modules = config.get("target_modules"),
|
||||
gradient_checkpointing = config.get("gradient_checkpointing", "unsloth"),
|
||||
optimizer = config.get("optim", "adamw_8bit"),
|
||||
)
|
||||
config["resolved_gpu_ids"] = resolved_gpu_ids
|
||||
config["gpu_selection"] = gpu_selection
|
||||
from utils.hardware import hardware as _hw
|
||||
|
||||
if _hw.DEVICE == _hw.DeviceType.MLX:
|
||||
config["resolved_gpu_ids"] = None
|
||||
config["gpu_selection"] = None
|
||||
else:
|
||||
resolved_gpu_ids, gpu_selection = prepare_gpu_selection(
|
||||
kwargs.get("gpu_ids"),
|
||||
model_name = config["model_name"],
|
||||
hf_token = config["hf_token"] or None,
|
||||
training_type = config["training_type"],
|
||||
load_in_4bit = config["load_in_4bit"],
|
||||
batch_size = config.get("batch_size", 4),
|
||||
max_seq_length = config.get("max_seq_length", 2048),
|
||||
lora_rank = config.get("lora_r", 16),
|
||||
target_modules = config.get("target_modules"),
|
||||
gradient_checkpointing = config.get("gradient_checkpointing", "unsloth"),
|
||||
optimizer = config.get("optim", "adamw_8bit"),
|
||||
)
|
||||
config["resolved_gpu_ids"] = resolved_gpu_ids
|
||||
config["gpu_selection"] = gpu_selection
|
||||
|
||||
from .worker import run_training_process
|
||||
|
||||
|
|
@ -512,6 +519,12 @@ class TrainingBackend:
|
|||
self._progress.grad_norm = event.get("grad_norm")
|
||||
self._progress.num_tokens = event.get("num_tokens")
|
||||
self._progress.eval_loss = event.get("eval_loss")
|
||||
_peak = event.get("peak_memory_gb")
|
||||
if _peak is not None:
|
||||
try:
|
||||
self._progress.peak_memory_gb = float(_peak)
|
||||
except (TypeError, ValueError):
|
||||
pass
|
||||
self._progress.is_training = True
|
||||
status = event.get("status_message", "")
|
||||
if status:
|
||||
|
|
|
|||
|
|
@ -338,6 +338,594 @@ def _activate_transformers_version(model_name: str) -> None:
|
|||
activate_transformers_for_subprocess(model_name)
|
||||
|
||||
|
||||
def _adapt_for_mlx_vlm(items):
|
||||
"""Adapt GPU-path VLM dataset output for mlx-vlm consumption.
|
||||
|
||||
The GPU path embeds PIL images inside messages content as
|
||||
{"type": "image", "image": PIL_Image}. mlx-vlm's prepare_inputs
|
||||
needs images at top-level to produce pixel_values — regardless of
|
||||
model type. Extract them and leave bare {"type": "image"} placeholders.
|
||||
"""
|
||||
adapted = []
|
||||
for item in items:
|
||||
images = []
|
||||
messages = []
|
||||
for msg in item.get("messages", []):
|
||||
content = msg.get("content", "")
|
||||
if isinstance(content, list):
|
||||
new_content = []
|
||||
for part in content:
|
||||
if isinstance(part, dict) and part.get("type") == "image":
|
||||
img = part.get("image")
|
||||
if img is not None:
|
||||
images.append(img)
|
||||
new_content.append({"type": "image"})
|
||||
else:
|
||||
new_content.append(part)
|
||||
messages.append({"role": msg["role"], "content": new_content})
|
||||
else:
|
||||
messages.append(msg)
|
||||
out = {"messages": messages}
|
||||
if images:
|
||||
out["image"] = images[0] if len(images) == 1 else images
|
||||
elif "image" in item:
|
||||
out["image"] = item["image"]
|
||||
elif "images" in item:
|
||||
out["images"] = item["images"]
|
||||
adapted.append(out)
|
||||
return adapted
|
||||
|
||||
|
||||
_MLX_STUDIO_OPTIM_MAP = {
|
||||
"adamw_8bit": "adamw",
|
||||
"paged_adamw_8bit": "adamw",
|
||||
"adamw_bnb_8bit": "adamw",
|
||||
"paged_adamw_32bit": "adamw",
|
||||
"adamw_torch": "adamw",
|
||||
"adamw_torch_fused": "adamw",
|
||||
"adamw": "adamw",
|
||||
"adafactor": "adafactor",
|
||||
"sgd": "sgd",
|
||||
"adam": "adam",
|
||||
"muon": "muon",
|
||||
"lion": "lion",
|
||||
}
|
||||
_MLX_STUDIO_LR_SCHEDULERS = {"linear", "cosine", "constant"}
|
||||
|
||||
|
||||
def _normalize_mlx_studio_optimizer(value):
|
||||
raw = str(value or "adamw_8bit").strip().lower()
|
||||
try:
|
||||
return _MLX_STUDIO_OPTIM_MAP[raw]
|
||||
except KeyError:
|
||||
supported = ", ".join(sorted(_MLX_STUDIO_OPTIM_MAP))
|
||||
raise ValueError(
|
||||
f"Unsupported optimizer for MLX training: {value!r}. "
|
||||
f"Supported values: {supported}."
|
||||
)
|
||||
|
||||
|
||||
def _normalize_mlx_studio_scheduler(value):
|
||||
raw = str(value or "linear").strip().lower()
|
||||
if raw not in _MLX_STUDIO_LR_SCHEDULERS:
|
||||
supported = ", ".join(sorted(_MLX_STUDIO_LR_SCHEDULERS))
|
||||
raise ValueError(
|
||||
f"Unsupported LR scheduler for MLX training: {value!r}. "
|
||||
f"Supported values: {supported}."
|
||||
)
|
||||
return raw
|
||||
|
||||
|
||||
def _run_mlx_training(event_queue, stop_queue, config):
|
||||
"""Self-contained MLX training path for Apple Silicon.
|
||||
|
||||
Uses MLXTrainer from unsloth_zoo directly -- no torch/SFTTrainer needed.
|
||||
Mirrors the event_queue protocol so the parent process pump works unchanged.
|
||||
"""
|
||||
import time
|
||||
import gc
|
||||
import math
|
||||
import threading
|
||||
import queue as _queue
|
||||
from pathlib import Path
|
||||
|
||||
def _send(event_type, **kwargs):
|
||||
if event_type == "status" and "message" not in kwargs:
|
||||
sm = kwargs.get("status_message")
|
||||
if sm is not None:
|
||||
kwargs["message"] = sm
|
||||
event_queue.put({"type": event_type, "ts": time.time(), **kwargs})
|
||||
|
||||
_send("status", status_message = "Loading MLX libraries...")
|
||||
|
||||
import mlx.core as mx
|
||||
|
||||
try:
|
||||
from unsloth_zoo.mlx_loader import FastMLXModel
|
||||
from unsloth_zoo.mlx_trainer import (
|
||||
MLXTrainer,
|
||||
MLXTrainingConfig,
|
||||
train_on_responses_only,
|
||||
)
|
||||
except ImportError as e:
|
||||
raise ImportError(
|
||||
"Unsloth: MLX training requires unsloth-zoo with the MLX modules "
|
||||
"(unsloth_zoo.mlx_loader / unsloth_zoo.mlx_trainer). Reinstall via "
|
||||
"install.sh on Apple Silicon."
|
||||
) from e
|
||||
from datasets import load_dataset
|
||||
|
||||
if mx.metal.is_available():
|
||||
info = mx.device_info()
|
||||
rec_bytes = info.get("max_recommended_working_set_size", 0) or 0
|
||||
if rec_bytes > 0:
|
||||
memory_cap = int(rec_bytes * 0.85)
|
||||
wired_cap = min(int(rec_bytes), memory_cap)
|
||||
mx.set_memory_limit(memory_cap)
|
||||
mx.set_wired_limit(wired_cap)
|
||||
|
||||
model_name = config["model_name"]
|
||||
hf_token = config.get("hf_token") or None
|
||||
if hf_token:
|
||||
os.environ["HF_TOKEN"] = hf_token
|
||||
|
||||
if config.get("use_loftq"):
|
||||
message = "LoftQ is not supported for MLX training yet."
|
||||
_send("error", error = message)
|
||||
raise NotImplementedError(message)
|
||||
|
||||
optim_name = _normalize_mlx_studio_optimizer(config.get("optim", "adamw_8bit"))
|
||||
lr_scheduler_type = _normalize_mlx_studio_scheduler(
|
||||
config.get("lr_scheduler_type", "linear")
|
||||
)
|
||||
|
||||
# ── 1. Load model ──
|
||||
# Force text-only if the dataset is not an image dataset, even if the model
|
||||
# has vision capabilities (e.g. Qwen3.5-VL trained on plain alpaca text).
|
||||
_send("status", status_message = f"Loading {model_name}...")
|
||||
is_dataset_image = bool(config.get("is_dataset_image", False))
|
||||
training_type = config.get("training_type", "LoRA/QLoRA")
|
||||
use_lora = training_type == "LoRA/QLoRA"
|
||||
model, tokenizer = FastMLXModel.from_pretrained(
|
||||
model_name,
|
||||
load_in_4bit = config.get("load_in_4bit", True),
|
||||
full_finetuning = not use_lora,
|
||||
text_only = None if is_dataset_image else True,
|
||||
token = hf_token,
|
||||
trust_remote_code = bool(config.get("trust_remote_code", False)),
|
||||
random_state = config.get("random_seed", 3407),
|
||||
)
|
||||
|
||||
is_vlm = bool(is_dataset_image and getattr(model, "_is_vlm_model", False))
|
||||
model._is_vlm_model = is_vlm
|
||||
|
||||
# ── 2. Apply LoRA / full FT ──
|
||||
# Pass gradient_checkpointing as string ("mlx"/"unsloth"/"none"/etc.)
|
||||
# get_peft_model and MLXTrainer both accept strings and handle them.
|
||||
gc_setting = config.get("gradient_checkpointing", "mlx")
|
||||
if isinstance(gc_setting, str):
|
||||
use_grad_checkpoint = (
|
||||
gc_setting if gc_setting.lower() not in ("false", "") else False
|
||||
)
|
||||
else:
|
||||
use_grad_checkpoint = gc_setting
|
||||
|
||||
if use_lora:
|
||||
_send("status", status_message = "Configuring LoRA adapters...")
|
||||
peft_kwargs = dict(
|
||||
r = config.get("lora_r", 16),
|
||||
lora_alpha = config.get("lora_alpha", 16),
|
||||
lora_dropout = config.get("lora_dropout", 0.0),
|
||||
use_rslora = config.get("use_rslora", False),
|
||||
init_lora_weights = config.get("init_lora_weights", True),
|
||||
random_state = config.get("random_seed", 3407),
|
||||
target_modules = config.get("target_modules")
|
||||
or [
|
||||
"q_proj",
|
||||
"k_proj",
|
||||
"v_proj",
|
||||
"o_proj",
|
||||
"gate_proj",
|
||||
"up_proj",
|
||||
"down_proj",
|
||||
],
|
||||
use_gradient_checkpointing = use_grad_checkpoint,
|
||||
)
|
||||
finetune_language = config.get("finetune_language_layers", True)
|
||||
finetune_attention = config.get("finetune_attention_modules", True)
|
||||
finetune_mlp = config.get("finetune_mlp_modules", True)
|
||||
finetune_vision = (
|
||||
config.get("finetune_vision_layers", False) if is_vlm else False
|
||||
)
|
||||
|
||||
if (
|
||||
(finetune_attention or finetune_mlp)
|
||||
and not finetune_language
|
||||
and not finetune_vision
|
||||
):
|
||||
finetune_language = True
|
||||
|
||||
peft_kwargs["finetune_language_layers"] = finetune_language
|
||||
peft_kwargs["finetune_attention_modules"] = finetune_attention
|
||||
peft_kwargs["finetune_mlp_modules"] = finetune_mlp
|
||||
if is_vlm:
|
||||
peft_kwargs["finetune_vision_layers"] = finetune_vision
|
||||
model = FastMLXModel.get_peft_model(model, **peft_kwargs)
|
||||
|
||||
# ── 3. Load dataset ──
|
||||
_send("status", status_message = "Loading dataset...")
|
||||
hf_dataset = config.get("hf_dataset", "")
|
||||
subset = config.get("subset")
|
||||
train_split = config.get("train_split", "train") or "train"
|
||||
eval_split = config.get("eval_split")
|
||||
slice_start = config.get("dataset_slice_start")
|
||||
slice_end = config.get("dataset_slice_end")
|
||||
|
||||
def _slice(ds):
|
||||
if slice_start is not None or slice_end is not None:
|
||||
start = slice_start if slice_start is not None else 0
|
||||
end = slice_end if slice_end is not None else len(ds) - 1
|
||||
if end < start:
|
||||
return ds.select([])
|
||||
ds = ds.select(range(start, min(end + 1, len(ds))))
|
||||
return ds
|
||||
|
||||
def _load_local(file_paths):
|
||||
from core.training.trainer import UnslothTrainer
|
||||
from datasets import load_from_disk
|
||||
|
||||
if len(file_paths) == 1:
|
||||
p = Path(file_paths[0])
|
||||
if p.is_dir() and (
|
||||
(p / "dataset_info.json").exists() or (p / "state.json").exists()
|
||||
):
|
||||
return load_from_disk(str(p))
|
||||
all_files = UnslothTrainer._resolve_local_files(file_paths)
|
||||
if not all_files:
|
||||
raise ValueError("No local dataset files found")
|
||||
loader = UnslothTrainer._loader_for_files(all_files)
|
||||
return load_dataset(loader, data_files = all_files, split = "train")
|
||||
|
||||
if hf_dataset:
|
||||
load_kwargs = {"split": train_split, "token": hf_token}
|
||||
if subset:
|
||||
load_kwargs["name"] = subset
|
||||
dataset = load_dataset(hf_dataset, **load_kwargs)
|
||||
dataset = _slice(dataset)
|
||||
elif config.get("local_datasets"):
|
||||
dataset = _load_local(config["local_datasets"])
|
||||
dataset = _slice(dataset)
|
||||
else:
|
||||
raise ValueError("No dataset specified")
|
||||
|
||||
# Eval dataset (separate split or local file)
|
||||
eval_dataset = None
|
||||
if eval_split and hf_dataset:
|
||||
eval_kwargs = {"split": eval_split, "token": hf_token}
|
||||
if subset:
|
||||
eval_kwargs["name"] = subset
|
||||
try:
|
||||
eval_dataset = load_dataset(hf_dataset, **eval_kwargs)
|
||||
except Exception as e:
|
||||
_send("status", status_message = f"Eval split load failed: {e}")
|
||||
eval_dataset = None
|
||||
elif config.get("local_eval_datasets"):
|
||||
eval_dataset = _load_local(config["local_eval_datasets"])
|
||||
|
||||
# ── 3b. Format dataset (VLM or text) ──
|
||||
# Reuse the GPU path's format pipeline for both VLM (auto-detects OCR/caption/
|
||||
# llava/sharegpt+images) and text (alpaca/sharegpt/chatml → "text" column).
|
||||
format_type = config.get("format_type", "")
|
||||
try:
|
||||
from utils.datasets import format_and_template_dataset
|
||||
|
||||
def _fmt_progress(status_message = "", **_kw):
|
||||
_send("status", status_message = status_message)
|
||||
|
||||
if is_vlm:
|
||||
_send("status", status_message = "Formatting VLM dataset...")
|
||||
vlm_info = format_and_template_dataset(
|
||||
dataset,
|
||||
model_name = model_name,
|
||||
tokenizer = tokenizer,
|
||||
is_vlm = True,
|
||||
dataset_name = hf_dataset or "local",
|
||||
progress_callback = _fmt_progress,
|
||||
)
|
||||
if vlm_info.get("success"):
|
||||
dataset = _adapt_for_mlx_vlm(vlm_info["dataset"])
|
||||
else:
|
||||
errors = vlm_info.get("errors", [])
|
||||
raise ValueError(
|
||||
f"VLM dataset format conversion failed: {'; '.join(errors)}"
|
||||
)
|
||||
if eval_dataset is not None:
|
||||
ev_info = format_and_template_dataset(
|
||||
eval_dataset,
|
||||
model_name = model_name,
|
||||
tokenizer = tokenizer,
|
||||
is_vlm = True,
|
||||
dataset_name = hf_dataset or "local",
|
||||
)
|
||||
if ev_info.get("success"):
|
||||
eval_dataset = _adapt_for_mlx_vlm(ev_info["dataset"])
|
||||
|
||||
elif format_type:
|
||||
_send("status", status_message = f"Formatting dataset ({format_type})...")
|
||||
info = format_and_template_dataset(
|
||||
dataset,
|
||||
model_name = model_name,
|
||||
tokenizer = tokenizer,
|
||||
is_vlm = False,
|
||||
format_type = format_type,
|
||||
dataset_name = hf_dataset or "local",
|
||||
)
|
||||
if info.get("success", True):
|
||||
dataset = info.get("dataset", dataset)
|
||||
if eval_dataset is not None:
|
||||
ev = format_and_template_dataset(
|
||||
eval_dataset,
|
||||
model_name = model_name,
|
||||
tokenizer = tokenizer,
|
||||
is_vlm = False,
|
||||
format_type = format_type,
|
||||
dataset_name = hf_dataset or "local",
|
||||
)
|
||||
if ev.get("success", True):
|
||||
eval_dataset = ev.get("dataset", eval_dataset)
|
||||
except ImportError:
|
||||
_send("status", status_message = "Format helper unavailable, using raw dataset")
|
||||
|
||||
# ── 4. Resolve training steps ──
|
||||
max_steps = config.get("max_steps", 0) or 0
|
||||
num_epochs = config.get("num_epochs", 3)
|
||||
max_seq_length = config.get("max_seq_length", 2048)
|
||||
batch_size = config.get("batch_size", 4)
|
||||
grad_accum = config.get("gradient_accumulation_steps", 4)
|
||||
|
||||
if max_steps <= 0:
|
||||
max_steps = max(
|
||||
1,
|
||||
math.ceil(len(dataset) / batch_size / grad_accum) * num_epochs,
|
||||
)
|
||||
|
||||
lr_value = float(config.get("learning_rate", "2e-4"))
|
||||
|
||||
# Warmup: prefer warmup_steps; fall back to warmup_ratio
|
||||
warmup_steps = config.get("warmup_steps")
|
||||
warmup_ratio = config.get("warmup_ratio")
|
||||
if warmup_steps is None and warmup_ratio is not None:
|
||||
warmup_steps = int(round(warmup_ratio * max_steps))
|
||||
if warmup_steps is None:
|
||||
warmup_steps = 5
|
||||
|
||||
# ── 5. Build output dir ──
|
||||
output_dir = config.get("output_dir", "")
|
||||
if not output_dir:
|
||||
output_dir = f"{model_name.replace('/', '_')}_{int(time.time())}"
|
||||
# Resolve to ~/.unsloth/studio/outputs/ so the export page can find it
|
||||
from utils.paths import resolve_output_dir, ensure_dir
|
||||
|
||||
output_dir = str(resolve_output_dir(output_dir))
|
||||
ensure_dir(Path(output_dir))
|
||||
|
||||
# ── 6. Create trainer ──
|
||||
eval_steps_val = config.get("eval_steps", 0) or 0
|
||||
if isinstance(eval_steps_val, float) and 0 < eval_steps_val < 1:
|
||||
# Studio sometimes sends fraction-of-total-steps
|
||||
eval_steps_val = max(1, int(eval_steps_val * max_steps))
|
||||
else:
|
||||
eval_steps_val = int(eval_steps_val)
|
||||
|
||||
trainer = MLXTrainer(
|
||||
model = model,
|
||||
tokenizer = tokenizer,
|
||||
train_dataset = dataset,
|
||||
eval_dataset = eval_dataset,
|
||||
args = MLXTrainingConfig(
|
||||
per_device_train_batch_size = batch_size,
|
||||
gradient_accumulation_steps = grad_accum,
|
||||
max_steps = max_steps,
|
||||
learning_rate = lr_value,
|
||||
warmup_steps = warmup_steps,
|
||||
lr_scheduler_type = lr_scheduler_type,
|
||||
optim = optim_name,
|
||||
weight_decay = float(config.get("weight_decay", 0.001) or 0.001),
|
||||
logging_steps = 1,
|
||||
max_seq_length = max_seq_length,
|
||||
seed = config.get("random_seed", 3407),
|
||||
use_cce = True,
|
||||
compile = True,
|
||||
gradient_checkpointing = use_grad_checkpoint,
|
||||
streaming = is_vlm,
|
||||
packing = bool(config.get("packing", False)),
|
||||
output_dir = output_dir,
|
||||
save_steps = int(config.get("save_steps", 0) or 0),
|
||||
eval_steps = eval_steps_val,
|
||||
),
|
||||
)
|
||||
|
||||
# Tell the parent that eval is configured so the frontend shows the eval chart
|
||||
if eval_dataset is not None and eval_steps_val > 0:
|
||||
_send("eval_configured")
|
||||
|
||||
# ── 7. Apply train_on_responses_only if requested ──
|
||||
if config.get("train_on_completions", False):
|
||||
_send("status", status_message = "Configuring response-only training...")
|
||||
try:
|
||||
from utils.datasets import (
|
||||
MODEL_TO_TEMPLATE_MAPPER,
|
||||
TEMPLATE_TO_RESPONSES_MAPPER,
|
||||
)
|
||||
|
||||
template_name = MODEL_TO_TEMPLATE_MAPPER.get(model_name.lower())
|
||||
markers = (
|
||||
TEMPLATE_TO_RESPONSES_MAPPER.get(template_name)
|
||||
if template_name
|
||||
else None
|
||||
)
|
||||
if markers:
|
||||
trainer = train_on_responses_only(
|
||||
trainer,
|
||||
instruction_part = markers["instruction"],
|
||||
response_part = markers["response"],
|
||||
)
|
||||
else:
|
||||
_send(
|
||||
"status",
|
||||
status_message = f"train_on_completions skipped (no template for {model_name})",
|
||||
)
|
||||
except Exception as e:
|
||||
_send("status", status_message = f"train_on_completions failed: {e}")
|
||||
|
||||
# ── 8. Setup wandb / tensorboard ──
|
||||
wandb_run = None
|
||||
tb_writer = None
|
||||
if config.get("enable_wandb", False):
|
||||
try:
|
||||
import wandb as _wandb
|
||||
|
||||
wandb_token = config.get("wandb_token")
|
||||
if wandb_token:
|
||||
os.environ["WANDB_API_KEY"] = wandb_token
|
||||
_wandb_sensitive = {"hf_token", "wandb_token"}
|
||||
wandb_run = _wandb.init(
|
||||
project = config.get("wandb_project") or "unsloth-mlx",
|
||||
config = {k: v for k, v in config.items() if k not in _wandb_sensitive},
|
||||
reinit = True,
|
||||
)
|
||||
except Exception as e:
|
||||
_send("status", status_message = f"wandb init failed: {e}")
|
||||
if config.get("enable_tensorboard", False):
|
||||
try:
|
||||
from tensorboardX import SummaryWriter
|
||||
except ImportError:
|
||||
try:
|
||||
from torch.utils.tensorboard import SummaryWriter
|
||||
except ImportError:
|
||||
SummaryWriter = None
|
||||
if SummaryWriter is not None:
|
||||
try:
|
||||
tb_dir = config.get("tensorboard_dir") or f"{output_dir}/runs"
|
||||
tb_writer = SummaryWriter(log_dir = tb_dir)
|
||||
except Exception as e:
|
||||
_send("status", status_message = f"tensorboard init failed: {e}")
|
||||
else:
|
||||
_send(
|
||||
"status",
|
||||
status_message = "tensorboard unavailable (install tensorboardX)",
|
||||
)
|
||||
|
||||
# ── 9. Real-time progress callback ──
|
||||
_send("status", status_message = f"Training {model_name}...")
|
||||
|
||||
def _on_step(step, total, loss, lr, tok_s, peak_gb, elapsed, num_tokens):
|
||||
eta = (elapsed / step * (total - step)) if step > 0 else 0
|
||||
_send(
|
||||
"progress",
|
||||
step = step,
|
||||
epoch = round(step / total * num_epochs, 2) if total > 0 else 0,
|
||||
loss = loss,
|
||||
learning_rate = lr,
|
||||
total_steps = total,
|
||||
elapsed_seconds = elapsed,
|
||||
eta_seconds = max(0, eta),
|
||||
grad_norm = None,
|
||||
num_tokens = num_tokens,
|
||||
eval_loss = None,
|
||||
status_message = None,
|
||||
peak_memory_gb = peak_gb,
|
||||
)
|
||||
if wandb_run is not None:
|
||||
try:
|
||||
wandb_run.log(
|
||||
{
|
||||
"train/loss": loss,
|
||||
"train/learning_rate": lr,
|
||||
"train/tokens_per_sec": tok_s,
|
||||
"train/peak_gb": peak_gb,
|
||||
"train/num_tokens": num_tokens,
|
||||
},
|
||||
step = step,
|
||||
)
|
||||
except Exception:
|
||||
pass
|
||||
if tb_writer is not None:
|
||||
try:
|
||||
tb_writer.add_scalar("train/loss", loss, step)
|
||||
tb_writer.add_scalar("train/learning_rate", lr, step)
|
||||
tb_writer.add_scalar("train/tokens_per_sec", tok_s, step)
|
||||
tb_writer.add_scalar("train/peak_gb", peak_gb, step)
|
||||
except Exception:
|
||||
pass
|
||||
|
||||
trainer.add_step_callback(_on_step)
|
||||
|
||||
def _on_eval(step, eval_loss, perplexity):
|
||||
_send("progress", step = step, eval_loss = eval_loss)
|
||||
if wandb_run is not None:
|
||||
try:
|
||||
wandb_run.log(
|
||||
{"eval/loss": eval_loss, "eval/perplexity": perplexity}, step = step
|
||||
)
|
||||
except Exception:
|
||||
pass
|
||||
if tb_writer is not None:
|
||||
try:
|
||||
tb_writer.add_scalar("eval/loss", eval_loss, step)
|
||||
tb_writer.add_scalar("eval/perplexity", perplexity, step)
|
||||
except Exception:
|
||||
pass
|
||||
|
||||
trainer.add_eval_callback(_on_eval)
|
||||
|
||||
# ── 10. Stop signal polling ──
|
||||
_stop_save = [True] # mutable so thread can update; [save_flag]
|
||||
|
||||
def _poll_stop():
|
||||
while True:
|
||||
try:
|
||||
msg = stop_queue.get(timeout = 1.0)
|
||||
if msg and msg.get("type") == "stop":
|
||||
_stop_save[0] = msg.get("save", True)
|
||||
trainer.stop_requested = True
|
||||
return
|
||||
except _queue.Empty:
|
||||
continue
|
||||
except (EOFError, OSError):
|
||||
# why safe: pipe permanently broken, no further messages can arrive
|
||||
return
|
||||
|
||||
stop_thread = threading.Thread(target = _poll_stop, daemon = True)
|
||||
stop_thread.start()
|
||||
|
||||
# ── 11. Run training ──
|
||||
gc.collect()
|
||||
mx.synchronize()
|
||||
trainer.train()
|
||||
|
||||
# ── 12. Save and finalize ──
|
||||
if trainer.stop_requested and not _stop_save[0]:
|
||||
# User clicked "Cancel" (save=False) — skip saving
|
||||
_send("complete", output_dir = None, status_message = "Training cancelled")
|
||||
else:
|
||||
_send("status", status_message = "Saving model...")
|
||||
mx.synchronize()
|
||||
trainer.save_model(output_dir)
|
||||
_send("complete", output_dir = output_dir, status_message = "Training completed")
|
||||
|
||||
if tb_writer is not None:
|
||||
try:
|
||||
tb_writer.close()
|
||||
except Exception:
|
||||
pass
|
||||
if wandb_run is not None:
|
||||
try:
|
||||
wandb_run.finish()
|
||||
except Exception:
|
||||
pass
|
||||
|
||||
|
||||
def run_training_process(
|
||||
*,
|
||||
event_queue: Any,
|
||||
|
|
@ -371,6 +959,46 @@ def run_training_process(
|
|||
|
||||
model_name = config["model_name"]
|
||||
|
||||
# ── 0. MLX FAST-PATH (must run before any torch/transformers imports) ──
|
||||
# Apple Silicon uses MLXTrainer directly -- skip transformers version
|
||||
# activation, causal-conv1d install, and torch imports entirely.
|
||||
backend_path = str(Path(__file__).resolve().parent.parent.parent)
|
||||
if backend_path not in sys.path:
|
||||
sys.path.insert(0, backend_path)
|
||||
|
||||
from utils.hardware import hardware as _hw
|
||||
|
||||
_hw.detect_hardware()
|
||||
if _hw.DEVICE == _hw.DeviceType.MLX:
|
||||
if config.get("is_dataset_audio"):
|
||||
event_queue.put(
|
||||
{
|
||||
"type": "error",
|
||||
"error": "Audio dataset training is not yet supported on Apple Silicon.",
|
||||
"stack": "",
|
||||
"ts": time.time(),
|
||||
}
|
||||
)
|
||||
return
|
||||
# Activate correct transformers version (Gemma-4 needs 5.5.0, etc.)
|
||||
# Must happen before any transformers/mlx-lm imports in _run_mlx_training.
|
||||
try:
|
||||
_activate_transformers_version(model_name)
|
||||
except Exception:
|
||||
pass # Non-fatal: fall through with whatever version is installed
|
||||
try:
|
||||
_run_mlx_training(event_queue, stop_queue, config)
|
||||
except Exception as exc:
|
||||
event_queue.put(
|
||||
{
|
||||
"type": "error",
|
||||
"error": str(exc),
|
||||
"stack": traceback.format_exc(limit = 20),
|
||||
"ts": time.time(),
|
||||
}
|
||||
)
|
||||
return
|
||||
|
||||
# ── 1. Activate correct transformers version BEFORE any ML imports ──
|
||||
try:
|
||||
_activate_transformers_version(model_name)
|
||||
|
|
|
|||
|
|
@ -23,12 +23,67 @@ if _backend_dir not in sys.path:
|
|||
# See: https://github.com/python/cpython/issues/102396
|
||||
import _platform_compat # noqa: F401
|
||||
|
||||
# Direct `uvicorn main:app` launches bypass run.py, so re-export here too
|
||||
# (mirrors run.py). Required BEFORE the unsloth-zoo import below, since
|
||||
# its LLAMA_CPP_DEFAULT_DIR binding is import-time.
|
||||
from utils.paths.storage_roots import studio_root as _studio_root
|
||||
|
||||
try:
|
||||
_LEGACY_STUDIO_ROOT = (_Path.home() / ".unsloth" / "studio").resolve()
|
||||
except (OSError, ValueError):
|
||||
_LEGACY_STUDIO_ROOT = _Path.home() / ".unsloth" / "studio"
|
||||
try:
|
||||
_STUDIO_ROOT_RESOLVED = _studio_root().resolve()
|
||||
except (OSError, ValueError):
|
||||
_STUDIO_ROOT_RESOLVED = _studio_root()
|
||||
if _STUDIO_ROOT_RESOLVED != _LEGACY_STUDIO_ROOT:
|
||||
if not os.environ.get("UNSLOTH_STUDIO_HOME"):
|
||||
os.environ["UNSLOTH_STUDIO_HOME"] = str(_STUDIO_ROOT_RESOLVED)
|
||||
if not os.environ.get("UNSLOTH_LLAMA_CPP_PATH"):
|
||||
os.environ["UNSLOTH_LLAMA_CPP_PATH"] = str(_STUDIO_ROOT_RESOLVED / "llama.cpp")
|
||||
|
||||
import mimetypes
|
||||
import re as _re
|
||||
import shutil
|
||||
import warnings
|
||||
from contextlib import asynccontextmanager
|
||||
from importlib.metadata import PackageNotFoundError, version as package_version
|
||||
|
||||
|
||||
_STUDIO_INSTALL_ID_RE = _re.compile(r"^[0-9a-f]{64}$")
|
||||
|
||||
|
||||
def _read_studio_install_id() -> str:
|
||||
"""Per-install opaque id written by install.sh / install.ps1 at
|
||||
$STUDIO_HOME/share/studio_install_id. Returns "" when the file is
|
||||
absent (pre-PR install, fresh tree never run through the installer)
|
||||
or contains anything other than a 64-char lowercase-hex token --
|
||||
in which case /api/health emits "" and the launcher's _check_health
|
||||
falls back to the existing "no baked id, accept any healthy
|
||||
Unsloth backend" path. This intentionally replaces a previous
|
||||
sha256(resolved_install_path) so the field carries no install-path
|
||||
information for callers reaching /api/health (relevant when Studio
|
||||
is run with -H 0.0.0.0)."""
|
||||
try:
|
||||
token = (
|
||||
(_STUDIO_ROOT_RESOLVED / "share" / "studio_install_id").read_text().strip()
|
||||
)
|
||||
except (OSError, ValueError):
|
||||
return ""
|
||||
return token if _STUDIO_INSTALL_ID_RE.fullmatch(token) else ""
|
||||
|
||||
|
||||
_STUDIO_ROOT_ID_CACHE: str = _read_studio_install_id()
|
||||
|
||||
|
||||
def _studio_root_id() -> str:
|
||||
"""Same-install discriminator for /api/health: a per-install opaque
|
||||
token written once by the installer and read once at module import.
|
||||
Empty when no installer-written token is present; the launcher
|
||||
contract treats "" as "no baked id, accept any healthy backend"."""
|
||||
return _STUDIO_ROOT_ID_CACHE
|
||||
|
||||
|
||||
# Fix broken Windows registry MIME types. Some Windows installs map .js to
|
||||
# "text/plain" in the registry (HKCR\.js\Content Type). Python's mimetypes
|
||||
# module reads from the registry, and FastAPI/Starlette's StaticFiles uses
|
||||
|
|
@ -252,6 +307,10 @@ async def health_check():
|
|||
"chat_only": _hw_module.CHAT_ONLY,
|
||||
"desktop_protocol_version": 1,
|
||||
"supports_desktop_auth": True,
|
||||
# why: launchers compare against an install-time hash so a sibling
|
||||
# Studio on the same port is rejected; hex digest avoids leaking the
|
||||
# raw install path on -H 0.0.0.0.
|
||||
"studio_root_id": _studio_root_id(),
|
||||
"native_path_leases_supported": native_path_leases_supported(),
|
||||
}
|
||||
|
||||
|
|
|
|||
|
|
@ -159,7 +159,27 @@ def _find_free_port(host: str, start: int, max_attempts: int = 20) -> int:
|
|||
)
|
||||
|
||||
|
||||
_PID_FILE = Path.home() / ".unsloth" / "studio" / "studio.pid"
|
||||
from utils.paths.storage_roots import studio_root as _studio_root
|
||||
|
||||
_PID_FILE = _studio_root() / "studio.pid"
|
||||
|
||||
# Direct backend launches bypass the CLI's env re-export; do it here for
|
||||
# real custom roots so unsloth-zoo's import-time LLAMA_CPP_DEFAULT_DIR
|
||||
# picks up the custom build. Skip for legacy-default to avoid flipping
|
||||
# default-mode installs into env-override.
|
||||
try:
|
||||
_LEGACY_STUDIO_ROOT = (Path.home() / ".unsloth" / "studio").resolve()
|
||||
except (OSError, ValueError):
|
||||
_LEGACY_STUDIO_ROOT = Path.home() / ".unsloth" / "studio"
|
||||
try:
|
||||
_STUDIO_ROOT_RESOLVED = _studio_root().resolve()
|
||||
except (OSError, ValueError):
|
||||
_STUDIO_ROOT_RESOLVED = _studio_root()
|
||||
if _STUDIO_ROOT_RESOLVED != _LEGACY_STUDIO_ROOT:
|
||||
if not os.environ.get("UNSLOTH_STUDIO_HOME"):
|
||||
os.environ["UNSLOTH_STUDIO_HOME"] = str(_STUDIO_ROOT_RESOLVED)
|
||||
if not os.environ.get("UNSLOTH_LLAMA_CPP_PATH"):
|
||||
os.environ["UNSLOTH_LLAMA_CPP_PATH"] = str(_STUDIO_ROOT_RESOLVED / "llama.cpp")
|
||||
|
||||
|
||||
def _write_pid_file():
|
||||
|
|
|
|||
157
studio/backend/tests/test_mlx_inference_backend.py
Normal file
157
studio/backend/tests/test_mlx_inference_backend.py
Normal file
|
|
@ -0,0 +1,157 @@
|
|||
# SPDX-License-Identifier: AGPL-3.0-only
|
||||
|
||||
import sys
|
||||
import types
|
||||
from types import SimpleNamespace
|
||||
|
||||
|
||||
class _DummyMetal:
|
||||
@staticmethod
|
||||
def is_available():
|
||||
return False
|
||||
|
||||
|
||||
class _DummyMX:
|
||||
metal = _DummyMetal()
|
||||
|
||||
@staticmethod
|
||||
def set_wired_limit(_limit):
|
||||
return None
|
||||
|
||||
@staticmethod
|
||||
def device_info():
|
||||
return {"max_recommended_working_set_size": 1024}
|
||||
|
||||
|
||||
class _DummyTokenizer:
|
||||
pass
|
||||
|
||||
|
||||
class _DummyProcessor:
|
||||
tokenizer = _DummyTokenizer()
|
||||
|
||||
|
||||
class _DummyModel:
|
||||
pass
|
||||
|
||||
|
||||
def _install_fake_mlx(monkeypatch):
|
||||
mlx_pkg = types.ModuleType("mlx")
|
||||
mlx_core = types.ModuleType("mlx.core")
|
||||
mlx_core.metal = _DummyMetal()
|
||||
mlx_core.set_wired_limit = _DummyMX.set_wired_limit
|
||||
mlx_core.device_info = _DummyMX.device_info
|
||||
mlx_pkg.core = mlx_core
|
||||
monkeypatch.setitem(sys.modules, "mlx", mlx_pkg)
|
||||
monkeypatch.setitem(sys.modules, "mlx.core", mlx_core)
|
||||
|
||||
|
||||
def _install_fake_fast_mlx(monkeypatch, calls):
|
||||
class _FastMLXModel:
|
||||
@staticmethod
|
||||
def from_pretrained(*args, **kwargs):
|
||||
calls.append((args, kwargs))
|
||||
if kwargs["text_only"] is False:
|
||||
return _DummyModel(), _DummyProcessor()
|
||||
return _DummyModel(), _DummyTokenizer()
|
||||
|
||||
unsloth_zoo_pkg = types.ModuleType("unsloth_zoo")
|
||||
mlx_loader = types.ModuleType("unsloth_zoo.mlx_loader")
|
||||
mlx_loader.FastMLXModel = _FastMLXModel
|
||||
unsloth_zoo_pkg.mlx_loader = mlx_loader
|
||||
monkeypatch.setitem(sys.modules, "unsloth_zoo", unsloth_zoo_pkg)
|
||||
monkeypatch.setitem(sys.modules, "unsloth_zoo.mlx_loader", mlx_loader)
|
||||
|
||||
|
||||
def test_mlx_inference_text_load_forwards_studio_settings(monkeypatch):
|
||||
_install_fake_mlx(monkeypatch)
|
||||
calls = []
|
||||
_install_fake_fast_mlx(monkeypatch, calls)
|
||||
|
||||
from core.inference.mlx_inference import MLXInferenceBackend
|
||||
|
||||
backend = MLXInferenceBackend()
|
||||
config = SimpleNamespace(identifier = "fake/text", is_vision = False, is_lora = False)
|
||||
|
||||
assert backend.load_model(
|
||||
config,
|
||||
max_seq_length = 4096,
|
||||
load_in_4bit = False,
|
||||
hf_token = "hf-token",
|
||||
trust_remote_code = True,
|
||||
dtype = "float16",
|
||||
)
|
||||
|
||||
assert calls == [
|
||||
(
|
||||
("fake/text",),
|
||||
{
|
||||
"max_seq_length": 4096,
|
||||
"dtype": "float16",
|
||||
"load_in_4bit": False,
|
||||
"token": "hf-token",
|
||||
"trust_remote_code": True,
|
||||
"text_only": True,
|
||||
},
|
||||
)
|
||||
]
|
||||
assert backend._is_vlm is False
|
||||
assert isinstance(backend._tokenizer, _DummyTokenizer)
|
||||
|
||||
|
||||
def test_mlx_inference_vlm_lora_uses_unsloth_loader_without_native_adapter_rewrite(
|
||||
monkeypatch,
|
||||
tmp_path,
|
||||
):
|
||||
_install_fake_mlx(monkeypatch)
|
||||
calls = []
|
||||
_install_fake_fast_mlx(monkeypatch, calls)
|
||||
|
||||
def _native_vlm_load(*_args, **_kwargs):
|
||||
raise AssertionError("Studio MLX VLM inference must use FastMLXModel")
|
||||
|
||||
mlx_vlm = types.ModuleType("mlx_vlm")
|
||||
mlx_vlm.load = _native_vlm_load
|
||||
monkeypatch.setitem(sys.modules, "mlx_vlm", mlx_vlm)
|
||||
|
||||
adapter_dir = tmp_path / "adapter"
|
||||
adapter_dir.mkdir()
|
||||
cfg_path = adapter_dir / "adapter_config.json"
|
||||
original_cfg = '{"base_model_name_or_path": "fake/base", "rank": 8}\n'
|
||||
cfg_path.write_text(original_cfg)
|
||||
|
||||
from core.inference.mlx_inference import MLXInferenceBackend
|
||||
|
||||
backend = MLXInferenceBackend()
|
||||
config = SimpleNamespace(
|
||||
identifier = str(adapter_dir),
|
||||
is_vision = True,
|
||||
is_lora = True,
|
||||
base_model = "fake/base",
|
||||
)
|
||||
|
||||
assert backend.load_model(
|
||||
config,
|
||||
max_seq_length = 8192,
|
||||
load_in_4bit = True,
|
||||
hf_token = "hf-token",
|
||||
trust_remote_code = True,
|
||||
)
|
||||
|
||||
assert calls == [
|
||||
(
|
||||
(str(adapter_dir),),
|
||||
{
|
||||
"max_seq_length": 8192,
|
||||
"dtype": None,
|
||||
"load_in_4bit": True,
|
||||
"token": "hf-token",
|
||||
"trust_remote_code": True,
|
||||
"text_only": False,
|
||||
},
|
||||
)
|
||||
]
|
||||
assert cfg_path.read_text() == original_cfg
|
||||
assert backend._is_vlm is True
|
||||
assert isinstance(backend._processor, _DummyProcessor)
|
||||
assert isinstance(backend._tokenizer, _DummyTokenizer)
|
||||
83
studio/backend/tests/test_mlx_training_worker_config.py
Normal file
83
studio/backend/tests/test_mlx_training_worker_config.py
Normal file
|
|
@ -0,0 +1,83 @@
|
|||
# SPDX-License-Identifier: AGPL-3.0-only
|
||||
|
||||
import importlib.util
|
||||
import sys
|
||||
import types
|
||||
from pathlib import Path
|
||||
|
||||
import pytest
|
||||
|
||||
|
||||
def _load_worker_module():
|
||||
stub_names = (
|
||||
"structlog",
|
||||
"loggers",
|
||||
"utils",
|
||||
"utils.hardware",
|
||||
"utils.wheel_utils",
|
||||
)
|
||||
previous_modules = {name: sys.modules.get(name) for name in stub_names}
|
||||
|
||||
try:
|
||||
sys.modules["structlog"] = types.ModuleType("structlog")
|
||||
|
||||
loggers = types.ModuleType("loggers")
|
||||
loggers.get_logger = lambda *_args, **_kwargs: None
|
||||
sys.modules["loggers"] = loggers
|
||||
|
||||
utils = types.ModuleType("utils")
|
||||
utils.__path__ = []
|
||||
sys.modules["utils"] = utils
|
||||
|
||||
hardware = types.ModuleType("utils.hardware")
|
||||
hardware.apply_gpu_ids = lambda *_args, **_kwargs: None
|
||||
sys.modules["utils.hardware"] = hardware
|
||||
|
||||
wheel_utils = types.ModuleType("utils.wheel_utils")
|
||||
for name in (
|
||||
"direct_wheel_url",
|
||||
"flash_attn_wheel_url",
|
||||
"install_wheel",
|
||||
"probe_torch_wheel_env",
|
||||
"url_exists",
|
||||
):
|
||||
setattr(wheel_utils, name, lambda *_args, **_kwargs: None)
|
||||
sys.modules["utils.wheel_utils"] = wheel_utils
|
||||
|
||||
worker_path = (
|
||||
Path(__file__).resolve().parents[1] / "core" / "training" / "worker.py"
|
||||
)
|
||||
spec = importlib.util.spec_from_file_location(
|
||||
"mlx_training_worker_under_test", worker_path
|
||||
)
|
||||
module = importlib.util.module_from_spec(spec)
|
||||
assert spec.loader is not None
|
||||
spec.loader.exec_module(module)
|
||||
return module
|
||||
finally:
|
||||
for name, module in previous_modules.items():
|
||||
if module is None:
|
||||
sys.modules.pop(name, None)
|
||||
else:
|
||||
sys.modules[name] = module
|
||||
|
||||
|
||||
_worker = _load_worker_module()
|
||||
_normalize_mlx_studio_optimizer = _worker._normalize_mlx_studio_optimizer
|
||||
_normalize_mlx_studio_scheduler = _worker._normalize_mlx_studio_scheduler
|
||||
|
||||
|
||||
def test_mlx_studio_optimizer_aliases_are_explicit():
|
||||
assert _normalize_mlx_studio_optimizer("adamw_8bit") == "adamw"
|
||||
assert _normalize_mlx_studio_optimizer("paged_adamw_8bit") == "adamw"
|
||||
assert _normalize_mlx_studio_optimizer("adafactor") == "adafactor"
|
||||
|
||||
|
||||
def test_mlx_studio_rejects_unknown_optimizer():
|
||||
with pytest.raises(ValueError, match = "Unsupported optimizer for MLX training"):
|
||||
_normalize_mlx_studio_optimizer("adamw_typo")
|
||||
|
||||
|
||||
def test_mlx_studio_rejects_unknown_scheduler():
|
||||
with pytest.raises(ValueError, match = "Unsupported LR scheduler for MLX training"):
|
||||
_normalize_mlx_studio_scheduler("linear_typo")
|
||||
|
|
@ -143,6 +143,7 @@ def detect_hardware() -> DeviceType:
|
|||
# --- MLX: Apple Silicon ---
|
||||
if is_apple_silicon() and _has_mlx():
|
||||
DEVICE = DeviceType.MLX
|
||||
CHAT_ONLY = False
|
||||
chip = platform.processor() or platform.machine()
|
||||
print(f"Hardware detected: MLX — Apple Silicon ({chip})")
|
||||
return DEVICE
|
||||
|
|
@ -270,19 +271,30 @@ def get_gpu_memory_info() -> Dict[str, Any]:
|
|||
import mlx.core as mx
|
||||
import psutil
|
||||
|
||||
# MLX uses unified memory — report system memory as the pool
|
||||
# MLX uses unified memory. Total = system RAM. GPU memory used
|
||||
# comes from IORegistry's AGXAccelerator (system-wide, no sudo).
|
||||
total = psutil.virtual_memory().total
|
||||
# MLX doesn't expose per-process GPU allocation; report 0 as allocated
|
||||
allocated = 0
|
||||
agx = _read_apple_gpu_stats()
|
||||
allocated = agx.get("vram_used_bytes", 0) if agx else 0
|
||||
|
||||
try:
|
||||
info = mx.device_info()
|
||||
gpu_name = (
|
||||
info.get("device_name")
|
||||
or platform.processor()
|
||||
or platform.machine()
|
||||
)
|
||||
except Exception:
|
||||
gpu_name = platform.processor() or platform.machine()
|
||||
|
||||
return {
|
||||
"available": True,
|
||||
"backend": _backend_label(device),
|
||||
"device": 0,
|
||||
"device_name": f"Apple Silicon ({platform.processor() or platform.machine()})",
|
||||
"device_name": f"Apple Silicon ({gpu_name})",
|
||||
"total_gb": total / (1024**3),
|
||||
"allocated_gb": allocated / (1024**3),
|
||||
"reserved_gb": 0,
|
||||
"reserved_gb": allocated / (1024**3),
|
||||
"free_gb": (total - allocated) / (1024**3),
|
||||
"utilization_pct": (allocated / total) * 100 if total else 0,
|
||||
}
|
||||
|
|
@ -460,6 +472,39 @@ def _smi_query(func_name: str, *args, **kwargs) -> Optional[Dict[str, Any]]:
|
|||
return None
|
||||
|
||||
|
||||
def _read_apple_gpu_stats() -> Dict[str, Any]:
|
||||
"""Query macOS IORegistry for AGX (Apple GPU) live stats. No sudo needed.
|
||||
|
||||
Returns dict with utilization_pct, vram_used_bytes (system-wide GPU memory).
|
||||
Returns empty dict on failure.
|
||||
"""
|
||||
import subprocess
|
||||
import re
|
||||
|
||||
try:
|
||||
result = subprocess.run(
|
||||
["ioreg", "-r", "-c", "AGXAccelerator"],
|
||||
capture_output = True,
|
||||
timeout = 2,
|
||||
)
|
||||
text = result.stdout.decode("utf-8", errors = "replace")
|
||||
except Exception:
|
||||
return {}
|
||||
|
||||
# PerformanceStatistics block has GPU utilization and in-use memory
|
||||
m = re.search(r'"PerformanceStatistics" = \{([^}]+)\}', text)
|
||||
if not m:
|
||||
return {}
|
||||
stats_str = m.group(1)
|
||||
pairs = re.findall(r'"([^"]+)"=(\d+)', stats_str)
|
||||
stats = {k: int(v) for k, v in pairs}
|
||||
|
||||
return {
|
||||
"utilization_pct": stats.get("Device Utilization %", 0),
|
||||
"vram_used_bytes": stats.get("In use system memory", 0),
|
||||
}
|
||||
|
||||
|
||||
def get_gpu_utilization() -> Dict[str, Any]:
|
||||
"""Return a live snapshot of device utilization information."""
|
||||
device = get_device()
|
||||
|
|
@ -470,6 +515,50 @@ def get_gpu_utilization() -> Dict[str, Any]:
|
|||
result["backend"] = _backend_label(device)
|
||||
return result
|
||||
|
||||
# MLX path: single _read_apple_gpu_stats() call carries both VRAM-used
|
||||
# bytes and GPU utilization %. psutil for unified-memory total is cheap.
|
||||
if device == DeviceType.MLX:
|
||||
try:
|
||||
import psutil
|
||||
|
||||
agx = _read_apple_gpu_stats()
|
||||
total_bytes = psutil.virtual_memory().total
|
||||
except Exception as e:
|
||||
logger.error(f"Error getting MLX GPU utilization: {e}")
|
||||
return {"available": False, "backend": device.value, "error": str(e)}
|
||||
if not agx:
|
||||
return {"available": False, "backend": device.value}
|
||||
allocated_bytes = agx.get("vram_used_bytes", 0) or 0
|
||||
vram_used_gb = allocated_bytes / (1024**3)
|
||||
total_gb = total_bytes / (1024**3)
|
||||
|
||||
try:
|
||||
from core.training import get_training_backend
|
||||
|
||||
tb = get_training_backend()
|
||||
tb_progress = getattr(tb, "_progress", None)
|
||||
if tb_progress is not None and getattr(tb_progress, "is_training", False):
|
||||
tb_peak = getattr(tb_progress, "peak_memory_gb", None)
|
||||
if tb_peak is not None and tb_peak > 0:
|
||||
vram_used_gb = float(tb_peak)
|
||||
except Exception:
|
||||
pass
|
||||
|
||||
return {
|
||||
"available": True,
|
||||
"backend": device.value,
|
||||
"gpu_utilization_pct": agx.get("utilization_pct") if agx else None,
|
||||
"temperature_c": None,
|
||||
"vram_used_gb": round(vram_used_gb, 2),
|
||||
"vram_total_gb": round(total_gb, 2),
|
||||
"vram_utilization_pct": (
|
||||
round((vram_used_gb / total_gb) * 100, 1) if total_gb > 0 else None
|
||||
),
|
||||
"power_draw_w": None,
|
||||
"power_limit_w": None,
|
||||
"power_utilization_pct": None,
|
||||
}
|
||||
|
||||
mem = get_gpu_memory_info()
|
||||
if device != DeviceType.CPU and mem.get("available"):
|
||||
return {
|
||||
|
|
|
|||
|
|
@ -500,7 +500,9 @@ _VLM_MODEL_TYPES = {
|
|||
|
||||
# Pre-computed .venv_t5 paths and backend dir for subprocess version switching.
|
||||
# Vision check uses 5.5.0 (newest, recognizes all architectures).
|
||||
_VENV_T5_DIR = str(Path.home() / ".unsloth" / "studio" / ".venv_t5_550")
|
||||
from utils.paths.storage_roots import studio_root as _studio_root # noqa: E402
|
||||
|
||||
_VENV_T5_DIR = str(_studio_root() / ".venv_t5_550")
|
||||
_BACKEND_DIR = str(Path(__file__).resolve().parent.parent.parent)
|
||||
|
||||
# Inline script executed in a subprocess with transformers 5.x activated.
|
||||
|
|
|
|||
|
|
@ -5,17 +5,59 @@ from __future__ import annotations
|
|||
|
||||
import json
|
||||
import os
|
||||
import sys
|
||||
from pathlib import Path
|
||||
import tempfile
|
||||
|
||||
|
||||
def _infer_studio_home_from_venv() -> Path | None:
|
||||
"""Return parent dir of sys.prefix as STUDIO_HOME if running from an
|
||||
installer-managed unsloth_studio venv. Sentinel-gated (share/studio.conf
|
||||
or bin shim) so a developer venv named unsloth_studio is not misidentified.
|
||||
"""
|
||||
try:
|
||||
prefix = Path(sys.prefix).resolve()
|
||||
except (OSError, ValueError):
|
||||
return None
|
||||
if prefix.name != "unsloth_studio":
|
||||
return None
|
||||
candidate = prefix.parent
|
||||
shim_name = "unsloth.exe" if os.name == "nt" else "unsloth"
|
||||
try:
|
||||
has_sentinel = (candidate / "share" / "studio.conf").is_file() or (
|
||||
candidate / "bin" / shim_name
|
||||
).is_file()
|
||||
except OSError:
|
||||
return None
|
||||
if has_sentinel:
|
||||
return candidate
|
||||
return None
|
||||
|
||||
|
||||
def studio_root() -> Path:
|
||||
"""Studio install root.
|
||||
|
||||
Priority: UNSLOTH_STUDIO_HOME, then STUDIO_HOME alias, then sys.prefix
|
||||
inference, then legacy ~/.unsloth/studio. UNSLOTH_STUDIO_HOME wins when
|
||||
both are set (the more specific signal beats the generic alias).
|
||||
"""
|
||||
override = (os.environ.get("UNSLOTH_STUDIO_HOME") or "").strip()
|
||||
if not override:
|
||||
override = (os.environ.get("STUDIO_HOME") or "").strip()
|
||||
if override:
|
||||
try:
|
||||
return Path(override).expanduser().resolve()
|
||||
except (OSError, ValueError):
|
||||
return Path(override).expanduser()
|
||||
inferred = _infer_studio_home_from_venv()
|
||||
if inferred is not None:
|
||||
return inferred
|
||||
return Path.home() / ".unsloth" / "studio"
|
||||
|
||||
|
||||
def cache_root() -> Path:
|
||||
"""Central cache directory for all studio downloads (models, datasets, etc.)."""
|
||||
return Path.home() / ".unsloth" / "studio" / "cache"
|
||||
return studio_root() / "cache"
|
||||
|
||||
|
||||
def assets_root() -> Path:
|
||||
|
|
|
|||
|
|
@ -95,9 +95,11 @@ TRANSFORMERS_DEFAULT_VERSION = "4.57.6"
|
|||
# Consumers should prefer TRANSFORMERS_530_VERSION / TRANSFORMERS_550_VERSION.
|
||||
TRANSFORMERS_5_VERSION = TRANSFORMERS_550_VERSION
|
||||
|
||||
# Pre-installed directories — created by setup.sh / setup.ps1
|
||||
_VENV_T5_530_DIR = str(Path.home() / ".unsloth" / "studio" / ".venv_t5_530")
|
||||
_VENV_T5_550_DIR = str(Path.home() / ".unsloth" / "studio" / ".venv_t5_550")
|
||||
# Pre-installed directories — created by setup.sh / setup.ps1.
|
||||
from utils.paths.storage_roots import studio_root as _studio_root # noqa: E402
|
||||
|
||||
_VENV_T5_530_DIR = str(_studio_root() / ".venv_t5_530")
|
||||
_VENV_T5_550_DIR = str(_studio_root() / ".venv_t5_550")
|
||||
# Backwards-compat alias
|
||||
_VENV_T5_DIR = _VENV_T5_550_DIR
|
||||
|
||||
|
|
|
|||
16817
studio/frontend/package-lock.json
generated
Normal file
16817
studio/frontend/package-lock.json
generated
Normal file
File diff suppressed because it is too large
Load diff
|
|
@ -17,9 +17,9 @@
|
|||
},
|
||||
"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",
|
||||
"@assistant-ui/react": "0.12.28",
|
||||
"@assistant-ui/react-markdown": "0.12.11",
|
||||
"@assistant-ui/react-streamdown": "0.1.11",
|
||||
"@base-ui/react": "^1.2.0",
|
||||
"@dagrejs/dagre": "^2.0.4",
|
||||
"@dagrejs/graphlib": "^3.0.4",
|
||||
|
|
@ -51,7 +51,7 @@
|
|||
"@toolwind/corner-shape": "^0.0.8-3",
|
||||
"@types/canvas-confetti": "^1.9.0",
|
||||
"@xyflow/react": "^12.10.0",
|
||||
"assistant-stream": "^0.3.2",
|
||||
"assistant-stream": "0.3.12",
|
||||
"canvas-confetti": "^1.9.4",
|
||||
"class-variance-authority": "^0.7.1",
|
||||
"clsx": "^2.1.1",
|
||||
|
|
|
|||
|
|
@ -37,13 +37,12 @@ import {
|
|||
Delete02Icon,
|
||||
Download03Icon,
|
||||
GemIcon,
|
||||
Globe02Icon,
|
||||
Search01Icon,
|
||||
PowerIcon,
|
||||
PencilEdit02Icon,
|
||||
LayoutAlignLeftIcon,
|
||||
HelpCircleIcon,
|
||||
Settings02Icon,
|
||||
SourceCodeSquareIcon,
|
||||
ZapIcon,
|
||||
} from "@hugeicons/core-free-icons";
|
||||
import {
|
||||
|
|
@ -528,7 +527,7 @@ export function AppSidebar() {
|
|||
</div>
|
||||
<div className="flex flex-col gap-0.5 leading-tight group-data-[collapsible=icon]:hidden">
|
||||
<span className="truncate font-heading text-[13px] tracking-[0.02em] font-semibold text-[#383835] dark:text-[#c7c7c4]">{displayTitle}</span>
|
||||
<span className="truncate text-[11px] tracking-[0.01em] text-muted-foreground">Unsloth</span>
|
||||
<span className="truncate text-[11px] tracking-[0.01em] text-muted-foreground">Studio</span>
|
||||
</div>
|
||||
<ChevronsUpDown strokeWidth={1.25} className="ml-auto size-4 text-muted-foreground group-data-[collapsible=icon]:hidden" />
|
||||
</SidebarMenuButton>
|
||||
|
|
@ -549,8 +548,8 @@ export function AppSidebar() {
|
|||
<DropdownMenuItem
|
||||
onSelect={() => useSettingsDialogStore.getState().openDialog("api-keys")}
|
||||
>
|
||||
<HugeiconsIcon icon={Globe02Icon} strokeWidth={1.75} className="size-[18px]" />
|
||||
<span>API</span>
|
||||
<HugeiconsIcon icon={SourceCodeSquareIcon} strokeWidth={1.75} className="size-[18px]" />
|
||||
<span>Developer</span>
|
||||
<span className="ml-auto rounded-[6px] border border-emerald-500/25 bg-emerald-500/10 px-1.5 py-0.5 text-[10px] leading-none font-semibold text-emerald-700 dark:text-emerald-300">
|
||||
New
|
||||
</span>
|
||||
|
|
@ -579,12 +578,6 @@ export function AppSidebar() {
|
|||
</DropdownMenuItem>
|
||||
</DropdownMenuGroup>
|
||||
<DropdownMenuSeparator className="mx-2.5! my-2.5! h-0! border-t border-border/70 bg-transparent!" />
|
||||
<DropdownMenuItem
|
||||
onSelect={() => useSettingsDialogStore.getState().openDialog("about")}
|
||||
>
|
||||
<HugeiconsIcon icon={HelpCircleIcon} strokeWidth={1.75} className="size-[18px]" />
|
||||
<span>Help</span>
|
||||
</DropdownMenuItem>
|
||||
<DropdownMenuItem onSelect={() => setShutdownOpen(true)}>
|
||||
<HugeiconsIcon icon={PowerIcon} strokeWidth={1.75} className="size-[18px]" />
|
||||
<span>Shutdown</span>
|
||||
|
|
|
|||
|
|
@ -50,7 +50,7 @@ export async function fetchDeviceType(): Promise<DeviceType> {
|
|||
if (res.ok) {
|
||||
const data = (await res.json()) as { device_type?: string; chat_only?: boolean };
|
||||
const deviceType = data.device_type ?? detectLocalPlatform();
|
||||
const chatOnly = data.chat_only ?? deviceType === "mac";
|
||||
const chatOnly = data.chat_only ?? false;
|
||||
usePlatformStore.setState({ deviceType, chatOnly, fetched: true });
|
||||
return deviceType;
|
||||
}
|
||||
|
|
|
|||
|
|
@ -9,8 +9,8 @@ import {
|
|||
ExportedMessageRepository,
|
||||
type ExportedMessageRepositoryItem,
|
||||
type PendingAttachment,
|
||||
RuntimeAdapterProvider,
|
||||
Suggestions,
|
||||
type LocalRuntimeOptions,
|
||||
type ThreadHistoryAdapter,
|
||||
type ThreadMessage,
|
||||
WebSpeechDictationAdapter,
|
||||
|
|
@ -583,9 +583,9 @@ function createDexieAdapter(
|
|||
};
|
||||
}
|
||||
|
||||
function ThreadHistoryProvider({
|
||||
children,
|
||||
}: { children?: ReactNode }): ReactElement {
|
||||
type StudioRuntimeAdapters = NonNullable<LocalRuntimeOptions["adapters"]>;
|
||||
|
||||
function useStudioRuntimeAdapters(): StudioRuntimeAdapters {
|
||||
const aui = useAui();
|
||||
|
||||
const history = useMemo<ThreadHistoryAdapter>(
|
||||
|
|
@ -711,17 +711,14 @@ function ThreadHistoryProvider({
|
|||
[history, dictation, attachments],
|
||||
);
|
||||
|
||||
return (
|
||||
<RuntimeAdapterProvider adapters={adapters}>
|
||||
{children}
|
||||
</RuntimeAdapterProvider>
|
||||
);
|
||||
return adapters;
|
||||
}
|
||||
|
||||
const chatAdapter = createOpenAIStreamAdapter();
|
||||
|
||||
function useRuntimeHook(): ReturnType<typeof useLocalRuntime> {
|
||||
return useLocalRuntime(chatAdapter);
|
||||
const adapters = useStudioRuntimeAdapters();
|
||||
return useLocalRuntime(chatAdapter, { adapters });
|
||||
}
|
||||
|
||||
function ThreadAutoSwitch({
|
||||
|
|
@ -898,10 +895,7 @@ export function ChatRuntimeProvider({
|
|||
}): ReactElement {
|
||||
const runtime = useRemoteThreadListRuntime({
|
||||
runtimeHook: useRuntimeHook,
|
||||
adapter: {
|
||||
...createDexieAdapter(modelType, pairId),
|
||||
unstable_Provider: ThreadHistoryProvider,
|
||||
},
|
||||
adapter: createDexieAdapter(modelType, pairId),
|
||||
});
|
||||
|
||||
const aui = useAui({
|
||||
|
|
|
|||
|
|
@ -10,11 +10,11 @@ import {
|
|||
import { cn } from "@/lib/utils";
|
||||
import {
|
||||
Cancel01Icon,
|
||||
Globe02Icon,
|
||||
HelpCircleIcon,
|
||||
Message01Icon,
|
||||
PaintBrush02Icon,
|
||||
Settings02Icon,
|
||||
SourceCodeSquareIcon,
|
||||
SparklesIcon,
|
||||
UserIcon,
|
||||
} from "@hugeicons/core-free-icons";
|
||||
import { HugeiconsIcon } from "@hugeicons/react";
|
||||
|
|
@ -40,8 +40,8 @@ const TABS: TabDef[] = [
|
|||
{ id: "profile", label: "Profile", icon: UserIcon },
|
||||
{ id: "appearance", label: "Appearance", icon: PaintBrush02Icon },
|
||||
{ id: "chat", label: "Chat", icon: Message01Icon },
|
||||
{ id: "api-keys", label: "API", icon: Globe02Icon, badge: "New" },
|
||||
{ id: "about", label: "Help", icon: HelpCircleIcon },
|
||||
{ id: "api-keys", label: "Developer", icon: SourceCodeSquareIcon, badge: "New" },
|
||||
{ id: "about", label: "Help", icon: SparklesIcon },
|
||||
];
|
||||
|
||||
function renderTab(tab: SettingsTab) {
|
||||
|
|
|
|||
|
|
@ -63,7 +63,7 @@ export function ApiKeysTab() {
|
|||
return (
|
||||
<div className="flex flex-col gap-6">
|
||||
<header className="flex flex-col gap-1">
|
||||
<h1 className="text-lg font-semibold font-heading">API</h1>
|
||||
<h1 className="text-lg font-semibold font-heading">Developer</h1>
|
||||
<p className="text-xs text-muted-foreground">
|
||||
Access Unsloth programmatically via the OpenAI-compatible API.{" "}
|
||||
<a
|
||||
|
|
|
|||
|
|
@ -1,6 +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
|
||||
|
||||
import { usePlatformStore } from "@/config/env";
|
||||
import { SectionCard } from "@/components/section-card";
|
||||
import { Checkbox } from "@/components/ui/checkbox";
|
||||
import {
|
||||
|
|
@ -124,6 +125,7 @@ function SliderRow({
|
|||
|
||||
export function ParamsSection(): ReactElement {
|
||||
const store = useTrainingConfigStore();
|
||||
const platformDeviceType = usePlatformStore((s) => s.deviceType);
|
||||
const isLora = store.trainingMethod !== "full";
|
||||
const showVisionLora = store.isVisionModel && store.isDatasetImage === true;
|
||||
const [loraOpen, setLoraOpen] = useState(false);
|
||||
|
|
@ -883,7 +885,11 @@ export function ParamsSection(): ReactElement {
|
|||
<SelectContent>
|
||||
<SelectItem value="none">None</SelectItem>
|
||||
<SelectItem value="true">Standard</SelectItem>
|
||||
<SelectItem value="unsloth">Unsloth</SelectItem>
|
||||
{platformDeviceType === "mac" ? (
|
||||
<SelectItem value="mlx">MLX</SelectItem>
|
||||
) : (
|
||||
<SelectItem value="unsloth">Unsloth</SelectItem>
|
||||
)}
|
||||
</SelectContent>
|
||||
</Select>
|
||||
</Row>
|
||||
|
|
|
|||
|
|
@ -3,6 +3,7 @@
|
|||
|
||||
import type { BackendModelConfig } from "../api/models-api";
|
||||
import type { TrainingConfigState } from "../types/config";
|
||||
import { usePlatformStore } from "@/config/env";
|
||||
|
||||
type ModelDefaultsPatch = Partial<
|
||||
Pick<
|
||||
|
|
@ -69,7 +70,13 @@ function toStringArray(value: unknown): string[] | undefined {
|
|||
function toGradientCheckpointing(
|
||||
value: unknown,
|
||||
): TrainingConfigState["gradientCheckpointing"] | undefined {
|
||||
if (value === "none" || value === "true" || value === "unsloth") return value;
|
||||
if (value === "none" || value === "true" || value === "unsloth" || value === "mlx") {
|
||||
// On Mac, map "unsloth" → "mlx" since Unsloth GC is GPU-only
|
||||
if (usePlatformStore.getState().deviceType === "mac" && value === "unsloth") {
|
||||
return "mlx";
|
||||
}
|
||||
return value;
|
||||
}
|
||||
return undefined;
|
||||
}
|
||||
|
||||
|
|
|
|||
|
|
@ -6,6 +6,7 @@ import { listModels } from "@huggingface/hub";
|
|||
import { type CachedResult, cachedModelInfo, primeCacheFromListing } from "@/lib/hf-cache";
|
||||
import { useCallback, useMemo } from "react";
|
||||
import { useHfPaginatedSearch } from "./use-hf-paginated-search";
|
||||
import { usePlatformStore } from "@/config/env";
|
||||
|
||||
export interface HfModelResult {
|
||||
id: string;
|
||||
|
|
@ -16,7 +17,8 @@ export interface HfModelResult {
|
|||
isGguf: boolean;
|
||||
}
|
||||
|
||||
const EXCLUDED_TAGS = new Set([
|
||||
/** Tags to exclude on GPU (CUDA/ROCm) — MLX models won't load on GPU. */
|
||||
const EXCLUDED_TAGS_GPU = new Set([
|
||||
"gptq",
|
||||
"awq",
|
||||
"exl2",
|
||||
|
|
@ -28,6 +30,18 @@ const EXCLUDED_TAGS = new Set([
|
|||
"ctranslate2",
|
||||
]);
|
||||
|
||||
/** Tags to exclude on MLX (Mac) — GPU-only quant formats won't load on MLX. */
|
||||
const EXCLUDED_TAGS_MLX = new Set([
|
||||
"gptq",
|
||||
"awq",
|
||||
"exl2",
|
||||
"onnx",
|
||||
"openvino",
|
||||
"coreml",
|
||||
"tflite",
|
||||
"ctranslate2",
|
||||
]);
|
||||
|
||||
// Embedding / sentence-transformer models ship with onnx/openvino as additional
|
||||
// export formats — they should not be excluded by the tag check above.
|
||||
const EMBEDDING_TAGS = new Set([
|
||||
|
|
@ -77,7 +91,7 @@ function estimateSizeFromDtypes(
|
|||
return total > 0 ? total : undefined;
|
||||
}
|
||||
|
||||
function makeMapModel(excludeGguf: boolean) {
|
||||
function makeMapModel(excludeGguf: boolean, excludedTags: Set<string>) {
|
||||
return (raw: unknown): HfModelResult | null => {
|
||||
const m = raw as {
|
||||
name: string;
|
||||
|
|
@ -87,7 +101,7 @@ function makeMapModel(excludeGguf: boolean) {
|
|||
tags?: string[];
|
||||
};
|
||||
const isEmbedding = m.tags?.some((t) => EMBEDDING_TAGS.has(t));
|
||||
if (!isEmbedding && m.tags?.some((t) => EXCLUDED_TAGS.has(t))) {
|
||||
if (!isEmbedding && m.tags?.some((t) => excludedTags.has(t))) {
|
||||
return null;
|
||||
}
|
||||
const isGguf =
|
||||
|
|
@ -314,7 +328,9 @@ export function useHfModelSearch(
|
|||
[trimmed, searchQuery, pinnedId, task, accessToken, priorityIds],
|
||||
);
|
||||
|
||||
const mapModel = useMemo(() => makeMapModel(excludeGguf), [excludeGguf]);
|
||||
const deviceType = usePlatformStore((s) => s.deviceType);
|
||||
const excludedTags = deviceType === "mac" ? EXCLUDED_TAGS_MLX : EXCLUDED_TAGS_GPU;
|
||||
const mapModel = useMemo(() => makeMapModel(excludeGguf, excludedTags), [excludeGguf, excludedTags]);
|
||||
const search = useHfPaginatedSearch(createIter, mapModel);
|
||||
|
||||
// Secondary sort guarantee: unsloth models always float to the top.
|
||||
|
|
|
|||
|
|
@ -10,7 +10,7 @@ export function isAdapterMethod(method: TrainingMethod): boolean {
|
|||
export type StepNumber = 1 | 2 | 3 | 4 | 5;
|
||||
export type DatasetSource = "huggingface" | "upload";
|
||||
export type DatasetFormat = "auto" | "alpaca" | "chatml" | "sharegpt";
|
||||
export type GradientCheckpointing = "none" | "true" | "unsloth";
|
||||
export type GradientCheckpointing = "none" | "true" | "unsloth" | "mlx";
|
||||
|
||||
export interface WizardState {
|
||||
currentStep: StepNumber;
|
||||
|
|
|
|||
211
studio/setup.ps1
211
studio/setup.ps1
|
|
@ -1492,9 +1492,79 @@ if (-not $PythonCmd) {
|
|||
|
||||
substep "Using $PythonCmd ($(& $PythonCmd --version 2>&1))"
|
||||
|
||||
# The venv must already exist (created by install.ps1).
|
||||
# This script (setup.ps1 / "unsloth studio update") only updates packages.
|
||||
$VenvDir = Join-Path $env:USERPROFILE ".unsloth\studio\unsloth_studio"
|
||||
# The venv must already exist (created by install.ps1); this script only
|
||||
# updates packages. UNSLOTH_STUDIO_HOME (or STUDIO_HOME alias) overrides the
|
||||
# root. UNSLOTH_STUDIO_HOME wins when both are set. Whitespace-only values
|
||||
# are treated as unset to match Python .strip() semantics.
|
||||
$_studioOverrideVar = $null
|
||||
$_studioOverride = $null
|
||||
if (-not [string]::IsNullOrWhiteSpace($env:UNSLOTH_STUDIO_HOME)) {
|
||||
$_studioOverrideVar = "UNSLOTH_STUDIO_HOME"
|
||||
$_studioOverride = $env:UNSLOTH_STUDIO_HOME.Trim()
|
||||
} elseif (-not [string]::IsNullOrWhiteSpace($env:STUDIO_HOME)) {
|
||||
$_studioOverrideVar = "STUDIO_HOME"
|
||||
$_studioOverride = $env:STUDIO_HOME.Trim()
|
||||
}
|
||||
if ($_studioOverride) {
|
||||
if ($_studioOverride -eq "~" -or $_studioOverride -like "~/*" -or $_studioOverride -like "~\*") {
|
||||
$_studioOverride = (Join-Path $env:USERPROFILE $_studioOverride.Substring(1).TrimStart('/','\'))
|
||||
}
|
||||
if (Test-Path -LiteralPath $_studioOverride -PathType Container) {
|
||||
$StudioHome = (Resolve-Path -LiteralPath $_studioOverride).Path
|
||||
# why: mirror setup.sh:417 and install.ps1:130 -- fail fast when the
|
||||
# custom root is read-only instead of erroring later while creating
|
||||
# sidecar venvs / installing packages.
|
||||
$_setupWriteProbe = Join-Path $StudioHome (".unsloth-write-probe-" + [guid]::NewGuid())
|
||||
try {
|
||||
[System.IO.File]::WriteAllText($_setupWriteProbe, "")
|
||||
Remove-Item -LiteralPath $_setupWriteProbe -Force -ErrorAction SilentlyContinue
|
||||
} catch {
|
||||
Write-Host "ERROR: $_studioOverrideVar=$StudioHome is not writable." -ForegroundColor Red
|
||||
exit 1
|
||||
}
|
||||
} else {
|
||||
Write-Host "ERROR: $_studioOverrideVar=$_studioOverride does not exist." -ForegroundColor Red
|
||||
Write-Host " Run install.ps1 to create the install root before 'unsloth studio update'." -ForegroundColor Red
|
||||
exit 1
|
||||
}
|
||||
} else {
|
||||
$StudioHome = Join-Path $env:USERPROFILE ".unsloth\studio"
|
||||
}
|
||||
$VenvDir = Join-Path $StudioHome "unsloth_studio"
|
||||
|
||||
# why: in env-override mode $StudioHome is user-chosen; require the
|
||||
# ownership marker before Remove-Item so unrelated dirs survive. Gated on
|
||||
# the canonical comparison so an override pointing at the legacy default
|
||||
# still behaves like a default install.
|
||||
$StudioOwnedMarker = ".unsloth-studio-owned"
|
||||
$LegacyStudioHome = Join-Path $env:USERPROFILE ".unsloth\studio"
|
||||
$_studioHomeCanon = $StudioHome
|
||||
if (Test-Path -LiteralPath $_studioHomeCanon -PathType Container) {
|
||||
$_studioHomeCanon = (Resolve-Path -LiteralPath $_studioHomeCanon).Path
|
||||
}
|
||||
if (Test-Path -LiteralPath $LegacyStudioHome -PathType Container) {
|
||||
$LegacyStudioHome = (Resolve-Path -LiteralPath $LegacyStudioHome).Path
|
||||
}
|
||||
$StudioHomeIsCustom = ($_studioHomeCanon -ne $LegacyStudioHome)
|
||||
function Assert-StudioOwnedOrAbsent {
|
||||
param(
|
||||
[Parameter(Mandatory = $true)][string]$Path,
|
||||
[Parameter(Mandatory = $true)][string]$Label
|
||||
)
|
||||
if (-not (Test-Path -LiteralPath $Path -PathType Container)) { return }
|
||||
if ($StudioHomeIsCustom -and -not (Test-Path -LiteralPath (Join-Path $Path $StudioOwnedMarker) -PathType Leaf)) {
|
||||
Write-Host "[ERROR] $Path already exists and is not marked as a Studio-owned $Label." -ForegroundColor Red
|
||||
Write-Host " Move it aside or choose an empty UNSLOTH_STUDIO_HOME before re-running." -ForegroundColor Yellow
|
||||
exit 1
|
||||
}
|
||||
}
|
||||
function Mark-StudioOwned {
|
||||
param([Parameter(Mandatory = $true)][string]$Path)
|
||||
if (-not (Test-Path -LiteralPath $Path -PathType Container)) { return }
|
||||
try {
|
||||
[System.IO.File]::WriteAllText((Join-Path $Path $StudioOwnedMarker), "")
|
||||
} catch {}
|
||||
}
|
||||
|
||||
# Stale-venv detection: if the venv exists but its torch flavor no longer
|
||||
# matches the current machine, repair according to invocation context.
|
||||
|
|
@ -1504,12 +1574,12 @@ $VenvDir = Join-Path $env:USERPROFILE ".unsloth\studio\unsloth_studio"
|
|||
# In no-torch mode, a missing torch package is expected.
|
||||
$NoTorchMode = $env:UNSLOTH_NO_TORCH -match '^(?i:true|1|yes)$'
|
||||
$InstallerManagedSetup = $env:UNSLOTH_INSTALL_ROLLBACK_MANAGED -match '^(?i:true|1|yes)$'
|
||||
if ((Test-Path $VenvDir -PathType Container) -and -not $NoTorchMode) {
|
||||
if ((Test-Path -LiteralPath $VenvDir -PathType Container) -and -not $NoTorchMode) {
|
||||
$VenvPyExe = Join-Path $VenvDir "Scripts\python.exe"
|
||||
$installedTorchTag = $null
|
||||
$shouldRebuild = $false
|
||||
|
||||
if (Test-Path $VenvPyExe) {
|
||||
if (Test-Path -LiteralPath $VenvPyExe) {
|
||||
try {
|
||||
$psi = New-Object System.Diagnostics.ProcessStartInfo
|
||||
$psi.FileName = $VenvPyExe
|
||||
|
|
@ -1558,8 +1628,21 @@ if ((Test-Path $VenvDir -PathType Container) -and -not $NoTorchMode) {
|
|||
exit 1
|
||||
}
|
||||
substep "Stale venv detected ($reason) -- rebuilding..." "Yellow"
|
||||
# why: mirror install.ps1 env-mode guard so an update against a custom
|
||||
# UNSLOTH_STUDIO_HOME never wipes an unrelated unsloth_studio venv;
|
||||
# -PathType Leaf rejects a directory masquerading as the sentinel.
|
||||
if (
|
||||
$StudioHomeIsCustom -and
|
||||
-not (Test-Path -LiteralPath (Join-Path $VenvDir $StudioOwnedMarker) -PathType Leaf) -and
|
||||
-not (Test-Path -LiteralPath (Join-Path $StudioHome "share\studio.conf") -PathType Leaf) -and
|
||||
-not (Test-Path -LiteralPath (Join-Path $StudioHome "bin\unsloth.exe") -PathType Leaf)
|
||||
) {
|
||||
Write-Host "[ERROR] $VenvDir already exists but does not look like an Unsloth Studio install." -ForegroundColor Red
|
||||
Write-Host " Move it aside or choose an empty UNSLOTH_STUDIO_HOME before re-running." -ForegroundColor Yellow
|
||||
exit 1
|
||||
}
|
||||
try {
|
||||
Remove-Item $VenvDir -Recurse -Force -ErrorAction Stop
|
||||
Remove-Item -LiteralPath $VenvDir -Recurse -Force -ErrorAction Stop
|
||||
} catch {
|
||||
Write-Host " [ERROR] Could not remove stale venv: $($_.Exception.Message)" -ForegroundColor Red
|
||||
Write-Host " Close any running Studio/Python processes and re-run setup." -ForegroundColor Red
|
||||
|
|
@ -1568,7 +1651,7 @@ if ((Test-Path $VenvDir -PathType Container) -and -not $NoTorchMode) {
|
|||
}
|
||||
}
|
||||
|
||||
if (-not (Test-Path $VenvDir)) {
|
||||
if (-not (Test-Path -LiteralPath $VenvDir)) {
|
||||
Write-Host "[ERROR] Virtual environment not found at $VenvDir" -ForegroundColor Red
|
||||
Write-Host " Run install.ps1 first to create the environment:" -ForegroundColor Yellow
|
||||
Write-Host " irm https://unsloth.ai/install.ps1 | iex" -ForegroundColor Yellow
|
||||
|
|
@ -1759,17 +1842,19 @@ if ($stackExit -ne 0) {
|
|||
# ── Pre-install transformers 5.x into .venv_t5_530/ and .venv_t5_550/ ──
|
||||
# Runs outside the deps fast-path gate so that upgrades from the legacy
|
||||
# single .venv_t5 are always migrated to the tiered layout.
|
||||
$VenvT5_530Dir = Join-Path $env:USERPROFILE ".unsloth\studio\.venv_t5_530"
|
||||
$VenvT5_550Dir = Join-Path $env:USERPROFILE ".unsloth\studio\.venv_t5_550"
|
||||
$VenvT5Legacy = Join-Path $env:USERPROFILE ".unsloth\studio\.venv_t5"
|
||||
# T5 sidecar venvs live under the resolved $StudioHome so custom installs are self-contained.
|
||||
$VenvT5_530Dir = Join-Path $StudioHome ".venv_t5_530"
|
||||
$VenvT5_550Dir = Join-Path $StudioHome ".venv_t5_550"
|
||||
$VenvT5Legacy = Join-Path $StudioHome ".venv_t5"
|
||||
|
||||
$_NeedT5Install = $false
|
||||
if (Test-Path $VenvT5Legacy) {
|
||||
Remove-Item -Recurse -Force $VenvT5Legacy
|
||||
if (Test-Path -LiteralPath $VenvT5Legacy) {
|
||||
Assert-StudioOwnedOrAbsent -Path $VenvT5Legacy -Label "legacy transformers sidecar venv"
|
||||
Remove-Item -LiteralPath $VenvT5Legacy -Recurse -Force
|
||||
$_NeedT5Install = $true
|
||||
}
|
||||
if (-not (Test-Path $VenvT5_530Dir)) { $_NeedT5Install = $true }
|
||||
if (-not (Test-Path $VenvT5_550Dir)) { $_NeedT5Install = $true }
|
||||
if (-not (Test-Path -LiteralPath $VenvT5_530Dir)) { $_NeedT5Install = $true }
|
||||
if (-not (Test-Path -LiteralPath $VenvT5_550Dir)) { $_NeedT5Install = $true }
|
||||
# Also reinstall when python deps were updated
|
||||
if (-not $SkipPythonDeps) { $_NeedT5Install = $true }
|
||||
|
||||
|
|
@ -1781,8 +1866,10 @@ $ErrorActionPreference = "Continue"
|
|||
|
||||
# --- .venv_t5_530 (transformers 5.3.0) ---
|
||||
substep "pre-installing transformers 5.3.0 for newer model support..."
|
||||
if (Test-Path $VenvT5_530Dir) { Remove-Item -Recurse -Force $VenvT5_530Dir }
|
||||
New-Item -ItemType Directory -Path $VenvT5_530Dir -Force | Out-Null
|
||||
Assert-StudioOwnedOrAbsent -Path $VenvT5_530Dir -Label "transformers 5.3 sidecar venv"
|
||||
if (Test-Path -LiteralPath $VenvT5_530Dir) { Remove-Item -LiteralPath $VenvT5_530Dir -Recurse -Force }
|
||||
[System.IO.Directory]::CreateDirectory($VenvT5_530Dir) | Out-Null
|
||||
Mark-StudioOwned -Path $VenvT5_530Dir
|
||||
foreach ($pkg in @("transformers==5.3.0", "huggingface_hub==1.8.0", "hf_xet==1.4.2")) {
|
||||
if ($script:UnslothVerbose) {
|
||||
Fast-Install --target $VenvT5_530Dir --no-deps $pkg
|
||||
|
|
@ -1814,8 +1901,10 @@ step "transformers" "5.3.0 pre-installed"
|
|||
|
||||
# --- .venv_t5_550 (transformers 5.5.0) ---
|
||||
substep "pre-installing transformers 5.5.0 for Gemma 4 support..."
|
||||
if (Test-Path $VenvT5_550Dir) { Remove-Item -Recurse -Force $VenvT5_550Dir }
|
||||
New-Item -ItemType Directory -Path $VenvT5_550Dir -Force | Out-Null
|
||||
Assert-StudioOwnedOrAbsent -Path $VenvT5_550Dir -Label "transformers 5.5 sidecar venv"
|
||||
if (Test-Path -LiteralPath $VenvT5_550Dir) { Remove-Item -LiteralPath $VenvT5_550Dir -Recurse -Force }
|
||||
[System.IO.Directory]::CreateDirectory($VenvT5_550Dir) | Out-Null
|
||||
Mark-StudioOwned -Path $VenvT5_550Dir
|
||||
foreach ($pkg in @("transformers==5.5.0", "huggingface_hub==1.8.0", "hf_xet==1.4.2")) {
|
||||
if ($script:UnslothVerbose) {
|
||||
Fast-Install --target $VenvT5_550Dir --no-deps $pkg
|
||||
|
|
@ -1851,8 +1940,15 @@ step "transformers" "5.5.0 pre-installed"
|
|||
# ==========================================================================
|
||||
# PHASE 3.4: Prefer prebuilt llama.cpp bundles before source build
|
||||
# ==========================================================================
|
||||
$UnslothHome = Join-Path $env:USERPROFILE ".unsloth"
|
||||
if (-not (Test-Path $UnslothHome)) { New-Item -ItemType Directory -Force $UnslothHome | Out-Null }
|
||||
# Nest llama.cpp under $StudioHome only for real env-overrides, never the
|
||||
# legacy default. Reuses $StudioHomeIsCustom from the canonical comparison
|
||||
# computed above so the llama.cpp nest matches ownership-guard semantics.
|
||||
if ($StudioHomeIsCustom) {
|
||||
$UnslothHome = $StudioHome
|
||||
} else {
|
||||
$UnslothHome = Join-Path $env:USERPROFILE ".unsloth"
|
||||
}
|
||||
if (-not (Test-Path -LiteralPath $UnslothHome)) { [System.IO.Directory]::CreateDirectory($UnslothHome) | Out-Null }
|
||||
$LlamaCppDir = Join-Path $UnslothHome "llama.cpp"
|
||||
$NeedLlamaSourceBuild = $false
|
||||
$SkipPrebuiltInstall = $false
|
||||
|
|
@ -1954,9 +2050,15 @@ if ($env:UNSLOTH_LLAMA_FORCE_COMPILE -eq "1") {
|
|||
} else {
|
||||
Write-Host ""
|
||||
substep "installing prebuilt llama.cpp bundle (preferred path)..."
|
||||
if (Test-Path $LlamaCppDir) {
|
||||
if (Test-Path -LiteralPath $LlamaCppDir) {
|
||||
substep "Existing llama.cpp install detected -- validating staged prebuilt update before replacement"
|
||||
}
|
||||
# why: install_llama_prebuilt.py uses os.replace(), which would displace
|
||||
# an unrelated $env:UNSLOTH_STUDIO_HOME\llama.cpp before the source-build
|
||||
# ownership check below ever runs.
|
||||
if ($StudioHomeIsCustom) {
|
||||
Assert-StudioOwnedOrAbsent -Path $LlamaCppDir -Label "llama.cpp install"
|
||||
}
|
||||
$prebuiltArgs = @(
|
||||
"$PSScriptRoot\install_llama_prebuilt.py",
|
||||
"--install-dir", $LlamaCppDir,
|
||||
|
|
@ -2001,6 +2103,9 @@ if ($env:UNSLOTH_LLAMA_FORCE_COMPILE -eq "1") {
|
|||
} else {
|
||||
step "llama.cpp" "prebuilt installed and validated"
|
||||
}
|
||||
if ($StudioHomeIsCustom -and (Test-Path -LiteralPath $LlamaCppDir -PathType Container)) {
|
||||
Mark-StudioOwned -Path $LlamaCppDir
|
||||
}
|
||||
$installedRelease = Get-InstalledLlamaPrebuiltRelease -InstallDir $LlamaCppDir
|
||||
if ($installedRelease) {
|
||||
substep $installedRelease
|
||||
|
|
@ -2008,7 +2113,7 @@ if ($env:UNSLOTH_LLAMA_FORCE_COMPILE -eq "1") {
|
|||
} elseif ($prebuiltExit -eq 3) {
|
||||
step "llama.cpp" "install blocked by active llama.cpp process" "Yellow"
|
||||
Write-LlamaFailureLog -Output $prebuiltOutput
|
||||
if (Test-Path $LlamaCppDir) {
|
||||
if (Test-Path -LiteralPath $LlamaCppDir) {
|
||||
substep "Existing install was restored" "Yellow"
|
||||
}
|
||||
substep "Close Studio or other llama.cpp users and retry" "Yellow"
|
||||
|
|
@ -2016,7 +2121,7 @@ if ($env:UNSLOTH_LLAMA_FORCE_COMPILE -eq "1") {
|
|||
} else {
|
||||
step "llama.cpp" "prebuilt install failed (continuing)" "Yellow"
|
||||
Write-LlamaFailureLog -Output $prebuiltOutput
|
||||
if (Test-Path $LlamaCppDir) {
|
||||
if (Test-Path -LiteralPath $LlamaCppDir) {
|
||||
substep "Prebuilt update failed; existing install was restored or cleaned before source build fallback" "Yellow"
|
||||
}
|
||||
substep "Prebuilt llama.cpp path unavailable or failed validation -- falling back to source build" "Yellow"
|
||||
|
|
@ -2092,10 +2197,10 @@ $HasCmakeForBuild = $null -ne (Get-Command cmake -ErrorAction SilentlyContinue)
|
|||
# Check if existing llama-server matches current GPU mode. A CUDA-built binary
|
||||
# on a now-CPU-only machine (or vice versa) needs to be rebuilt.
|
||||
$NeedRebuild = $false
|
||||
if (Test-Path $LlamaServerBin) {
|
||||
if (Test-Path -LiteralPath $LlamaServerBin) {
|
||||
$CmakeCacheFile = Join-Path $BuildDir "CMakeCache.txt"
|
||||
if (Test-Path $CmakeCacheFile) {
|
||||
$cachedCuda = Select-String -Path $CmakeCacheFile -Pattern 'GGML_CUDA:BOOL=ON' -Quiet
|
||||
if (Test-Path -LiteralPath $CmakeCacheFile) {
|
||||
$cachedCuda = Select-String -LiteralPath $CmakeCacheFile -Pattern 'GGML_CUDA:BOOL=ON' -Quiet
|
||||
if ($HasNvidiaSmi -and -not $cachedCuda) {
|
||||
Write-Host " Existing llama-server is CPU-only but GPU is available -- rebuilding" -ForegroundColor Yellow
|
||||
$NeedRebuild = $true
|
||||
|
|
@ -2109,7 +2214,7 @@ if (Test-Path $LlamaServerBin) {
|
|||
if (-not $NeedLlamaSourceBuild) {
|
||||
Write-Host ""
|
||||
step "llama.cpp" "prebuilt (validated)"
|
||||
} elseif ((Test-Path $LlamaServerBin) -and -not $NeedRebuild -and $RequestedLlamaTag -ne "master") {
|
||||
} elseif ((Test-Path -LiteralPath $LlamaServerBin) -and -not $NeedRebuild -and $RequestedLlamaTag -ne "master") {
|
||||
# Skip rebuild only for pinned tags (e.g. b8635). When the requested
|
||||
# tag is "master" (a moving target), always rebuild so the binary picks
|
||||
# up new model architecture support (e.g. Gemma 4).
|
||||
|
|
@ -2211,7 +2316,13 @@ if (-not $NeedLlamaSourceBuild) {
|
|||
|
||||
$UseConcreteRef = ($ResolvedSourceRef -ne "latest" -and -not [string]::IsNullOrWhiteSpace($ResolvedSourceRef))
|
||||
|
||||
if (Test-Path (Join-Path $LlamaCppDir ".git")) {
|
||||
if (Test-Path -LiteralPath (Join-Path $LlamaCppDir ".git")) {
|
||||
# why: in-place git mutation (remote set-url, checkout -B, clean -fdx)
|
||||
# rewrites $LlamaCppDir; mirror the prebuilt and temp-dir-swap guards
|
||||
# so an unrelated workspace .git tree is never silently overwritten.
|
||||
if ($StudioHomeIsCustom) {
|
||||
Assert-StudioOwnedOrAbsent -Path $LlamaCppDir -Label "llama.cpp install"
|
||||
}
|
||||
Write-Host " Syncing llama.cpp to $ResolvedSourceRef..." -ForegroundColor Gray
|
||||
# Always sync the remote URL so switching between default/fork sources works
|
||||
Invoke-SetupCommand -AlwaysQuiet { git -C $LlamaCppDir remote set-url origin "$ResolvedSourceUrl.git" } | Out-Null
|
||||
|
|
@ -2282,24 +2393,30 @@ if (-not $NeedLlamaSourceBuild) {
|
|||
}
|
||||
}
|
||||
}
|
||||
# why: in-place git-sync (the temp-dir clone path calls Mark-StudioOwned
|
||||
# at swap-time) must mark the existing tree so a subsequent prebuilt
|
||||
# update path's Assert-StudioOwnedOrAbsent does not exit on the same root.
|
||||
if ($BuildOk -and $StudioHomeIsCustom) {
|
||||
Mark-StudioOwned -Path $LlamaCppDir
|
||||
}
|
||||
} else {
|
||||
Write-Host " Cloning llama.cpp @ $ResolvedSourceRef..." -ForegroundColor Gray
|
||||
$buildTmp = "$LlamaCppDir.build.$PID"
|
||||
$null = New-Item -ItemType Directory -Force -Path (Split-Path $LlamaCppDir -Parent)
|
||||
if (Test-Path $buildTmp) { Remove-Item -Recurse -Force $buildTmp }
|
||||
$null = [System.IO.Directory]::CreateDirectory((Split-Path -LiteralPath $LlamaCppDir))
|
||||
if (Test-Path -LiteralPath $buildTmp) { Remove-Item -LiteralPath $buildTmp -Recurse -Force }
|
||||
if ($LlamaPr) {
|
||||
$cloneExit = Invoke-SetupCommand -AlwaysQuiet { git clone --depth 1 "$LlamaSource.git" $buildTmp }
|
||||
if ($cloneExit -ne 0) {
|
||||
$BuildOk = $false
|
||||
$FailedStep = "git clone"
|
||||
if (Test-Path $buildTmp) { Remove-Item -Recurse -Force $buildTmp }
|
||||
if (Test-Path -LiteralPath $buildTmp) { Remove-Item -LiteralPath $buildTmp -Recurse -Force }
|
||||
}
|
||||
if ($BuildOk) {
|
||||
$fetchExit = Invoke-SetupCommand -AlwaysQuiet { git -C $buildTmp fetch --depth 1 origin "pull/$LlamaPr/head:pr-$LlamaPr" }
|
||||
if ($fetchExit -ne 0) {
|
||||
$BuildOk = $false
|
||||
$FailedStep = "git fetch PR #$LlamaPr"
|
||||
if (Test-Path $buildTmp) { Remove-Item -Recurse -Force $buildTmp }
|
||||
if (Test-Path -LiteralPath $buildTmp) { Remove-Item -LiteralPath $buildTmp -Recurse -Force }
|
||||
}
|
||||
}
|
||||
if ($BuildOk) {
|
||||
|
|
@ -2307,7 +2424,7 @@ if (-not $NeedLlamaSourceBuild) {
|
|||
if ($checkoutExit -ne 0) {
|
||||
$BuildOk = $false
|
||||
$FailedStep = "git checkout PR #$LlamaPr"
|
||||
if (Test-Path $buildTmp) { Remove-Item -Recurse -Force $buildTmp }
|
||||
if (Test-Path -LiteralPath $buildTmp) { Remove-Item -LiteralPath $buildTmp -Recurse -Force }
|
||||
}
|
||||
}
|
||||
} elseif ($ResolvedSourceRefKind -eq "pull") {
|
||||
|
|
@ -2315,14 +2432,14 @@ if (-not $NeedLlamaSourceBuild) {
|
|||
if ($cloneExit -ne 0) {
|
||||
$BuildOk = $false
|
||||
$FailedStep = "git clone"
|
||||
if (Test-Path $buildTmp) { Remove-Item -Recurse -Force $buildTmp }
|
||||
if (Test-Path -LiteralPath $buildTmp) { Remove-Item -LiteralPath $buildTmp -Recurse -Force }
|
||||
}
|
||||
if ($BuildOk) {
|
||||
$fetchExit = Invoke-SetupCommand -AlwaysQuiet { git -C $buildTmp fetch --depth 1 origin $ResolvedSourceRef }
|
||||
if ($fetchExit -ne 0) {
|
||||
$BuildOk = $false
|
||||
$FailedStep = "git fetch source PR ref"
|
||||
if (Test-Path $buildTmp) { Remove-Item -Recurse -Force $buildTmp }
|
||||
if (Test-Path -LiteralPath $buildTmp) { Remove-Item -LiteralPath $buildTmp -Recurse -Force }
|
||||
}
|
||||
}
|
||||
if ($BuildOk) {
|
||||
|
|
@ -2330,7 +2447,7 @@ if (-not $NeedLlamaSourceBuild) {
|
|||
if ($checkoutExit -ne 0) {
|
||||
$BuildOk = $false
|
||||
$FailedStep = "git checkout source PR ref"
|
||||
if (Test-Path $buildTmp) { Remove-Item -Recurse -Force $buildTmp }
|
||||
if (Test-Path -LiteralPath $buildTmp) { Remove-Item -LiteralPath $buildTmp -Recurse -Force }
|
||||
}
|
||||
}
|
||||
} elseif ($ResolvedSourceRefKind -eq "commit") {
|
||||
|
|
@ -2338,14 +2455,14 @@ if (-not $NeedLlamaSourceBuild) {
|
|||
if ($cloneExit -ne 0) {
|
||||
$BuildOk = $false
|
||||
$FailedStep = "git clone"
|
||||
if (Test-Path $buildTmp) { Remove-Item -Recurse -Force $buildTmp }
|
||||
if (Test-Path -LiteralPath $buildTmp) { Remove-Item -LiteralPath $buildTmp -Recurse -Force }
|
||||
}
|
||||
if ($BuildOk) {
|
||||
$fetchExit = Invoke-SetupCommand -AlwaysQuiet { git -C $buildTmp fetch --depth 1 origin $ResolvedSourceRef }
|
||||
if ($fetchExit -ne 0) {
|
||||
$BuildOk = $false
|
||||
$FailedStep = "git fetch source commit"
|
||||
if (Test-Path $buildTmp) { Remove-Item -Recurse -Force $buildTmp }
|
||||
if (Test-Path -LiteralPath $buildTmp) { Remove-Item -LiteralPath $buildTmp -Recurse -Force }
|
||||
}
|
||||
}
|
||||
if ($BuildOk) {
|
||||
|
|
@ -2353,7 +2470,7 @@ if (-not $NeedLlamaSourceBuild) {
|
|||
if ($checkoutExit -ne 0) {
|
||||
$BuildOk = $false
|
||||
$FailedStep = "git checkout source commit"
|
||||
if (Test-Path $buildTmp) { Remove-Item -Recurse -Force $buildTmp }
|
||||
if (Test-Path -LiteralPath $buildTmp) { Remove-Item -LiteralPath $buildTmp -Recurse -Force }
|
||||
}
|
||||
}
|
||||
} else {
|
||||
|
|
@ -2366,7 +2483,7 @@ if (-not $NeedLlamaSourceBuild) {
|
|||
if ($cloneExit -ne 0) {
|
||||
$BuildOk = $false
|
||||
$FailedStep = "git clone"
|
||||
if (Test-Path $buildTmp) { Remove-Item -Recurse -Force $buildTmp }
|
||||
if (Test-Path -LiteralPath $buildTmp) { Remove-Item -LiteralPath $buildTmp -Recurse -Force }
|
||||
}
|
||||
}
|
||||
# Use temp dir for build; swap into $LlamaCppDir only after build succeeds
|
||||
|
|
@ -2482,14 +2599,16 @@ if (-not $NeedLlamaSourceBuild) {
|
|||
|
||||
# Swap temp build dir into final location (only if we built in a temp dir)
|
||||
if ($BuildOk -and $LlamaCppDir -ne $OriginalLlamaCppDir) {
|
||||
if (Test-Path $OriginalLlamaCppDir) { Remove-Item -Recurse -Force $OriginalLlamaCppDir }
|
||||
Move-Item $LlamaCppDir $OriginalLlamaCppDir
|
||||
Assert-StudioOwnedOrAbsent -Path $OriginalLlamaCppDir -Label "llama.cpp install"
|
||||
if (Test-Path -LiteralPath $OriginalLlamaCppDir) { Remove-Item -LiteralPath $OriginalLlamaCppDir -Recurse -Force }
|
||||
Move-Item -LiteralPath $LlamaCppDir -Destination $OriginalLlamaCppDir
|
||||
$LlamaCppDir = $OriginalLlamaCppDir
|
||||
$BuildDir = Join-Path $LlamaCppDir "build"
|
||||
$LlamaServerBin = Join-Path $BuildDir "bin\Release\llama-server.exe"
|
||||
Mark-StudioOwned -Path $LlamaCppDir
|
||||
} elseif (-not $BuildOk -and $LlamaCppDir -ne $OriginalLlamaCppDir) {
|
||||
# Build failed -- clean up temp dir, preserve existing install
|
||||
if (Test-Path $LlamaCppDir) { Remove-Item -Recurse -Force $LlamaCppDir }
|
||||
if (Test-Path -LiteralPath $LlamaCppDir) { Remove-Item -LiteralPath $LlamaCppDir -Recurse -Force }
|
||||
$LlamaCppDir = $OriginalLlamaCppDir
|
||||
$BuildDir = Join-Path $LlamaCppDir "build"
|
||||
$LlamaServerBin = Join-Path $BuildDir "bin\Release\llama-server.exe"
|
||||
|
|
@ -2504,16 +2623,16 @@ if (-not $NeedLlamaSourceBuild) {
|
|||
$totalSec = [math]::Round($totalSw.Elapsed.TotalSeconds % 60, 1)
|
||||
|
||||
# -- Summary --
|
||||
if ($BuildOk -and (Test-Path $LlamaServerBin)) {
|
||||
if ($BuildOk -and (Test-Path -LiteralPath $LlamaServerBin)) {
|
||||
step "llama.cpp" "built"
|
||||
$QuantizeBin = Join-Path $BuildDir "bin\Release\llama-quantize.exe"
|
||||
if (Test-Path $QuantizeBin) {
|
||||
if (Test-Path -LiteralPath $QuantizeBin) {
|
||||
step "llama-quantize" "built"
|
||||
}
|
||||
step "build time" "${totalMin}m ${totalSec}s" "DarkGray"
|
||||
} else {
|
||||
$altBin = Join-Path $BuildDir "bin\llama-server.exe"
|
||||
if ($BuildOk -and (Test-Path $altBin)) {
|
||||
if ($BuildOk -and (Test-Path -LiteralPath $altBin)) {
|
||||
step "llama.cpp" "built"
|
||||
step "build time" "${totalMin}m ${totalSec}s" "DarkGray"
|
||||
} else {
|
||||
|
|
|
|||
103
studio/setup.sh
103
studio/setup.sh
|
|
@ -417,7 +417,36 @@ if [ -d "$SCRIPT_DIR/backend/core/data_recipe/oxc-validator" ] && command -v npm
|
|||
fi
|
||||
|
||||
# ── Python venv + deps ──
|
||||
STUDIO_HOME="$HOME/.unsloth/studio"
|
||||
# UNSLOTH_STUDIO_HOME (or STUDIO_HOME alias) overrides the install root
|
||||
# (mirrors install.sh). UNSLOTH_STUDIO_HOME wins when both are set.
|
||||
_studio_override_var=""
|
||||
_studio_override="${UNSLOTH_STUDIO_HOME:-}"
|
||||
if [ -n "$_studio_override" ]; then
|
||||
_studio_override_var="UNSLOTH_STUDIO_HOME"
|
||||
else
|
||||
_studio_override="${STUDIO_HOME:-}"
|
||||
[ -n "$_studio_override" ] && _studio_override_var="STUDIO_HOME"
|
||||
fi
|
||||
# Strip whitespace so " " is treated as unset (matches Python .strip()).
|
||||
_studio_override=$(printf '%s' "$_studio_override" | sed -e 's/^[[:space:]]*//' -e 's/[[:space:]]*$//')
|
||||
case "$_studio_override" in
|
||||
"~") _studio_override="$HOME" ;;
|
||||
"~/"*) _studio_override="$HOME/${_studio_override#'~/'}" ;;
|
||||
esac
|
||||
if [ -n "$_studio_override" ]; then
|
||||
# setup.sh runs against an existing install (via 'unsloth studio update');
|
||||
# a typo in the override must fail fast instead of materializing an
|
||||
# empty workspace dir. Mirrors setup.ps1 behavior.
|
||||
if [ ! -d "$_studio_override" ]; then
|
||||
echo "ERROR: $_studio_override_var=$_studio_override does not exist." >&2
|
||||
echo " Run install.sh to create the install root before 'unsloth studio update'." >&2
|
||||
exit 1
|
||||
fi
|
||||
[ -w "$_studio_override" ] || { echo "ERROR: $_studio_override_var=$_studio_override is not writable." >&2; exit 1; }
|
||||
STUDIO_HOME="$(CDPATH= cd -P -- "$_studio_override" && pwd -P)" || exit 1
|
||||
else
|
||||
STUDIO_HOME="$HOME/.unsloth/studio"
|
||||
fi
|
||||
VENV_DIR="$STUDIO_HOME/unsloth_studio"
|
||||
VENV_T5_530_DIR="$STUDIO_HOME/.venv_t5_530"
|
||||
VENV_T5_550_DIR="$STUDIO_HOME/.venv_t5_550"
|
||||
|
|
@ -542,9 +571,39 @@ fi
|
|||
#
|
||||
# Runs outside the _SKIP_PYTHON_DEPS gate so that upgrades from legacy
|
||||
# single .venv_t5 are always migrated to the tiered layout.
|
||||
# why: in env-override mode $STUDIO_HOME is user-chosen; require the
|
||||
# ownership marker before rm -rf so unrelated dirs survive. Gated on the
|
||||
# canonical comparison so an override pointing at the legacy default still
|
||||
# behaves like a default install.
|
||||
_STUDIO_OWNED_MARKER=".unsloth-studio-owned"
|
||||
_LEGACY_STUDIO_HOME="$HOME/.unsloth/studio"
|
||||
_studio_home_canon="$STUDIO_HOME"
|
||||
if [ -d "$_studio_home_canon" ]; then
|
||||
_studio_home_canon=$(CDPATH= cd -P -- "$_studio_home_canon" 2>/dev/null && pwd -P) \
|
||||
|| _studio_home_canon="$STUDIO_HOME"
|
||||
fi
|
||||
if [ -d "$_LEGACY_STUDIO_HOME" ]; then
|
||||
_LEGACY_STUDIO_HOME=$(CDPATH= cd -P -- "$_LEGACY_STUDIO_HOME" 2>/dev/null && pwd -P) \
|
||||
|| _LEGACY_STUDIO_HOME="$HOME/.unsloth/studio"
|
||||
fi
|
||||
_STUDIO_HOME_IS_CUSTOM=false
|
||||
if [ "$_studio_home_canon" != "$_LEGACY_STUDIO_HOME" ]; then
|
||||
_STUDIO_HOME_IS_CUSTOM=true
|
||||
fi
|
||||
_assert_studio_owned_or_absent() {
|
||||
_aso_dir="$1"
|
||||
_aso_label="$2"
|
||||
[ -d "$_aso_dir" ] || return 0
|
||||
if [ "$_STUDIO_HOME_IS_CUSTOM" = true ] && [ ! -f "$_aso_dir/$_STUDIO_OWNED_MARKER" ]; then
|
||||
echo "ERROR: $_aso_dir already exists and is not marked as a Studio-owned $_aso_label." >&2
|
||||
echo " Move it aside or choose an empty UNSLOTH_STUDIO_HOME before re-running." >&2
|
||||
exit 1
|
||||
fi
|
||||
}
|
||||
_NEED_T5_INSTALL=false
|
||||
if [ -d "$STUDIO_HOME/.venv_t5" ]; then
|
||||
# Legacy layout — migrate
|
||||
_assert_studio_owned_or_absent "$STUDIO_HOME/.venv_t5" "legacy transformers sidecar venv"
|
||||
rm -rf "$STUDIO_HOME/.venv_t5"
|
||||
_NEED_T5_INSTALL=true
|
||||
fi
|
||||
|
|
@ -554,16 +613,20 @@ fi
|
|||
[ "$_SKIP_PYTHON_DEPS" = false ] && _NEED_T5_INSTALL=true
|
||||
|
||||
if [ "$_NEED_T5_INSTALL" = true ]; then
|
||||
_assert_studio_owned_or_absent "$VENV_T5_530_DIR" "transformers 5.3 sidecar venv"
|
||||
[ -d "$VENV_T5_530_DIR" ] && rm -rf "$VENV_T5_530_DIR"
|
||||
mkdir -p "$VENV_T5_530_DIR"
|
||||
: > "$VENV_T5_530_DIR/$_STUDIO_OWNED_MARKER" 2>/dev/null || true
|
||||
run_quiet "install transformers 5.3.0" fast_install --target "$VENV_T5_530_DIR" --no-deps "transformers==5.3.0"
|
||||
run_quiet "install huggingface_hub for t5_530" fast_install --target "$VENV_T5_530_DIR" --no-deps "huggingface_hub==1.8.0"
|
||||
run_quiet "install hf_xet for t5_530" fast_install --target "$VENV_T5_530_DIR" --no-deps "hf_xet==1.4.2"
|
||||
run_quiet "install tiktoken for t5_530" fast_install --target "$VENV_T5_530_DIR" "tiktoken"
|
||||
step "transformers" "5.3.0 pre-installed"
|
||||
|
||||
_assert_studio_owned_or_absent "$VENV_T5_550_DIR" "transformers 5.5 sidecar venv"
|
||||
[ -d "$VENV_T5_550_DIR" ] && rm -rf "$VENV_T5_550_DIR"
|
||||
mkdir -p "$VENV_T5_550_DIR"
|
||||
: > "$VENV_T5_550_DIR/$_STUDIO_OWNED_MARKER" 2>/dev/null || true
|
||||
run_quiet "install transformers 5.5.0" fast_install --target "$VENV_T5_550_DIR" --no-deps "transformers==5.5.0"
|
||||
run_quiet "install huggingface_hub for t5_550" fast_install --target "$VENV_T5_550_DIR" --no-deps "huggingface_hub==1.8.0"
|
||||
run_quiet "install hf_xet for t5_550" fast_install --target "$VENV_T5_550_DIR" --no-deps "hf_xet==1.4.2"
|
||||
|
|
@ -573,7 +636,13 @@ fi
|
|||
fi
|
||||
|
||||
# ── 7. Prefer prebuilt llama.cpp bundles before any source build path ──
|
||||
UNSLOTH_HOME="$HOME/.unsloth"
|
||||
# Nest llama.cpp under $STUDIO_HOME only for real env-overrides; legacy
|
||||
# default keeps ~/.unsloth/llama.cpp so pre-PR builds are still discovered.
|
||||
if [ "$_STUDIO_HOME_IS_CUSTOM" = true ]; then
|
||||
UNSLOTH_HOME="$STUDIO_HOME"
|
||||
else
|
||||
UNSLOTH_HOME="$HOME/.unsloth"
|
||||
fi
|
||||
mkdir -p "$UNSLOTH_HOME"
|
||||
LLAMA_CPP_DIR="$UNSLOTH_HOME/llama.cpp"
|
||||
LLAMA_SERVER_BIN="$LLAMA_CPP_DIR/build/bin/llama-server"
|
||||
|
|
@ -582,11 +651,30 @@ _LLAMA_CPP_DEGRADED=false
|
|||
_LLAMA_FORCE_COMPILE="${UNSLOTH_LLAMA_FORCE_COMPILE:-0}"
|
||||
_REQUESTED_LLAMA_TAG="${UNSLOTH_LLAMA_TAG:-${_DEFAULT_LLAMA_TAG}}"
|
||||
_HOST_SYSTEM="$(uname -s 2>/dev/null || true)"
|
||||
_HOST_MACHINE="$(uname -m 2>/dev/null || true)"
|
||||
|
||||
# Pick the release repo install_llama_prebuilt.py plans against.
|
||||
# unslothai/llama.cpp ships only Linux CUDA bundles, so CPU-only Linux
|
||||
# x86_64 routes to ggml-org for bin-ubuntu-x64.tar.gz. Anything with a
|
||||
# GPU tool installed stays on unslothai (CUDA bundle / ROCm source build).
|
||||
_LINUX_HAS_GPU=false
|
||||
for _GPU_TOOL in nvidia-smi rocminfo amd-smi hipconfig hipinfo; do
|
||||
if command -v "$_GPU_TOOL" >/dev/null 2>&1; then
|
||||
_LINUX_HAS_GPU=true
|
||||
break
|
||||
fi
|
||||
done
|
||||
|
||||
if [ "$_HOST_SYSTEM" = "Darwin" ]; then
|
||||
_HELPER_RELEASE_REPO="ggml-org/llama.cpp"
|
||||
elif [ "$_HOST_SYSTEM" = "Linux" ] \
|
||||
&& [ "$_HOST_MACHINE" = "x86_64" ] \
|
||||
&& [ "$_LINUX_HAS_GPU" = false ]; then
|
||||
_HELPER_RELEASE_REPO="ggml-org/llama.cpp"
|
||||
else
|
||||
_HELPER_RELEASE_REPO="unslothai/llama.cpp"
|
||||
fi
|
||||
unset _GPU_TOOL
|
||||
_LLAMA_PR="${UNSLOTH_LLAMA_PR:-}"
|
||||
_SKIP_PREBUILT_INSTALL=false
|
||||
_LLAMA_PR_FORCE="${UNSLOTH_LLAMA_PR_FORCE:-${_DEFAULT_LLAMA_PR_FORCE}}"
|
||||
|
|
@ -635,6 +723,12 @@ else
|
|||
if [ -d "$LLAMA_CPP_DIR" ]; then
|
||||
substep "existing install detected -- validating update"
|
||||
fi
|
||||
# why: install_llama_prebuilt.py uses os.replace(), which would displace
|
||||
# an unrelated $UNSLOTH_STUDIO_HOME/llama.cpp before the source-build
|
||||
# ownership check below ever runs.
|
||||
if [ "$_STUDIO_HOME_IS_CUSTOM" = true ]; then
|
||||
_assert_studio_owned_or_absent "$LLAMA_CPP_DIR" "llama.cpp install"
|
||||
fi
|
||||
_PREBUILT_CMD=(
|
||||
python "$SCRIPT_DIR/install_llama_prebuilt.py"
|
||||
--install-dir "$LLAMA_CPP_DIR"
|
||||
|
|
@ -662,6 +756,9 @@ else
|
|||
else
|
||||
step "llama.cpp" "prebuilt installed and validated"
|
||||
fi
|
||||
if [ "$_STUDIO_HOME_IS_CUSTOM" = true ] && [ -d "$LLAMA_CPP_DIR" ]; then
|
||||
: > "$LLAMA_CPP_DIR/$_STUDIO_OWNED_MARKER" 2>/dev/null || true
|
||||
fi
|
||||
print_installed_llama_prebuilt_release "$LLAMA_CPP_DIR"
|
||||
verbose_substep "llama.cpp install dir: $LLAMA_CPP_DIR"
|
||||
rm -f "$_PREBUILT_LOG"
|
||||
|
|
@ -1032,8 +1129,10 @@ else
|
|||
|
||||
# Swap only after build succeeds -- preserves existing install on failure
|
||||
if [ "$BUILD_OK" = true ]; then
|
||||
_assert_studio_owned_or_absent "$LLAMA_CPP_DIR" "llama.cpp install"
|
||||
rm -rf "$LLAMA_CPP_DIR"
|
||||
mv "$_BUILD_TMP" "$LLAMA_CPP_DIR"
|
||||
: > "$LLAMA_CPP_DIR/$_STUDIO_OWNED_MARKER" 2>/dev/null || true
|
||||
# Symlink to llama.cpp root -- check_llama_cpp() looks for the binary there
|
||||
QUANTIZE_BIN="$LLAMA_CPP_DIR/build/bin/llama-quantize"
|
||||
if [ -f "$QUANTIZE_BIN" ]; then
|
||||
|
|
|
|||
|
|
@ -60,6 +60,11 @@ pub async fn check_install_status() -> bool {
|
|||
cmd.env_remove("PYTHONPATH");
|
||||
}
|
||||
|
||||
// Tauri uses the legacy root regardless of UNSLOTH_STUDIO_HOME / STUDIO_HOME;
|
||||
// probe subprocesses must follow the same isolation as process.rs.
|
||||
cmd.env_remove("UNSLOTH_STUDIO_HOME");
|
||||
cmd.env_remove("STUDIO_HOME");
|
||||
|
||||
let mut child = match cmd.spawn() {
|
||||
Ok(c) => c,
|
||||
Err(e) => {
|
||||
|
|
|
|||
|
|
@ -203,6 +203,11 @@ async fn provision_desktop_auth() -> Result<(), String> {
|
|||
cmd.env_remove("PYTHONHOME");
|
||||
cmd.env_remove("PYTHONPATH");
|
||||
}
|
||||
|
||||
// Tauri uses the legacy root regardless of UNSLOTH_STUDIO_HOME / STUDIO_HOME.
|
||||
// Scrub so provisioning writes match what the Rust auth code reads.
|
||||
cmd.env_remove("UNSLOTH_STUDIO_HOME");
|
||||
cmd.env_remove("STUDIO_HOME");
|
||||
#[cfg(windows)]
|
||||
{
|
||||
use std::os::windows::process::CommandExt;
|
||||
|
|
|
|||
|
|
@ -196,6 +196,11 @@ fn spawn_script(
|
|||
cmd.env_remove("PYTHONPATH");
|
||||
}
|
||||
|
||||
// Tauri only does default-root installs; install.sh / install.ps1 reject
|
||||
// these under --tauri. Scrub so an inherited value can't trip the guard.
|
||||
cmd.env_remove("UNSLOTH_STUDIO_HOME");
|
||||
cmd.env_remove("STUDIO_HOME");
|
||||
|
||||
// On Windows, launch the installer directly with CREATE_NO_WINDOW.
|
||||
// The app process is assigned to a KILL_ON_JOB_CLOSE job in main.rs, so
|
||||
// child cleanup on crash comes from inherited job membership instead.
|
||||
|
|
|
|||
|
|
@ -102,6 +102,11 @@ async fn run_cli_probe(bin: &std::path::Path, args: &[&str]) -> bool {
|
|||
cmd.env_remove("PYTHONPATH");
|
||||
}
|
||||
|
||||
// Tauri uses the legacy root regardless of UNSLOTH_STUDIO_HOME / STUDIO_HOME;
|
||||
// probe subprocesses must follow the same isolation as process.rs.
|
||||
cmd.env_remove("UNSLOTH_STUDIO_HOME");
|
||||
cmd.env_remove("STUDIO_HOME");
|
||||
|
||||
#[cfg(windows)]
|
||||
{
|
||||
use std::os::windows::process::CommandExt;
|
||||
|
|
@ -135,6 +140,11 @@ async fn probe_cli_capability(bin: &std::path::Path) -> Option<DesktopCapability
|
|||
cmd.env_remove("PYTHONPATH");
|
||||
}
|
||||
|
||||
// Tauri uses the legacy root regardless of UNSLOTH_STUDIO_HOME / STUDIO_HOME;
|
||||
// probe subprocesses must follow the same isolation as process.rs.
|
||||
cmd.env_remove("UNSLOTH_STUDIO_HOME");
|
||||
cmd.env_remove("STUDIO_HOME");
|
||||
|
||||
#[cfg(windows)]
|
||||
{
|
||||
use std::os::windows::process::CommandExt;
|
||||
|
|
|
|||
|
|
@ -316,6 +316,12 @@ pub fn start_backend(
|
|||
cmd.env_remove("PYTHONPATH");
|
||||
}
|
||||
|
||||
// Tauri uses the legacy root regardless of UNSLOTH_STUDIO_HOME / STUDIO_HOME;
|
||||
// scrub so the spawned Python backend can't diverge. UNSLOTH_LLAMA_CPP_PATH
|
||||
// is a pre-existing user-controlled llama.cpp dir override; keep it.
|
||||
cmd.env_remove("UNSLOTH_STUDIO_HOME");
|
||||
cmd.env_remove("STUDIO_HOME");
|
||||
|
||||
// On Windows, launch the backend directly with hidden-window flags.
|
||||
// The app process is assigned to a KILL_ON_JOB_CLOSE job in main.rs, so
|
||||
// children inherit crash-safe cleanup without the buggy per-child JobObject wrapper.
|
||||
|
|
|
|||
|
|
@ -61,6 +61,11 @@ fn spawn_update(
|
|||
cmd.env_remove("PYTHONPATH");
|
||||
}
|
||||
|
||||
// Tauri manages the legacy root; scrub so 'unsloth studio update' targets
|
||||
// the same install the desktop app uses, not an inherited custom root.
|
||||
cmd.env_remove("UNSLOTH_STUDIO_HOME");
|
||||
cmd.env_remove("STUDIO_HOME");
|
||||
|
||||
#[cfg(windows)]
|
||||
let mut child: Box<dyn ChildWrapper + Send> = {
|
||||
use std::os::windows::process::CommandExt;
|
||||
|
|
|
|||
46
tests/python/test_gpu_init_ldconfig_guard.py
Normal file
46
tests/python/test_gpu_init_ldconfig_guard.py
Normal file
|
|
@ -0,0 +1,46 @@
|
|||
import ast
|
||||
from pathlib import Path
|
||||
|
||||
REPO_ROOT = Path(__file__).resolve().parents[2]
|
||||
GPU_INIT = REPO_ROOT / "unsloth" / "_gpu_init.py"
|
||||
|
||||
|
||||
def _find_geteuid_guard(tree: ast.AST):
|
||||
for node in ast.walk(tree):
|
||||
if not isinstance(node, ast.If):
|
||||
continue
|
||||
for sub in ast.walk(node.test):
|
||||
if isinstance(sub, ast.Call) and isinstance(sub.func, ast.Attribute):
|
||||
if sub.func.attr == "geteuid":
|
||||
return node
|
||||
return None
|
||||
|
||||
|
||||
def test_gpu_init_has_geteuid_guard():
|
||||
tree = ast.parse(GPU_INIT.read_text())
|
||||
guard = _find_geteuid_guard(tree)
|
||||
assert (
|
||||
guard is not None
|
||||
), "_gpu_init.py must guard ldconfig recovery on os.geteuid()"
|
||||
|
||||
|
||||
def test_ldconfig_calls_only_inside_geteuid_guard():
|
||||
src = GPU_INIT.read_text()
|
||||
tree = ast.parse(src)
|
||||
guard = _find_geteuid_guard(tree)
|
||||
assert guard is not None
|
||||
guard_src = ast.get_source_segment(src, guard) or ""
|
||||
ldconfig_lines = [
|
||||
line for line in src.splitlines() if "ldconfig" in line and "os.system" in line
|
||||
]
|
||||
for line in ldconfig_lines:
|
||||
assert line.strip() in guard_src, (
|
||||
"os.system('ldconfig ...') must live inside the geteuid guard, "
|
||||
f"but found unguarded: {line!r}"
|
||||
)
|
||||
|
||||
|
||||
def test_non_root_branch_warns_when_bnb_present():
|
||||
src = GPU_INIT.read_text()
|
||||
assert "elif bnb is not None" in src
|
||||
assert "sudo ldconfig" in src
|
||||
122
tests/studio/test_export_output_path_contract.py
Normal file
122
tests/studio/test_export_output_path_contract.py
Normal file
|
|
@ -0,0 +1,122 @@
|
|||
# SPDX-License-Identifier: AGPL-3.0-only
|
||||
|
||||
import ast
|
||||
from pathlib import Path
|
||||
|
||||
REPO_ROOT = Path(__file__).resolve().parents[2]
|
||||
EXPORT = REPO_ROOT / "studio" / "backend" / "core" / "export" / "export.py"
|
||||
|
||||
EXPORT_FNS = (
|
||||
"export_merged_model",
|
||||
"export_base_model",
|
||||
"export_gguf",
|
||||
"export_lora_adapter",
|
||||
)
|
||||
|
||||
|
||||
def _find_method(tree, cls_name, method_name):
|
||||
for cls in ast.walk(tree):
|
||||
if isinstance(cls, ast.ClassDef) and cls.name == cls_name:
|
||||
for item in cls.body:
|
||||
if isinstance(item, ast.FunctionDef) and item.name == method_name:
|
||||
return item
|
||||
return None
|
||||
|
||||
|
||||
def _return_tuple_arity(fn):
|
||||
arities = []
|
||||
for node in ast.walk(fn):
|
||||
if isinstance(node, ast.Return) and isinstance(node.value, ast.Tuple):
|
||||
arities.append(len(node.value.elts))
|
||||
return arities
|
||||
|
||||
|
||||
def test_export_methods_return_three_tuple_annotation():
|
||||
tree = ast.parse(EXPORT.read_text())
|
||||
for fn_name in EXPORT_FNS:
|
||||
fn = _find_method(tree, "ExportBackend", fn_name)
|
||||
assert fn is not None, f"missing ExportBackend.{fn_name}"
|
||||
ret = fn.returns
|
||||
assert isinstance(ret, ast.Subscript), f"{fn_name} return must be Tuple[...]"
|
||||
slc = ret.slice
|
||||
elts = slc.elts if isinstance(slc, ast.Tuple) else None
|
||||
assert (
|
||||
elts is not None and len(elts) == 3
|
||||
), f"{fn_name} return annotation must be a 3-tuple, got {ast.dump(ret)}"
|
||||
|
||||
|
||||
def test_export_methods_return_three_element_tuples():
|
||||
tree = ast.parse(EXPORT.read_text())
|
||||
for fn_name in EXPORT_FNS:
|
||||
fn = _find_method(tree, "ExportBackend", fn_name)
|
||||
assert fn is not None
|
||||
arities = _return_tuple_arity(fn)
|
||||
assert arities, f"{fn_name} has no tuple-return statements"
|
||||
for arity in arities:
|
||||
assert arity == 3, f"{fn_name} return tuple arity {arity}, expected 3"
|
||||
|
||||
|
||||
def test_local_save_assigns_output_path():
|
||||
tree = ast.parse(EXPORT.read_text())
|
||||
for fn_name in EXPORT_FNS:
|
||||
fn = _find_method(tree, "ExportBackend", fn_name)
|
||||
assert fn is not None
|
||||
assigns = []
|
||||
for node in ast.walk(fn):
|
||||
if isinstance(node, ast.Assign):
|
||||
for tgt in node.targets:
|
||||
if isinstance(tgt, ast.Name) and tgt.id == "output_path":
|
||||
assigns.append(node)
|
||||
non_none = [
|
||||
a
|
||||
for a in assigns
|
||||
if not (isinstance(a.value, ast.Constant) and a.value.value is None)
|
||||
]
|
||||
assert non_none, f"{fn_name} never assigns a non-None output_path"
|
||||
|
||||
|
||||
def test_gpu_save_method_bound_for_hub_only():
|
||||
tree = ast.parse(EXPORT.read_text())
|
||||
fn = _find_method(tree, "ExportBackend", "export_merged_model")
|
||||
assert fn is not None
|
||||
found_pre_save_method = False
|
||||
for node in ast.walk(fn):
|
||||
if isinstance(node, ast.Try):
|
||||
for stmt in node.body:
|
||||
if isinstance(stmt, ast.If):
|
||||
test = stmt.test
|
||||
if isinstance(test, ast.Name) and test.id == "_IS_MLX":
|
||||
for sub in ast.walk(
|
||||
ast.Module(body = stmt.orelse, type_ignores = [])
|
||||
):
|
||||
if isinstance(sub, ast.Assign) and any(
|
||||
isinstance(t, ast.Name) and t.id == "save_method"
|
||||
for t in sub.targets
|
||||
):
|
||||
found_pre_save_method = True
|
||||
break
|
||||
if found_pre_save_method:
|
||||
break
|
||||
if found_pre_save_method:
|
||||
break
|
||||
assert found_pre_save_method, (
|
||||
"GPU save_method must be assigned at the top of the try block, "
|
||||
"before the `if save_directory:` guard, so Hub-only export does not "
|
||||
"raise UnboundLocalError."
|
||||
)
|
||||
|
||||
|
||||
def test_mlx_hub_only_uses_temp_directory():
|
||||
src = EXPORT.read_text()
|
||||
assert (
|
||||
src.count("tempfile.TemporaryDirectory") >= 3
|
||||
), "expected TemporaryDirectory in merged, base, and lora hub-push paths"
|
||||
assert "import tempfile" in src.split("class ExportBackend")[0]
|
||||
|
||||
|
||||
def test_is_mlx_imported_from_unsloth():
|
||||
src = EXPORT.read_text()
|
||||
assert "from unsloth import" in src
|
||||
head = src.split("class ExportBackend")[0]
|
||||
assert "_IS_MLX" in head
|
||||
assert "_IS_MLX = platform.system()" not in src
|
||||
213
tests/studio/test_is_mlx_dispatch_gate.py
Normal file
213
tests/studio/test_is_mlx_dispatch_gate.py
Normal file
|
|
@ -0,0 +1,213 @@
|
|||
# SPDX-License-Identifier: AGPL-3.0-only
|
||||
|
||||
"""
|
||||
Regression tests for the CUDA-vs-MLX dispatch gates Studio relies on.
|
||||
|
||||
Two gates drive every dispatch decision in Studio's MLX path:
|
||||
|
||||
1. ``unsloth._IS_MLX`` at the top of ``unsloth/__init__.py`` -- evaluated
|
||||
once at import time and read by Studio worker code to choose between
|
||||
the GPU and MLX trainer / inference / export paths. Defined as
|
||||
``Darwin AND arm64 AND find_spec("mlx") is not None``.
|
||||
|
||||
2. ``utils.hardware.detect_hardware()`` -- runtime probe in the Studio
|
||||
backend. Priority order: CUDA -> XPU -> MLX -> CPU. The MLX branch is
|
||||
reached only when both CUDA and XPU are unavailable AND the host is
|
||||
Apple Silicon AND mlx is importable.
|
||||
|
||||
These gates are the canaries for "MLX support accidentally hijacks
|
||||
CUDA/AMD/Intel users". The tests here:
|
||||
|
||||
* verify the source-level structure of the ``_IS_MLX`` expression so an
|
||||
accidental rewrite (e.g. dropping the ``arm64`` check) is caught,
|
||||
* exercise the runtime gate logic under a spoofed Darwin+arm64 platform
|
||||
with a fake ``mlx`` module in ``sys.modules`` to confirm both gates
|
||||
flip True together,
|
||||
* confirm that on the actual Linux+CUDA test host both gates remain in
|
||||
their CUDA-side state.
|
||||
|
||||
No real MLX install is required; uses the same ``monkeypatch.setitem``
|
||||
fake-mlx pattern as ``test_mlx_inference_backend.py``.
|
||||
"""
|
||||
|
||||
import ast
|
||||
import importlib
|
||||
import sys
|
||||
import types
|
||||
from pathlib import Path
|
||||
|
||||
|
||||
REPO_ROOT = Path(__file__).resolve().parents[2]
|
||||
UNSLOTH_INIT = REPO_ROOT / "unsloth" / "__init__.py"
|
||||
|
||||
|
||||
# ---------------------------------------------------------------------------
|
||||
# 1. Source-level structure check on _IS_MLX (no platform dependencies).
|
||||
# ---------------------------------------------------------------------------
|
||||
|
||||
|
||||
def test_is_mlx_gate_uses_three_required_predicates():
|
||||
"""The _IS_MLX assignment must AND together exactly the three checks
|
||||
that Studio depends on: Darwin OS, arm64 machine, and an importable
|
||||
mlx package. Dropping any one of them silently breaks dispatch.
|
||||
"""
|
||||
tree = ast.parse(UNSLOTH_INIT.read_text())
|
||||
|
||||
target = None
|
||||
for node in ast.walk(tree):
|
||||
if (
|
||||
isinstance(node, ast.Assign)
|
||||
and len(node.targets) == 1
|
||||
and isinstance(node.targets[0], ast.Name)
|
||||
and node.targets[0].id == "_IS_MLX"
|
||||
):
|
||||
target = node.value
|
||||
break
|
||||
assert target is not None, "_IS_MLX assignment not found in unsloth/__init__.py"
|
||||
assert isinstance(target, ast.BoolOp) and isinstance(
|
||||
target.op, ast.And
|
||||
), "_IS_MLX must be a BoolOp(And) of platform + mlx checks"
|
||||
|
||||
expr_src = ast.unparse(target)
|
||||
assert (
|
||||
"platform.system()" in expr_src and "Darwin" in expr_src
|
||||
), "_IS_MLX must check platform.system() == 'Darwin'"
|
||||
assert (
|
||||
"platform.machine()" in expr_src and "arm64" in expr_src
|
||||
), "_IS_MLX must check platform.machine() == 'arm64'"
|
||||
assert (
|
||||
"find_spec" in expr_src and "'mlx'" in expr_src
|
||||
), "_IS_MLX must check importlib.util.find_spec('mlx')"
|
||||
|
||||
|
||||
# ---------------------------------------------------------------------------
|
||||
# 2. Runtime gate behavior with the platform spoofed to Apple Silicon and a
|
||||
# fake mlx module in sys.modules. Re-evaluates the same expression
|
||||
# rather than reloading unsloth (which would cascade-reload torch).
|
||||
# ---------------------------------------------------------------------------
|
||||
|
||||
|
||||
def _evaluate_is_mlx_gate(platform_module, importlib_util):
|
||||
"""Re-evaluate the _IS_MLX expression using injected dependencies.
|
||||
|
||||
Mirrors the assignment in unsloth/__init__.py exactly.
|
||||
"""
|
||||
return (
|
||||
platform_module.system() == "Darwin"
|
||||
and platform_module.machine() == "arm64"
|
||||
and importlib_util.find_spec("mlx") is not None
|
||||
)
|
||||
|
||||
|
||||
def test_is_mlx_gate_true_on_apple_silicon_with_mlx_present(monkeypatch):
|
||||
import platform
|
||||
import importlib.util
|
||||
|
||||
# Inject a fake mlx package so find_spec returns a non-None ModuleSpec.
|
||||
fake_mlx = types.ModuleType("mlx")
|
||||
fake_mlx.__spec__ = importlib.machinery.ModuleSpec("mlx", loader = None)
|
||||
fake_mlx.__path__ = []
|
||||
monkeypatch.setitem(sys.modules, "mlx", fake_mlx)
|
||||
|
||||
monkeypatch.setattr(platform, "system", lambda: "Darwin")
|
||||
monkeypatch.setattr(platform, "machine", lambda: "arm64")
|
||||
|
||||
assert _evaluate_is_mlx_gate(platform, importlib.util) is True
|
||||
|
||||
|
||||
def test_is_mlx_gate_false_when_mlx_missing(monkeypatch):
|
||||
import platform
|
||||
import importlib.util
|
||||
|
||||
# Apple Silicon platform but no mlx package -> gate must be False.
|
||||
monkeypatch.delitem(sys.modules, "mlx", raising = False)
|
||||
monkeypatch.setattr(platform, "system", lambda: "Darwin")
|
||||
monkeypatch.setattr(platform, "machine", lambda: "arm64")
|
||||
|
||||
real_find_spec = importlib.util.find_spec
|
||||
|
||||
def _no_mlx(name, *args, **kwargs):
|
||||
if name == "mlx":
|
||||
return None
|
||||
return real_find_spec(name, *args, **kwargs)
|
||||
|
||||
monkeypatch.setattr(importlib.util, "find_spec", _no_mlx)
|
||||
|
||||
assert _evaluate_is_mlx_gate(platform, importlib.util) is False
|
||||
|
||||
|
||||
def test_is_mlx_gate_false_on_non_apple_silicon():
|
||||
"""On the real Linux+CUDA / AMD / Intel test host, the gate stays False."""
|
||||
import platform
|
||||
import importlib.util
|
||||
|
||||
if platform.system() == "Darwin" and platform.machine() == "arm64":
|
||||
# On a Mac CI runner this assertion would not apply; skip there.
|
||||
import pytest
|
||||
|
||||
pytest.skip("Test host is Apple Silicon; CUDA-side canary doesn't apply.")
|
||||
|
||||
assert _evaluate_is_mlx_gate(platform, importlib.util) is False
|
||||
|
||||
|
||||
# ---------------------------------------------------------------------------
|
||||
# 3. Studio's runtime detect_hardware() picks MLX only when CUDA + XPU are
|
||||
# both unavailable AND the host is Apple Silicon AND mlx is importable.
|
||||
# ---------------------------------------------------------------------------
|
||||
|
||||
|
||||
def _import_studio_hardware():
|
||||
"""Lazy import for the Studio hardware module, with the bare-imports
|
||||
convention that Studio uses (studio/backend on sys.path).
|
||||
"""
|
||||
studio_backend = REPO_ROOT / "studio" / "backend"
|
||||
if str(studio_backend) not in sys.path:
|
||||
sys.path.insert(0, str(studio_backend))
|
||||
from utils.hardware import hardware as hw # type: ignore
|
||||
|
||||
return hw
|
||||
|
||||
|
||||
def test_detect_hardware_picks_mlx_when_only_apple_silicon_available(monkeypatch):
|
||||
hw = _import_studio_hardware()
|
||||
|
||||
# Force CUDA + XPU paths off so detect_hardware falls through to MLX.
|
||||
import torch
|
||||
|
||||
monkeypatch.setattr(torch.cuda, "is_available", lambda: False)
|
||||
if hasattr(torch, "xpu"):
|
||||
monkeypatch.setattr(torch.xpu, "is_available", lambda: False)
|
||||
|
||||
# Spoof Apple Silicon and provide an importable mlx.core for _has_mlx().
|
||||
import platform
|
||||
|
||||
monkeypatch.setattr(platform, "system", lambda: "Darwin")
|
||||
monkeypatch.setattr(platform, "machine", lambda: "arm64")
|
||||
|
||||
fake_mlx = types.ModuleType("mlx")
|
||||
fake_mlx_core = types.ModuleType("mlx.core")
|
||||
fake_mlx.core = fake_mlx_core
|
||||
monkeypatch.setitem(sys.modules, "mlx", fake_mlx)
|
||||
monkeypatch.setitem(sys.modules, "mlx.core", fake_mlx_core)
|
||||
|
||||
detected = hw.detect_hardware()
|
||||
assert detected == hw.DeviceType.MLX, f"expected MLX, got {detected!r}"
|
||||
|
||||
|
||||
def test_detect_hardware_picks_cuda_on_real_host():
|
||||
"""Canary: on a real CUDA host the MLX branch must NOT be taken even
|
||||
if mlx happens to be importable. Protects CUDA/AMD/Intel users from
|
||||
accidental MLX dispatch when MLX support is added.
|
||||
"""
|
||||
import torch
|
||||
|
||||
if not torch.cuda.is_available():
|
||||
import pytest
|
||||
|
||||
pytest.skip("No CUDA available on this host; canary not applicable.")
|
||||
|
||||
hw = _import_studio_hardware()
|
||||
detected = hw.detect_hardware()
|
||||
assert (
|
||||
detected == hw.DeviceType.CUDA
|
||||
), f"CUDA host must dispatch to CUDA, got {detected!r}"
|
||||
90
tests/studio/test_mlx_training_worker_behaviors.py
Normal file
90
tests/studio/test_mlx_training_worker_behaviors.py
Normal file
|
|
@ -0,0 +1,90 @@
|
|||
# SPDX-License-Identifier: AGPL-3.0-only
|
||||
|
||||
import ast
|
||||
from pathlib import Path
|
||||
|
||||
REPO_ROOT = Path(__file__).resolve().parents[2]
|
||||
WORKER = REPO_ROOT / "studio" / "backend" / "core" / "training" / "worker.py"
|
||||
|
||||
|
||||
def _find_func(tree, name):
|
||||
for node in ast.walk(tree):
|
||||
if isinstance(node, ast.FunctionDef) and node.name == name:
|
||||
return node
|
||||
return None
|
||||
|
||||
|
||||
def test_run_mlx_training_passes_token_to_from_pretrained():
|
||||
tree = ast.parse(WORKER.read_text())
|
||||
fn = _find_func(tree, "_run_mlx_training")
|
||||
assert fn is not None
|
||||
found = False
|
||||
for node in ast.walk(fn):
|
||||
if (
|
||||
isinstance(node, ast.Call)
|
||||
and isinstance(node.func, ast.Attribute)
|
||||
and node.func.attr == "from_pretrained"
|
||||
and isinstance(node.func.value, ast.Name)
|
||||
and node.func.value.id == "FastMLXModel"
|
||||
):
|
||||
kwarg_names = {kw.arg for kw in node.keywords if kw.arg}
|
||||
assert (
|
||||
"token" in kwarg_names
|
||||
), f"FastMLXModel.from_pretrained must forward token=hf_token; got {kwarg_names!r}"
|
||||
found = True
|
||||
assert found, "FastMLXModel.from_pretrained call not found in _run_mlx_training"
|
||||
|
||||
|
||||
def test_wandb_init_strips_secret_keys():
|
||||
src = WORKER.read_text()
|
||||
assert "_wandb_sensitive" in src, "expected a sensitive-key set near wandb.init"
|
||||
assert '"hf_token"' in src and '"wandb_token"' in src
|
||||
assert (
|
||||
"config = dict(config)" not in src
|
||||
), "wandb.init received raw config dict; secrets would leak"
|
||||
|
||||
|
||||
def test_local_dataset_loader_uses_load_dataset_path():
|
||||
src = WORKER.read_text()
|
||||
assert "_resolve_local_files" in src
|
||||
assert "_loader_for_files" in src
|
||||
assert "data_files = all_files" in src or "data_files=all_files" in src
|
||||
|
||||
|
||||
def test_send_aliases_status_message_to_message():
|
||||
src = WORKER.read_text()
|
||||
assert 'kwargs["message"] = sm' in src or 'kwargs["message"]=sm' in src
|
||||
|
||||
|
||||
def test_slice_uses_inclusive_end_and_handles_zero():
|
||||
src = WORKER.read_text()
|
||||
assert "min(end + 1, len(ds))" in src or "min(end+1, len(ds))" in src
|
||||
assert "slice_start if slice_start is not None else 0" in src
|
||||
assert "slice_end if slice_end is not None else len(ds) - 1" in src
|
||||
|
||||
|
||||
def test_poll_stop_returns_on_broken_pipe():
|
||||
src = WORKER.read_text()
|
||||
assert "except (EOFError, OSError)" in src
|
||||
lines = src.splitlines()
|
||||
for i, line in enumerate(lines):
|
||||
if "except (EOFError, OSError)" in line:
|
||||
for j in range(i + 1, min(i + 6, len(lines))):
|
||||
stripped = lines[j].strip()
|
||||
if not stripped or stripped.startswith("#"):
|
||||
continue
|
||||
assert stripped.startswith(
|
||||
"return"
|
||||
), f"expected return after EOFError/OSError, got {stripped!r}"
|
||||
break
|
||||
break
|
||||
else:
|
||||
raise AssertionError("EOFError/OSError handler not found in worker.py")
|
||||
|
||||
|
||||
def test_unsloth_zoo_mlx_imports_have_friendly_error():
|
||||
src = WORKER.read_text()
|
||||
assert "from unsloth_zoo.mlx_loader import FastMLXModel" in src
|
||||
assert "from unsloth_zoo.mlx_trainer import" in src
|
||||
assert "raise ImportError" in src
|
||||
assert "install.sh" in src
|
||||
1021
tests/test_studio_install_workspace_guard.py
Normal file
1021
tests/test_studio_install_workspace_guard.py
Normal file
File diff suppressed because it is too large
Load diff
154
tests/test_studio_root_resilience.py
Normal file
154
tests/test_studio_root_resilience.py
Normal file
|
|
@ -0,0 +1,154 @@
|
|||
"""Resilience checks for Studio install-root inference under hostile
|
||||
filesystem conditions:
|
||||
- _infer_studio_home_from_venv must NOT propagate PermissionError /
|
||||
OSError out through studio_root() (it would crash module import in
|
||||
run.py / main.py / transformers_version.py / model_config.py).
|
||||
- _kill_orphaned_servers must catch (ImportError, OSError, ValueError)
|
||||
on the studio_root() probe so a transient resolve / sentinel failure
|
||||
cannot crash server startup.
|
||||
- _find_llama_server_binary must keep the custom-root in search_roots
|
||||
when the inner resolve() comparison itself fails."""
|
||||
|
||||
from __future__ import annotations
|
||||
|
||||
import importlib.util
|
||||
import re
|
||||
import sys
|
||||
import textwrap
|
||||
from pathlib import Path
|
||||
from unittest import mock
|
||||
|
||||
import pytest
|
||||
|
||||
|
||||
REPO_ROOT = Path(__file__).resolve().parents[1]
|
||||
STORAGE_ROOTS = (
|
||||
REPO_ROOT / "studio" / "backend" / "utils" / "paths" / "storage_roots.py"
|
||||
)
|
||||
LLAMA_CPP = REPO_ROOT / "studio" / "backend" / "core" / "inference" / "llama_cpp.py"
|
||||
|
||||
|
||||
def _load(name: str, path: Path):
|
||||
spec = importlib.util.spec_from_file_location(name, path)
|
||||
assert spec is not None and spec.loader is not None
|
||||
mod = importlib.util.module_from_spec(spec)
|
||||
sys.modules[name] = mod
|
||||
spec.loader.exec_module(mod)
|
||||
return mod
|
||||
|
||||
|
||||
def test_infer_studio_home_swallows_permission_error(tmp_path, monkeypatch):
|
||||
candidate = tmp_path / "fake_root"
|
||||
venv = candidate / "unsloth_studio"
|
||||
venv.mkdir(parents = True)
|
||||
monkeypatch.setattr(sys, "prefix", str(venv))
|
||||
sys.modules.pop("sr_perm", None)
|
||||
mod = _load("sr_perm", STORAGE_ROOTS)
|
||||
with mock.patch.object(Path, "is_file", side_effect = PermissionError("denied")):
|
||||
# Must NOT raise.
|
||||
assert mod._infer_studio_home_from_venv() is None
|
||||
|
||||
|
||||
def test_studio_root_does_not_crash_on_permission_error(tmp_path, monkeypatch):
|
||||
"""studio_root() must remain callable even when the venv inference
|
||||
encounters a restricted filesystem; it should fall through to the
|
||||
legacy default."""
|
||||
candidate = tmp_path / "fake_root"
|
||||
venv = candidate / "unsloth_studio"
|
||||
venv.mkdir(parents = True)
|
||||
monkeypatch.setattr(sys, "prefix", str(venv))
|
||||
monkeypatch.delenv("UNSLOTH_STUDIO_HOME", raising = False)
|
||||
monkeypatch.delenv("STUDIO_HOME", raising = False)
|
||||
sys.modules.pop("sr_studio_perm", None)
|
||||
mod = _load("sr_studio_perm", STORAGE_ROOTS)
|
||||
with mock.patch.object(Path, "is_file", side_effect = OSError("ebusy")):
|
||||
result = mod.studio_root()
|
||||
assert result == Path.home() / ".unsloth" / "studio"
|
||||
|
||||
|
||||
def test_kill_orphan_catches_oserror_from_studio_root():
|
||||
"""_kill_orphaned_servers must catch (ImportError, OSError, ValueError)
|
||||
on the studio_root() probe specifically; the sister function
|
||||
_find_llama_server_binary uses the same broader catch on its own probe."""
|
||||
src = LLAMA_CPP.read_text()
|
||||
fn_start = src.index("def _kill_orphaned_servers")
|
||||
fn_body = src[fn_start : fn_start + 4000]
|
||||
# The studio_root() probe in this fn is the one that imports as `_sr`
|
||||
# and assigns `_resolved_sr = _sr()`. Find the except that closes it.
|
||||
probe_idx = fn_body.index("storage_roots import studio_root as _sr")
|
||||
# The matching except is the next `except ...:` after the inner
|
||||
# OSError/ValueError block that wraps resolve().
|
||||
after = fn_body[probe_idx:]
|
||||
# Skip over the inner `except (OSError, ValueError):` that wraps resolve().
|
||||
inner_idx = after.index("except (OSError, ValueError):")
|
||||
after_inner = after[inner_idx + len("except (OSError, ValueError):") :]
|
||||
outer_match = re.search(r"except\s*\(?[^)]*?\)?:", after_inner)
|
||||
assert outer_match, "outer except for studio_root probe missing"
|
||||
clause = outer_match.group(0)
|
||||
assert (
|
||||
"OSError" in clause and "ValueError" in clause
|
||||
), f"_kill_orphaned_servers studio_root probe catch too narrow: {clause!r}"
|
||||
|
||||
|
||||
def _exec_search_roots_block(
|
||||
home: Path, studio_root_value: Path, resolve_raises: bool
|
||||
) -> list[Path]:
|
||||
"""Extract _find_llama_server_binary's env-mode search_roots block
|
||||
and execute it with controlled inputs."""
|
||||
src = LLAMA_CPP.read_text()
|
||||
block_start = src.index('legacy_llama = Path.home() / ".unsloth" / "llama.cpp"')
|
||||
block_end = src.index("_seen_roots: set[str]", block_start)
|
||||
raw = src[block_start:block_end]
|
||||
indent = " " * 8
|
||||
block = textwrap.dedent(indent + raw)
|
||||
fake_module = type(sys)("fake_storage_roots")
|
||||
fake_module.studio_root = lambda: studio_root_value
|
||||
sys.modules["utils.paths.storage_roots"] = fake_module
|
||||
try:
|
||||
original_resolve = Path.resolve
|
||||
|
||||
def _resolve(self, *a, **k):
|
||||
if resolve_raises:
|
||||
raise OSError("ebusy")
|
||||
return original_resolve(self, *a, **k)
|
||||
|
||||
with (
|
||||
mock.patch.object(Path, "home", classmethod(lambda cls: home)),
|
||||
mock.patch.object(Path, "resolve", _resolve),
|
||||
):
|
||||
ns: dict = {"Path": Path}
|
||||
exec(block, ns) # noqa: S102
|
||||
return ns["search_roots"]
|
||||
finally:
|
||||
sys.modules.pop("utils.paths.storage_roots", None)
|
||||
|
||||
|
||||
def test_search_roots_keeps_custom_when_resolve_fails(tmp_path):
|
||||
home = tmp_path / "home"
|
||||
home.mkdir()
|
||||
custom = tmp_path / "custom_studio"
|
||||
custom.mkdir()
|
||||
roots = _exec_search_roots_block(
|
||||
home = home, studio_root_value = custom, resolve_raises = True
|
||||
)
|
||||
# On resolve() failure, the inner except falls back to direct equality;
|
||||
# custom != legacy_studio so the custom root must remain in search_roots.
|
||||
assert (
|
||||
custom / "llama.cpp" in roots
|
||||
), f"custom root dropped on resolve() failure: {roots}"
|
||||
# custom-mode discovery excludes the legacy tree to match _kill_orphaned_servers.
|
||||
assert (
|
||||
(home / ".unsloth" / "llama.cpp") not in roots
|
||||
), f"legacy llama path must not appear in custom-mode search_roots: {roots}"
|
||||
|
||||
|
||||
def test_search_roots_default_mode_uses_legacy_only(tmp_path):
|
||||
home = tmp_path / "home"
|
||||
home.mkdir()
|
||||
legacy = home / ".unsloth" / "studio"
|
||||
legacy.mkdir(parents = True)
|
||||
roots = _exec_search_roots_block(
|
||||
home = home, studio_root_value = legacy, resolve_raises = False
|
||||
)
|
||||
# Default mode: only legacy_llama.
|
||||
assert roots == [home / ".unsloth" / "llama.cpp"]
|
||||
|
|
@ -12,348 +12,117 @@
|
|||
# See the License for the specific language governing permissions and
|
||||
# limitations under the License.
|
||||
|
||||
import warnings, importlib, sys
|
||||
from packaging.version import Version
|
||||
import os, re, subprocess, inspect, functools
|
||||
import numpy as np
|
||||
import os, platform, importlib.util
|
||||
|
||||
# Log Unsloth is being used
|
||||
os.environ["UNSLOTH_IS_PRESENT"] = "1"
|
||||
|
||||
# Check if modules that need patching are already imported
|
||||
critical_modules = ["trl", "transformers", "peft"]
|
||||
already_imported = [mod for mod in critical_modules if mod in sys.modules]
|
||||
|
||||
# Fix some issues before importing other packages
|
||||
from .import_fixes import (
|
||||
fix_message_factory_issue,
|
||||
check_fbgemm_gpu_version,
|
||||
disable_broken_causal_conv1d,
|
||||
disable_broken_vllm,
|
||||
configure_amdgpu_asic_id_table_path,
|
||||
torchvision_compatibility_check,
|
||||
fix_diffusers_warnings,
|
||||
fix_huggingface_hub,
|
||||
# Detect Apple Silicon + MLX before any torch/numpy imports
|
||||
_IS_MLX = (
|
||||
platform.system() == "Darwin"
|
||||
and platform.machine() == "arm64"
|
||||
and importlib.util.find_spec("mlx") is not None
|
||||
)
|
||||
|
||||
# Configure libdrm ids table path early so ROCm can resolve AMD GPU names.
|
||||
configure_amdgpu_asic_id_table_path()
|
||||
disable_broken_causal_conv1d()
|
||||
disable_broken_vllm()
|
||||
fix_message_factory_issue()
|
||||
check_fbgemm_gpu_version()
|
||||
torchvision_compatibility_check()
|
||||
fix_diffusers_warnings()
|
||||
fix_huggingface_hub()
|
||||
del configure_amdgpu_asic_id_table_path
|
||||
del disable_broken_causal_conv1d
|
||||
del disable_broken_vllm
|
||||
del fix_message_factory_issue
|
||||
del check_fbgemm_gpu_version
|
||||
del torchvision_compatibility_check
|
||||
del fix_diffusers_warnings
|
||||
del fix_huggingface_hub
|
||||
|
||||
# This check is critical because Unsloth optimizes these libraries by modifying
|
||||
# their code at import time. If they're imported first, the original (slower,
|
||||
# more memory-intensive) implementations will be used instead of Unsloth's
|
||||
# optimized versions, potentially causing OOM errors or slower training.
|
||||
if already_imported:
|
||||
# stacklevel=2 makes warning point to user's import line rather than this library code,
|
||||
# showing them exactly where to fix the import order in their script
|
||||
warnings.warn(
|
||||
f"WARNING: Unsloth should be imported before [{', '.join(already_imported)}] "
|
||||
f"to ensure all optimizations are applied. Your code may run slower or encounter "
|
||||
f"memory issues without these optimizations.\n\n"
|
||||
f"Please restructure your imports with 'import unsloth' at the top of your file.",
|
||||
stacklevel = 2,
|
||||
)
|
||||
del already_imported, critical_modules
|
||||
|
||||
# Unsloth currently does not work on multi GPU setups - sadly we are a 2 brother team so
|
||||
# enabling it will require much more work, so we have to prioritize. Please understand!
|
||||
# We do have a beta version, which you can contact us about!
|
||||
# Thank you for your understanding and we appreciate it immensely!
|
||||
|
||||
# Fixes https://github.com/unslothai/unsloth/issues/1266
|
||||
os.environ["PROTOCOL_BUFFERS_PYTHON_IMPLEMENTATION"] = "python"
|
||||
|
||||
# [TODO] Check why some GPUs don't work
|
||||
# "pinned_use_cuda_host_register:True,"\
|
||||
# "pinned_num_register_threads:8"
|
||||
|
||||
|
||||
from importlib.metadata import version as importlib_version
|
||||
from importlib.metadata import PackageNotFoundError
|
||||
|
||||
# Check for unsloth_zoo
|
||||
try:
|
||||
unsloth_zoo_version = importlib_version("unsloth_zoo")
|
||||
if Version(unsloth_zoo_version) < Version("2026.3.4"):
|
||||
print(
|
||||
"Unsloth: Please update Unsloth and Unsloth-Zoo to the latest version!\n"
|
||||
"Do this via `pip install --upgrade --force-reinstall --no-cache-dir --no-deps unsloth unsloth_zoo`"
|
||||
)
|
||||
# if os.environ.get("UNSLOTH_DISABLE_AUTO_UPDATES", "0") == "0":
|
||||
# try:
|
||||
# os.system("pip install --upgrade --no-cache-dir --no-deps unsloth_zoo")
|
||||
# except:
|
||||
# try:
|
||||
# os.system("pip install --upgrade --no-cache-dir --no-deps --user unsloth_zoo")
|
||||
# except:
|
||||
# raise ImportError("Unsloth: Please update unsloth_zoo via `pip install --upgrade --no-cache-dir --no-deps unsloth_zoo`")
|
||||
import unsloth_zoo
|
||||
except PackageNotFoundError:
|
||||
raise ImportError(
|
||||
f"Unsloth: Please install unsloth_zoo via `pip install unsloth_zoo` then retry!"
|
||||
)
|
||||
except:
|
||||
raise
|
||||
del PackageNotFoundError, importlib_version
|
||||
|
||||
# Try importing PyTorch and check version
|
||||
try:
|
||||
import torch
|
||||
except ModuleNotFoundError:
|
||||
raise ImportError(
|
||||
"Unsloth: Pytorch is not installed. Go to https://pytorch.org/.\n"
|
||||
"We have some installation instructions on our Github page."
|
||||
)
|
||||
except:
|
||||
raise
|
||||
|
||||
from unsloth_zoo.device_type import (
|
||||
is_hip,
|
||||
get_device_type,
|
||||
DEVICE_TYPE,
|
||||
DEVICE_TYPE_TORCH,
|
||||
DEVICE_COUNT,
|
||||
ALLOW_PREQUANTIZED_MODELS,
|
||||
)
|
||||
|
||||
# Fix other issues
|
||||
from .import_fixes import (
|
||||
fix_xformers_performance_issue,
|
||||
fix_vllm_aimv2_issue,
|
||||
check_vllm_torch_sm100_compatibility,
|
||||
fix_vllm_guided_decoding_params,
|
||||
fix_trl_vllm_ascend,
|
||||
fix_vllm_pdl_blackwell,
|
||||
fix_triton_compiled_kernel_missing_attrs,
|
||||
patch_trunc_normal_precision_issue,
|
||||
ignore_logger_messages,
|
||||
patch_ipykernel_hf_xet,
|
||||
patch_trackio,
|
||||
patch_datasets,
|
||||
patch_enable_input_require_grads,
|
||||
fix_openenv_no_vllm,
|
||||
patch_openspiel_env_async,
|
||||
fix_executorch,
|
||||
patch_vllm_for_notebooks,
|
||||
patch_torchcodec_audio_decoder,
|
||||
disable_torchcodec_if_broken,
|
||||
disable_broken_wandb,
|
||||
patch_peft_weight_converter_compatibility,
|
||||
)
|
||||
|
||||
fix_xformers_performance_issue()
|
||||
fix_vllm_aimv2_issue()
|
||||
# Check vLLM + torch < 2.9.0 + SM100 compatibility BEFORE importing vLLM
|
||||
check_vllm_torch_sm100_compatibility()
|
||||
fix_vllm_guided_decoding_params()
|
||||
fix_trl_vllm_ascend()
|
||||
fix_vllm_pdl_blackwell()
|
||||
fix_triton_compiled_kernel_missing_attrs()
|
||||
patch_trunc_normal_precision_issue()
|
||||
ignore_logger_messages()
|
||||
patch_ipykernel_hf_xet()
|
||||
patch_trackio()
|
||||
patch_datasets()
|
||||
patch_enable_input_require_grads()
|
||||
fix_openenv_no_vllm()
|
||||
patch_openspiel_env_async()
|
||||
fix_executorch()
|
||||
patch_vllm_for_notebooks()
|
||||
patch_torchcodec_audio_decoder()
|
||||
disable_torchcodec_if_broken()
|
||||
disable_broken_wandb()
|
||||
patch_peft_weight_converter_compatibility()
|
||||
|
||||
del fix_xformers_performance_issue
|
||||
del fix_vllm_aimv2_issue
|
||||
del check_vllm_torch_sm100_compatibility
|
||||
del fix_vllm_guided_decoding_params
|
||||
del fix_trl_vllm_ascend
|
||||
del fix_vllm_pdl_blackwell
|
||||
del fix_triton_compiled_kernel_missing_attrs
|
||||
del patch_trunc_normal_precision_issue
|
||||
del ignore_logger_messages
|
||||
del patch_ipykernel_hf_xet
|
||||
del patch_trackio
|
||||
del patch_datasets
|
||||
del patch_enable_input_require_grads
|
||||
del fix_openenv_no_vllm
|
||||
del patch_openspiel_env_async
|
||||
del fix_executorch
|
||||
del patch_vllm_for_notebooks
|
||||
del patch_torchcodec_audio_decoder
|
||||
del disable_torchcodec_if_broken
|
||||
del disable_broken_wandb
|
||||
del patch_peft_weight_converter_compatibility
|
||||
|
||||
# Torch 2.4 has including_emulation
|
||||
if DEVICE_TYPE == "cuda":
|
||||
major_version, minor_version = torch.cuda.get_device_capability()
|
||||
SUPPORTS_BFLOAT16 = major_version >= 8
|
||||
|
||||
old_is_bf16_supported = torch.cuda.is_bf16_supported
|
||||
if "including_emulation" in str(inspect.signature(old_is_bf16_supported)):
|
||||
|
||||
def is_bf16_supported(including_emulation = False):
|
||||
return old_is_bf16_supported(including_emulation)
|
||||
|
||||
torch.cuda.is_bf16_supported = is_bf16_supported
|
||||
else:
|
||||
|
||||
def is_bf16_supported():
|
||||
return SUPPORTS_BFLOAT16
|
||||
|
||||
torch.cuda.is_bf16_supported = is_bf16_supported
|
||||
del major_version, minor_version
|
||||
elif DEVICE_TYPE == "hip":
|
||||
SUPPORTS_BFLOAT16 = torch.cuda.is_bf16_supported()
|
||||
elif DEVICE_TYPE == "xpu":
|
||||
# torch.xpu.is_bf16_supported() does not have including_emulation
|
||||
# set SUPPORTS_BFLOAT16 as torch.xpu.is_bf16_supported()
|
||||
SUPPORTS_BFLOAT16 = torch.xpu.is_bf16_supported()
|
||||
|
||||
# For Gradio HF Spaces?
|
||||
# if "SPACE_AUTHOR_NAME" not in os.environ and "SPACE_REPO_NAME" not in os.environ:
|
||||
import triton
|
||||
|
||||
if DEVICE_TYPE == "cuda":
|
||||
libcuda_dirs = lambda: None
|
||||
if Version(triton.__version__) >= Version("3.0.0"):
|
||||
try:
|
||||
from triton.backends.nvidia.driver import libcuda_dirs
|
||||
except:
|
||||
pass
|
||||
else:
|
||||
from triton.common.build import libcuda_dirs
|
||||
|
||||
# Try loading bitsandbytes and triton
|
||||
if _IS_MLX:
|
||||
try:
|
||||
import bitsandbytes as bnb
|
||||
except:
|
||||
print(
|
||||
"Unsloth: `bitsandbytes` is not installed - 4bit QLoRA unallowed, but 16bit and full finetuning works!"
|
||||
)
|
||||
bnb = None
|
||||
import unsloth_zoo
|
||||
except ImportError as _e:
|
||||
raise ImportError(
|
||||
"Unsloth: MLX support requires `unsloth-zoo` with MLX modules. "
|
||||
"Reinstall with `pip install unsloth-zoo` or rerun install.sh."
|
||||
) from _e
|
||||
# The mlx_trainer / mlx_loader submodules ship with unsloth-zoo's MLX
|
||||
# support. An older installed unsloth-zoo (e.g. from PyPI before the
|
||||
# MLX release lands) will satisfy `import unsloth_zoo` but be missing
|
||||
# these submodules. Surface the same friendly install hint instead of
|
||||
# a raw ImportError on the submodule path.
|
||||
try:
|
||||
cdequantize_blockwise_fp32 = bnb.functional.lib.cdequantize_blockwise_fp32
|
||||
libcuda_dirs()
|
||||
except:
|
||||
# Only run the ldconfig recovery when we can actually run
|
||||
# ldconfig (root). On non-root environments (shared HPC,
|
||||
# locked-down containers, CI runners, etc.) the recovery would
|
||||
# shell out to `ldconfig` and fail with "Permission denied",
|
||||
# which is especially noisy for users who don't even have
|
||||
# bitsandbytes installed and are just doing 16bit/full
|
||||
# finetuning. libcuda_dirs() is used by both triton and bnb,
|
||||
# so we still run the recovery whenever we're root, regardless
|
||||
# of whether bnb is installed.
|
||||
if hasattr(os, "geteuid") and os.geteuid() == 0:
|
||||
warnings.warn("Unsloth: Running `ldconfig /usr/lib64-nvidia` to link CUDA.")
|
||||
from unsloth_zoo.mlx_trainer import MLXTrainer, MLXTrainingConfig
|
||||
from unsloth_zoo.mlx_loader import FastMLXModel
|
||||
except ImportError as _e:
|
||||
raise ImportError(
|
||||
"Unsloth: MLX support requires an unsloth-zoo build that includes "
|
||||
"`unsloth_zoo.mlx_trainer` and `unsloth_zoo.mlx_loader`. Upgrade with "
|
||||
"`pip install -U unsloth-zoo` or rerun install.sh."
|
||||
) from _e
|
||||
|
||||
if os.path.exists("/usr/lib64-nvidia"):
|
||||
os.system("ldconfig /usr/lib64-nvidia")
|
||||
elif os.path.exists("/usr/local"):
|
||||
# Sometimes bitsandbytes cannot be linked properly in Runpod for example
|
||||
possible_cudas = (
|
||||
subprocess.check_output(["ls", "-al", "/usr/local"])
|
||||
.decode("utf-8")
|
||||
.split("\n")
|
||||
)
|
||||
find_cuda = re.compile(r"[\s](cuda\-[\d\.]{2,})$")
|
||||
possible_cudas = [find_cuda.search(x) for x in possible_cudas]
|
||||
possible_cudas = [x.group(1) for x in possible_cudas if x is not None]
|
||||
# Load raw_text helpers without executing dataprep/__init__.py, which
|
||||
# imports synthetic.py -> torch and would defeat the torch-free MLX path.
|
||||
from pathlib import Path as _Path
|
||||
|
||||
# Try linking cuda folder, or everything in local
|
||||
if len(possible_cudas) == 0:
|
||||
os.system("ldconfig /usr/local/")
|
||||
else:
|
||||
find_number = re.compile(r"([\d\.]{2,})")
|
||||
latest_cuda = np.argsort(
|
||||
[float(find_number.search(x).group(1)) for x in possible_cudas]
|
||||
)[::-1][0]
|
||||
latest_cuda = possible_cudas[latest_cuda]
|
||||
os.system(f"ldconfig /usr/local/{latest_cuda}")
|
||||
del find_number, latest_cuda
|
||||
del possible_cudas, find_cuda
|
||||
_raw_text_path = _Path(__file__).resolve().parent / "dataprep" / "raw_text.py"
|
||||
_raw_text_spec = importlib.util.spec_from_file_location(
|
||||
"unsloth._mlx_raw_text", _raw_text_path
|
||||
)
|
||||
if _raw_text_spec is None or _raw_text_spec.loader is None:
|
||||
raise ImportError("Unsloth: could not load MLX raw_text dataprep helpers.")
|
||||
_raw_text = importlib.util.module_from_spec(_raw_text_spec)
|
||||
_raw_text_spec.loader.exec_module(_raw_text)
|
||||
RawTextDataLoader = _raw_text.RawTextDataLoader
|
||||
TextPreprocessor = _raw_text.TextPreprocessor
|
||||
del _raw_text, _raw_text_spec, _raw_text_path, _Path
|
||||
|
||||
if bnb is not None:
|
||||
importlib.reload(bnb)
|
||||
importlib.reload(triton)
|
||||
try:
|
||||
libcuda_dirs = lambda: None
|
||||
if Version(triton.__version__) >= Version("3.0.0"):
|
||||
try:
|
||||
from triton.backends.nvidia.driver import libcuda_dirs
|
||||
except:
|
||||
pass
|
||||
else:
|
||||
from triton.common.build import libcuda_dirs
|
||||
cdequantize_blockwise_fp32 = (
|
||||
bnb.functional.lib.cdequantize_blockwise_fp32
|
||||
)
|
||||
libcuda_dirs()
|
||||
except:
|
||||
warnings.warn(
|
||||
"Unsloth: CUDA is not linked properly.\n"
|
||||
"Try running `python -m bitsandbytes` then `python -m xformers.info`\n"
|
||||
"We tried running `ldconfig /usr/lib64-nvidia` ourselves, but it didn't work.\n"
|
||||
"You need to run in your terminal `sudo ldconfig /usr/lib64-nvidia` yourself, then import Unsloth.\n"
|
||||
"Also try `sudo ldconfig /usr/local/cuda-xx.x` - find the latest cuda version.\n"
|
||||
"Unsloth will still run for now, but maybe it might crash - let's hope it works!"
|
||||
)
|
||||
elif bnb is not None:
|
||||
# Non-root + bnb installed: we can't run ldconfig ourselves,
|
||||
# but bnb is going to crash later when the user actually uses
|
||||
# 4bit quantization - tell them how to fix it manually so
|
||||
# they're not surprised by an opaque error down the road.
|
||||
warnings.warn(
|
||||
"Unsloth: CUDA is not linked properly.\n"
|
||||
"You need to run in your terminal `sudo ldconfig /usr/lib64-nvidia` yourself, then import Unsloth.\n"
|
||||
"Also try `sudo ldconfig /usr/local/cuda-xx.x` - find the latest cuda version.\n"
|
||||
"Unsloth will still run for now, but maybe it might crash - let's hope it works!"
|
||||
__version__ = unsloth_zoo.__version__
|
||||
DEVICE_TYPE = "mlx"
|
||||
|
||||
class FastLanguageModel:
|
||||
@staticmethod
|
||||
def from_pretrained(*args, **kwargs):
|
||||
return FastMLXModel.from_pretrained(*args, **kwargs)
|
||||
|
||||
@staticmethod
|
||||
def get_peft_model(*args, **kwargs):
|
||||
return FastMLXModel.get_peft_model(*args, **kwargs)
|
||||
|
||||
@staticmethod
|
||||
def for_inference(*args, **kwargs):
|
||||
return args[0] if args else None
|
||||
|
||||
class FastVisionModel(FastLanguageModel):
|
||||
@staticmethod
|
||||
def from_pretrained(*args, **kwargs):
|
||||
kwargs.setdefault("text_only", False)
|
||||
return FastMLXModel.from_pretrained(*args, **kwargs)
|
||||
|
||||
@staticmethod
|
||||
def for_training(*args, **kwargs):
|
||||
return args[0] if args else None
|
||||
|
||||
FastTextModel = FastLanguageModel
|
||||
FastModel = FastLanguageModel
|
||||
|
||||
class FastSentenceTransformer:
|
||||
@staticmethod
|
||||
def from_pretrained(*args, **kwargs):
|
||||
raise NotImplementedError(
|
||||
"Unsloth: FastSentenceTransformer is not yet supported on MLX."
|
||||
)
|
||||
del libcuda_dirs
|
||||
elif DEVICE_TYPE == "hip":
|
||||
# NO-OP for rocm device
|
||||
pass
|
||||
elif DEVICE_TYPE == "xpu":
|
||||
import bitsandbytes as bnb
|
||||
|
||||
# TODO: check triton for intel installed properly.
|
||||
pass
|
||||
@staticmethod
|
||||
def get_peft_model(*args, **kwargs):
|
||||
raise NotImplementedError(
|
||||
"Unsloth: FastSentenceTransformer is not yet supported on MLX."
|
||||
)
|
||||
|
||||
from .models import *
|
||||
from .models import __version__
|
||||
from .save import *
|
||||
from .chat_templates import *
|
||||
from .tokenizer_utils import *
|
||||
from .trainer import *
|
||||
def is_bfloat16_supported():
|
||||
try:
|
||||
import mlx.core as mx
|
||||
|
||||
# Export dataprep utilities for CLI and downstream users
|
||||
from .dataprep.raw_text import RawTextDataLoader, TextPreprocessor
|
||||
from unsloth_zoo.rl_environments import (
|
||||
check_python_modules,
|
||||
create_locked_down_function,
|
||||
execute_with_time_limit,
|
||||
Benchmarker,
|
||||
is_port_open,
|
||||
launch_openenv,
|
||||
)
|
||||
name = mx.device_info().get("device_name", "") or ""
|
||||
return not name.startswith(("Apple M1", "Apple M2"))
|
||||
except Exception:
|
||||
return True
|
||||
|
||||
# Patch TRL trainers for backwards compatibility
|
||||
_patch_trl_trainer()
|
||||
is_bf16_supported = is_bfloat16_supported
|
||||
|
||||
class UnslothVisionDataCollator:
|
||||
def __init__(self, *args, **kwargs):
|
||||
raise NotImplementedError(
|
||||
"Unsloth: UnslothVisionDataCollator is not used on MLX. "
|
||||
"Use the MLX trainer/data path instead."
|
||||
)
|
||||
|
||||
else:
|
||||
# GPU path: load everything from _gpu_init
|
||||
from ._gpu_init import *
|
||||
from ._gpu_init import __version__
|
||||
|
|
|
|||
346
unsloth/_gpu_init.py
Normal file
346
unsloth/_gpu_init.py
Normal file
|
|
@ -0,0 +1,346 @@
|
|||
# Copyright 2023-present Daniel Han-Chen & the Unsloth team. All rights reserved.
|
||||
#
|
||||
# Licensed under the Apache License, Version 2.0 (the "License");
|
||||
# you may not use this file except in compliance with the License.
|
||||
# You may obtain a copy of the License at
|
||||
#
|
||||
# http://www.apache.org/licenses/LICENSE-2.0
|
||||
#
|
||||
# Unless required by applicable law or agreed to in writing, software
|
||||
# distributed under the License is distributed on an "AS IS" BASIS,
|
||||
# WITHOUT WARRANTIES OR CONDITIONS OF ANY KIND, either express or implied.
|
||||
# See the License for the specific language governing permissions and
|
||||
# limitations under the License.
|
||||
|
||||
import warnings, importlib, sys
|
||||
from packaging.version import Version
|
||||
import os, re, subprocess, inspect, functools
|
||||
import numpy as np
|
||||
|
||||
# Log Unsloth is being used
|
||||
os.environ["UNSLOTH_IS_PRESENT"] = "1"
|
||||
|
||||
# Check if modules that need patching are already imported
|
||||
critical_modules = ["trl", "transformers", "peft"]
|
||||
already_imported = [mod for mod in critical_modules if mod in sys.modules]
|
||||
|
||||
# Fix some issues before importing other packages
|
||||
from .import_fixes import (
|
||||
fix_message_factory_issue,
|
||||
check_fbgemm_gpu_version,
|
||||
disable_broken_causal_conv1d,
|
||||
disable_broken_vllm,
|
||||
configure_amdgpu_asic_id_table_path,
|
||||
torchvision_compatibility_check,
|
||||
fix_diffusers_warnings,
|
||||
fix_huggingface_hub,
|
||||
)
|
||||
|
||||
# Configure libdrm ids table path early so ROCm can resolve AMD GPU names.
|
||||
configure_amdgpu_asic_id_table_path()
|
||||
disable_broken_causal_conv1d()
|
||||
disable_broken_vllm()
|
||||
fix_message_factory_issue()
|
||||
check_fbgemm_gpu_version()
|
||||
torchvision_compatibility_check()
|
||||
fix_diffusers_warnings()
|
||||
fix_huggingface_hub()
|
||||
del configure_amdgpu_asic_id_table_path
|
||||
del disable_broken_causal_conv1d
|
||||
del disable_broken_vllm
|
||||
del fix_message_factory_issue
|
||||
del check_fbgemm_gpu_version
|
||||
del torchvision_compatibility_check
|
||||
del fix_diffusers_warnings
|
||||
del fix_huggingface_hub
|
||||
|
||||
# This check is critical because Unsloth optimizes these libraries by modifying
|
||||
# their code at import time. If they're imported first, the original (slower,
|
||||
# more memory-intensive) implementations will be used instead of Unsloth's
|
||||
# optimized versions, potentially causing OOM errors or slower training.
|
||||
if already_imported:
|
||||
# stacklevel=2 makes warning point to user's import line rather than this library code,
|
||||
# showing them exactly where to fix the import order in their script
|
||||
warnings.warn(
|
||||
f"WARNING: Unsloth should be imported before [{', '.join(already_imported)}] "
|
||||
f"to ensure all optimizations are applied. Your code may run slower or encounter "
|
||||
f"memory issues without these optimizations.\n\n"
|
||||
f"Please restructure your imports with 'import unsloth' at the top of your file.",
|
||||
stacklevel = 2,
|
||||
)
|
||||
del already_imported, critical_modules
|
||||
|
||||
# Unsloth currently does not work on multi GPU setups - sadly we are a 2 brother team so
|
||||
# enabling it will require much more work, so we have to prioritize. Please understand!
|
||||
# We do have a beta version, which you can contact us about!
|
||||
# Thank you for your understanding and we appreciate it immensely!
|
||||
|
||||
# Fixes https://github.com/unslothai/unsloth/issues/1266
|
||||
os.environ["PROTOCOL_BUFFERS_PYTHON_IMPLEMENTATION"] = "python"
|
||||
|
||||
# [TODO] Check why some GPUs don't work
|
||||
# "pinned_use_cuda_host_register:True,"\
|
||||
# "pinned_num_register_threads:8"
|
||||
|
||||
|
||||
from importlib.metadata import version as importlib_version
|
||||
from importlib.metadata import PackageNotFoundError
|
||||
|
||||
# Check for unsloth_zoo
|
||||
try:
|
||||
unsloth_zoo_version = importlib_version("unsloth_zoo")
|
||||
if Version(unsloth_zoo_version) < Version("2026.3.4"):
|
||||
print(
|
||||
"Unsloth: Please update Unsloth and Unsloth-Zoo to the latest version!\n"
|
||||
"Do this via `pip install --upgrade --force-reinstall --no-cache-dir --no-deps unsloth unsloth_zoo`"
|
||||
)
|
||||
# if os.environ.get("UNSLOTH_DISABLE_AUTO_UPDATES", "0") == "0":
|
||||
# try:
|
||||
# os.system("pip install --upgrade --no-cache-dir --no-deps unsloth_zoo")
|
||||
# except:
|
||||
# try:
|
||||
# os.system("pip install --upgrade --no-cache-dir --no-deps --user unsloth_zoo")
|
||||
# except:
|
||||
# raise ImportError("Unsloth: Please update unsloth_zoo via `pip install --upgrade --no-cache-dir --no-deps unsloth_zoo`")
|
||||
import unsloth_zoo
|
||||
except PackageNotFoundError:
|
||||
raise ImportError(
|
||||
f"Unsloth: Please install unsloth_zoo via `pip install unsloth_zoo` then retry!"
|
||||
)
|
||||
except:
|
||||
raise
|
||||
del PackageNotFoundError, importlib_version
|
||||
|
||||
# Try importing PyTorch and check version
|
||||
try:
|
||||
import torch
|
||||
except ModuleNotFoundError:
|
||||
raise ImportError(
|
||||
"Unsloth: Pytorch is not installed. Go to https://pytorch.org/.\n"
|
||||
"We have some installation instructions on our Github page."
|
||||
)
|
||||
except:
|
||||
raise
|
||||
|
||||
from unsloth_zoo.device_type import (
|
||||
is_hip,
|
||||
get_device_type,
|
||||
DEVICE_TYPE,
|
||||
DEVICE_TYPE_TORCH,
|
||||
DEVICE_COUNT,
|
||||
ALLOW_PREQUANTIZED_MODELS,
|
||||
)
|
||||
|
||||
# Fix other issues
|
||||
from .import_fixes import (
|
||||
fix_xformers_performance_issue,
|
||||
fix_vllm_aimv2_issue,
|
||||
check_vllm_torch_sm100_compatibility,
|
||||
fix_vllm_guided_decoding_params,
|
||||
fix_vllm_pdl_blackwell,
|
||||
fix_triton_compiled_kernel_missing_attrs,
|
||||
patch_trunc_normal_precision_issue,
|
||||
ignore_logger_messages,
|
||||
patch_ipykernel_hf_xet,
|
||||
patch_trackio,
|
||||
patch_datasets,
|
||||
patch_enable_input_require_grads,
|
||||
fix_openenv_no_vllm,
|
||||
patch_openspiel_env_async,
|
||||
fix_executorch,
|
||||
patch_vllm_for_notebooks,
|
||||
patch_torchcodec_audio_decoder,
|
||||
disable_torchcodec_if_broken,
|
||||
disable_broken_wandb,
|
||||
fix_trl_vllm_ascend,
|
||||
patch_peft_weight_converter_compatibility,
|
||||
)
|
||||
|
||||
fix_xformers_performance_issue()
|
||||
fix_vllm_aimv2_issue()
|
||||
# Check vLLM + torch < 2.9.0 + SM100 compatibility BEFORE importing vLLM
|
||||
check_vllm_torch_sm100_compatibility()
|
||||
fix_vllm_guided_decoding_params()
|
||||
fix_trl_vllm_ascend()
|
||||
fix_vllm_pdl_blackwell()
|
||||
fix_triton_compiled_kernel_missing_attrs()
|
||||
patch_trunc_normal_precision_issue()
|
||||
ignore_logger_messages()
|
||||
patch_ipykernel_hf_xet()
|
||||
patch_trackio()
|
||||
patch_datasets()
|
||||
patch_enable_input_require_grads()
|
||||
fix_openenv_no_vllm()
|
||||
patch_openspiel_env_async()
|
||||
fix_executorch()
|
||||
patch_vllm_for_notebooks()
|
||||
patch_torchcodec_audio_decoder()
|
||||
disable_torchcodec_if_broken()
|
||||
disable_broken_wandb()
|
||||
patch_peft_weight_converter_compatibility()
|
||||
|
||||
del fix_xformers_performance_issue
|
||||
del fix_vllm_aimv2_issue
|
||||
del check_vllm_torch_sm100_compatibility
|
||||
del fix_vllm_guided_decoding_params
|
||||
del fix_trl_vllm_ascend
|
||||
del fix_vllm_pdl_blackwell
|
||||
del fix_triton_compiled_kernel_missing_attrs
|
||||
del patch_trunc_normal_precision_issue
|
||||
del ignore_logger_messages
|
||||
del patch_ipykernel_hf_xet
|
||||
del patch_trackio
|
||||
del patch_datasets
|
||||
del patch_enable_input_require_grads
|
||||
del fix_openenv_no_vllm
|
||||
del patch_openspiel_env_async
|
||||
del fix_executorch
|
||||
del patch_vllm_for_notebooks
|
||||
del patch_torchcodec_audio_decoder
|
||||
del disable_torchcodec_if_broken
|
||||
del disable_broken_wandb
|
||||
del patch_peft_weight_converter_compatibility
|
||||
|
||||
# Torch 2.4 has including_emulation
|
||||
if DEVICE_TYPE == "cuda":
|
||||
major_version, minor_version = torch.cuda.get_device_capability()
|
||||
SUPPORTS_BFLOAT16 = major_version >= 8
|
||||
|
||||
old_is_bf16_supported = torch.cuda.is_bf16_supported
|
||||
if "including_emulation" in str(inspect.signature(old_is_bf16_supported)):
|
||||
|
||||
def is_bf16_supported(including_emulation = False):
|
||||
return old_is_bf16_supported(including_emulation)
|
||||
|
||||
torch.cuda.is_bf16_supported = is_bf16_supported
|
||||
else:
|
||||
|
||||
def is_bf16_supported():
|
||||
return SUPPORTS_BFLOAT16
|
||||
|
||||
torch.cuda.is_bf16_supported = is_bf16_supported
|
||||
del major_version, minor_version
|
||||
elif DEVICE_TYPE == "hip":
|
||||
SUPPORTS_BFLOAT16 = torch.cuda.is_bf16_supported()
|
||||
elif DEVICE_TYPE == "xpu":
|
||||
# torch.xpu.is_bf16_supported() does not have including_emulation
|
||||
# set SUPPORTS_BFLOAT16 as torch.xpu.is_bf16_supported()
|
||||
SUPPORTS_BFLOAT16 = torch.xpu.is_bf16_supported()
|
||||
|
||||
# For Gradio HF Spaces?
|
||||
# if "SPACE_AUTHOR_NAME" not in os.environ and "SPACE_REPO_NAME" not in os.environ:
|
||||
import triton
|
||||
|
||||
if DEVICE_TYPE == "cuda":
|
||||
libcuda_dirs = lambda: None
|
||||
if Version(triton.__version__) >= Version("3.0.0"):
|
||||
try:
|
||||
from triton.backends.nvidia.driver import libcuda_dirs
|
||||
except:
|
||||
pass
|
||||
else:
|
||||
from triton.common.build import libcuda_dirs
|
||||
|
||||
# Try loading bitsandbytes and triton
|
||||
try:
|
||||
import bitsandbytes as bnb
|
||||
except:
|
||||
print(
|
||||
"Unsloth: `bitsandbytes` is not installed - 4bit QLoRA unallowed, but 16bit and full finetuning works!"
|
||||
)
|
||||
bnb = None
|
||||
try:
|
||||
cdequantize_blockwise_fp32 = bnb.functional.lib.cdequantize_blockwise_fp32
|
||||
libcuda_dirs()
|
||||
except:
|
||||
if hasattr(os, "geteuid") and os.geteuid() == 0:
|
||||
warnings.warn("Unsloth: Running `ldconfig /usr/lib64-nvidia` to link CUDA.")
|
||||
|
||||
if os.path.exists("/usr/lib64-nvidia"):
|
||||
os.system("ldconfig /usr/lib64-nvidia")
|
||||
elif os.path.exists("/usr/local"):
|
||||
# Sometimes bitsandbytes cannot be linked properly in Runpod for example
|
||||
possible_cudas = (
|
||||
subprocess.check_output(["ls", "-al", "/usr/local"])
|
||||
.decode("utf-8")
|
||||
.split("\n")
|
||||
)
|
||||
find_cuda = re.compile(r"[\s](cuda\-[\d\.]{2,})$")
|
||||
possible_cudas = [find_cuda.search(x) for x in possible_cudas]
|
||||
possible_cudas = [x.group(1) for x in possible_cudas if x is not None]
|
||||
|
||||
# Try linking cuda folder, or everything in local
|
||||
if len(possible_cudas) == 0:
|
||||
os.system("ldconfig /usr/local/")
|
||||
else:
|
||||
find_number = re.compile(r"([\d\.]{2,})")
|
||||
latest_cuda = np.argsort(
|
||||
[float(find_number.search(x).group(1)) for x in possible_cudas]
|
||||
)[::-1][0]
|
||||
latest_cuda = possible_cudas[latest_cuda]
|
||||
os.system(f"ldconfig /usr/local/{latest_cuda}")
|
||||
del find_number, latest_cuda
|
||||
del possible_cudas, find_cuda
|
||||
|
||||
if bnb is not None:
|
||||
importlib.reload(bnb)
|
||||
importlib.reload(triton)
|
||||
try:
|
||||
libcuda_dirs = lambda: None
|
||||
if Version(triton.__version__) >= Version("3.0.0"):
|
||||
try:
|
||||
from triton.backends.nvidia.driver import libcuda_dirs
|
||||
except:
|
||||
pass
|
||||
else:
|
||||
from triton.common.build import libcuda_dirs
|
||||
cdequantize_blockwise_fp32 = (
|
||||
bnb.functional.lib.cdequantize_blockwise_fp32
|
||||
)
|
||||
libcuda_dirs()
|
||||
except:
|
||||
warnings.warn(
|
||||
"Unsloth: CUDA is not linked properly.\n"
|
||||
"Try running `python -m bitsandbytes` then `python -m xformers.info`\n"
|
||||
"We tried running `ldconfig /usr/lib64-nvidia` ourselves, but it didn't work.\n"
|
||||
"You need to run in your terminal `sudo ldconfig /usr/lib64-nvidia` yourself, then import Unsloth.\n"
|
||||
"Also try `sudo ldconfig /usr/local/cuda-xx.x` - find the latest cuda version.\n"
|
||||
"Unsloth will still run for now, but maybe it might crash - let's hope it works!"
|
||||
)
|
||||
elif bnb is not None:
|
||||
warnings.warn(
|
||||
"Unsloth: CUDA is not linked properly.\n"
|
||||
"You need to run in your terminal `sudo ldconfig /usr/lib64-nvidia` yourself, then import Unsloth.\n"
|
||||
"Also try `sudo ldconfig /usr/local/cuda-xx.x` - find the latest cuda version.\n"
|
||||
"Unsloth will still run for now, but maybe it might crash - let's hope it works!"
|
||||
)
|
||||
del libcuda_dirs
|
||||
elif DEVICE_TYPE == "hip":
|
||||
# NO-OP for rocm device
|
||||
pass
|
||||
elif DEVICE_TYPE == "xpu":
|
||||
import bitsandbytes as bnb
|
||||
|
||||
# TODO: check triton for intel installed properly.
|
||||
pass
|
||||
|
||||
from .models import *
|
||||
from .models import __version__
|
||||
from .save import *
|
||||
from .chat_templates import *
|
||||
from .tokenizer_utils import *
|
||||
from .trainer import *
|
||||
|
||||
# Export dataprep utilities for CLI and downstream users
|
||||
from .dataprep.raw_text import RawTextDataLoader, TextPreprocessor
|
||||
from unsloth_zoo.rl_environments import (
|
||||
check_python_modules,
|
||||
create_locked_down_function,
|
||||
execute_with_time_limit,
|
||||
Benchmarker,
|
||||
is_port_open,
|
||||
launch_openenv,
|
||||
)
|
||||
|
||||
# Patch TRL trainers for backwards compatibility
|
||||
_patch_trl_trainer()
|
||||
|
|
@ -12,7 +12,7 @@
|
|||
# See the License for the specific language governing permissions and
|
||||
# limitations under the License.
|
||||
|
||||
__version__ = "2026.5.1"
|
||||
__version__ = "2026.5.2"
|
||||
|
||||
__all__ = [
|
||||
"SUPPORTS_BFLOAT16",
|
||||
|
|
|
|||
|
|
@ -23,7 +23,9 @@ from functools import wraps
|
|||
import trl
|
||||
import inspect
|
||||
from trl import SFTTrainer
|
||||
from . import is_bfloat16_supported
|
||||
|
||||
# why: bypass partially-initialised unsloth ns during _gpu_init load
|
||||
from .models._utils import is_bfloat16_supported
|
||||
from unsloth.utils import (
|
||||
configure_padding_free,
|
||||
configure_sample_packing,
|
||||
|
|
|
|||
|
|
@ -20,7 +20,73 @@ import typer
|
|||
|
||||
studio_app = typer.Typer(help = "Unsloth Studio commands.")
|
||||
|
||||
STUDIO_HOME = Path.home() / ".unsloth" / "studio"
|
||||
|
||||
# Resolve install root: UNSLOTH_STUDIO_HOME, then STUDIO_HOME alias, then
|
||||
# sys.prefix inference (so a direct call to <root>/bin/unsloth resolves after
|
||||
# the installer's env var has expired), then legacy ~/.unsloth/studio.
|
||||
# UNSLOTH_STUDIO_HOME wins when both env vars are set.
|
||||
def _looks_like_installer_managed_studio_home(candidate: Path) -> bool:
|
||||
"""Sentinel check (studio.conf or bin shim) so a dev venv named
|
||||
unsloth_studio is not misidentified as a custom Studio root.
|
||||
"""
|
||||
shim_name = "unsloth.exe" if platform.system() == "Windows" else "unsloth"
|
||||
return (candidate / "share" / "studio.conf").is_file() or (
|
||||
candidate / "bin" / shim_name
|
||||
).is_file()
|
||||
|
||||
|
||||
def _resolve_studio_home() -> tuple[Path, bool]:
|
||||
override = (os.environ.get("UNSLOTH_STUDIO_HOME") or "").strip()
|
||||
if not override:
|
||||
override = (os.environ.get("STUDIO_HOME") or "").strip()
|
||||
if override:
|
||||
try:
|
||||
return Path(override).expanduser().resolve(), True
|
||||
except (OSError, ValueError):
|
||||
return Path(override).expanduser(), True
|
||||
try:
|
||||
prefix = Path(sys.prefix).resolve()
|
||||
if prefix.name == "unsloth_studio":
|
||||
inferred = prefix.parent
|
||||
legacy = (Path.home() / ".unsloth" / "studio").resolve()
|
||||
if inferred != legacy and _looks_like_installer_managed_studio_home(
|
||||
inferred
|
||||
):
|
||||
return inferred, True
|
||||
except (OSError, ValueError):
|
||||
pass
|
||||
return Path.home() / ".unsloth" / "studio", False
|
||||
|
||||
|
||||
STUDIO_HOME, _STUDIO_HOME_IS_CUSTOM = _resolve_studio_home()
|
||||
|
||||
|
||||
def _ensure_studio_env_exported() -> None:
|
||||
"""Re-export UNSLOTH_STUDIO_HOME / UNSLOTH_LLAMA_CPP_PATH only for real
|
||||
custom roots so subprocesses inherit the right install. Called from each
|
||||
studio subcommand entry rather than at import time, to avoid leaking env
|
||||
state into unrelated importers (tests, --help, CLI introspection).
|
||||
"""
|
||||
if not _STUDIO_HOME_IS_CUSTOM:
|
||||
return
|
||||
# Truthy-check (not setdefault) so a blank UNSLOTH_STUDIO_HOME= does not
|
||||
# suppress the inferred custom root.
|
||||
if not os.environ.get("UNSLOTH_STUDIO_HOME"):
|
||||
os.environ["UNSLOTH_STUDIO_HOME"] = str(STUDIO_HOME)
|
||||
# When override == legacy default, llama.cpp stays at ~/.unsloth/llama.cpp.
|
||||
try:
|
||||
_legacy_studio = (Path.home() / ".unsloth" / "studio").resolve()
|
||||
_is_legacy = STUDIO_HOME.resolve() == _legacy_studio
|
||||
except (OSError, ValueError):
|
||||
_is_legacy = STUDIO_HOME == (Path.home() / ".unsloth" / "studio")
|
||||
if _is_legacy:
|
||||
_llama_dir = Path.home() / ".unsloth" / "llama.cpp"
|
||||
else:
|
||||
_llama_dir = STUDIO_HOME / "llama.cpp"
|
||||
if not os.environ.get("UNSLOTH_LLAMA_CPP_PATH"):
|
||||
os.environ["UNSLOTH_LLAMA_CPP_PATH"] = str(_llama_dir)
|
||||
|
||||
|
||||
BOOTSTRAP_PASSWORD_FILE = ".bootstrap_password"
|
||||
DESKTOP_SECRET_FILE = ".desktop_secret"
|
||||
DEFAULT_ADMIN_USERNAME = "unsloth"
|
||||
|
|
@ -427,6 +493,8 @@ def studio_default(
|
|||
),
|
||||
):
|
||||
"""Launch the Unsloth Studio server."""
|
||||
# Runs before any subcommand; covers run/setup/update/etc in one place.
|
||||
_ensure_studio_env_exported()
|
||||
if ctx.invoked_subcommand is not None:
|
||||
return
|
||||
|
||||
|
|
|
|||
Loading…
Add table
Add a link
Reference in a new issue