mirror of
https://github.com/jamiepine/voicebox.git
synced 2026-10-03 00:55:14 -07:00
fix: download progress tracking for all engines and inline progress UI
- Add HFProgressTracker to LuxTTS and Chatterbox backends so tqdm-based file-level download progress reaches the frontend (previously only Qwen had this, LuxTTS/Chatterbox showed a static spinner) - Add progress/current/total/filename fields to ActiveDownloadTask so the /tasks/active polling endpoint carries progress data - Show inline progress bar + bytes in the model list and detail modal, poll at 1s during active downloads (5s otherwise) - Fix GpuAcceleration crash: cudaStatusLoading was referenced before initialization in its own useQuery declaration
This commit is contained in:
@@ -32,7 +32,7 @@ export function GpuAcceleration() {
|
|||||||
} = useQuery({
|
} = useQuery({
|
||||||
queryKey: ['cuda-status', serverUrl],
|
queryKey: ['cuda-status', serverUrl],
|
||||||
queryFn: () => apiClient.getCudaStatus(),
|
queryFn: () => apiClient.getCudaStatus(),
|
||||||
refetchInterval: cudaStatusLoading ? false : 10000,
|
refetchInterval: (query) => (query.state.status === 'pending' ? false : 10000),
|
||||||
retry: 1,
|
retry: 1,
|
||||||
enabled: !!health, // Only fetch when backend is reachable
|
enabled: !!health, // Only fetch when backend is reachable
|
||||||
});
|
});
|
||||||
|
|||||||
@@ -16,7 +16,7 @@ import {
|
|||||||
X,
|
X,
|
||||||
Zap,
|
Zap,
|
||||||
} from 'lucide-react';
|
} from 'lucide-react';
|
||||||
import { useCallback, useState } from 'react';
|
import { useCallback, useMemo, useState } from 'react';
|
||||||
import {
|
import {
|
||||||
AlertDialog,
|
AlertDialog,
|
||||||
AlertDialogAction,
|
AlertDialogAction,
|
||||||
@@ -36,6 +36,7 @@ import {
|
|||||||
DialogHeader,
|
DialogHeader,
|
||||||
DialogTitle,
|
DialogTitle,
|
||||||
} from '@/components/ui/dialog';
|
} from '@/components/ui/dialog';
|
||||||
|
import { Progress } from '@/components/ui/progress';
|
||||||
import { useToast } from '@/components/ui/use-toast';
|
import { useToast } from '@/components/ui/use-toast';
|
||||||
import { apiClient } from '@/lib/api/client';
|
import { apiClient } from '@/lib/api/client';
|
||||||
import type { ActiveDownloadTask, HuggingFaceModelInfo, ModelStatus } from '@/lib/api/types';
|
import type { ActiveDownloadTask, HuggingFaceModelInfo, ModelStatus } from '@/lib/api/types';
|
||||||
@@ -73,6 +74,14 @@ function formatPipelineTag(tag: string): string {
|
|||||||
.join(' ');
|
.join(' ');
|
||||||
}
|
}
|
||||||
|
|
||||||
|
function formatBytes(bytes: number): string {
|
||||||
|
if (bytes === 0) return '0 B';
|
||||||
|
const k = 1024;
|
||||||
|
const sizes = ['B', 'KB', 'MB', 'GB'];
|
||||||
|
const i = Math.floor(Math.log(bytes) / Math.log(k));
|
||||||
|
return `${(bytes / k ** i).toFixed(1)} ${sizes[i]}`;
|
||||||
|
}
|
||||||
|
|
||||||
export function ModelManagement() {
|
export function ModelManagement() {
|
||||||
const { toast } = useToast();
|
const { toast } = useToast();
|
||||||
const queryClient = useQueryClient();
|
const queryClient = useQueryClient();
|
||||||
@@ -98,7 +107,11 @@ export function ModelManagement() {
|
|||||||
const { data: activeTasks } = useQuery({
|
const { data: activeTasks } = useQuery({
|
||||||
queryKey: ['activeTasks'],
|
queryKey: ['activeTasks'],
|
||||||
queryFn: () => apiClient.getActiveTasks(),
|
queryFn: () => apiClient.getActiveTasks(),
|
||||||
refetchInterval: 5000,
|
refetchInterval: (query) => {
|
||||||
|
const data = query.state.data;
|
||||||
|
const hasActive = data?.downloads.some((d) => d.status === 'downloading');
|
||||||
|
return hasActive ? 1000 : 5000;
|
||||||
|
},
|
||||||
});
|
});
|
||||||
|
|
||||||
// HuggingFace model card query - only fetches when modal is open and model has a repo ID
|
// HuggingFace model card query - only fetches when modal is open and model has a repo ID
|
||||||
@@ -133,6 +146,19 @@ export function ModelManagement() {
|
|||||||
|
|
||||||
const errorCount = erroredDownloads.size;
|
const errorCount = erroredDownloads.size;
|
||||||
|
|
||||||
|
// Build progress map from active tasks for inline display
|
||||||
|
const downloadProgressMap = useMemo(() => {
|
||||||
|
const map = new Map<string, ActiveDownloadTask>();
|
||||||
|
if (activeTasks?.downloads) {
|
||||||
|
for (const dl of activeTasks.downloads) {
|
||||||
|
if (dl.status === 'downloading') {
|
||||||
|
map.set(dl.model_name, dl);
|
||||||
|
}
|
||||||
|
}
|
||||||
|
}
|
||||||
|
return map;
|
||||||
|
}, [activeTasks]);
|
||||||
|
|
||||||
const handleDownloadComplete = useCallback(() => {
|
const handleDownloadComplete = useCallback(() => {
|
||||||
setDownloadingModel(null);
|
setDownloadingModel(null);
|
||||||
setDownloadingDisplayName(null);
|
setDownloadingDisplayName(null);
|
||||||
@@ -371,16 +397,29 @@ export function ModelManagement() {
|
|||||||
)}
|
)}
|
||||||
</div>
|
</div>
|
||||||
|
|
||||||
{/* Name + meta */}
|
{/* Name + inline progress */}
|
||||||
<div className="flex-1 min-w-0">
|
<div className="flex-1 min-w-0">
|
||||||
<span className="text-sm font-medium">{model.display_name}</span>
|
<span className="text-sm font-medium">{model.display_name}</span>
|
||||||
|
{isDownloading &&
|
||||||
|
(() => {
|
||||||
|
const dl = downloadProgressMap.get(model.model_name);
|
||||||
|
const pct = dl?.progress ?? 0;
|
||||||
|
const hasProgress = dl && dl.total && dl.total > 0;
|
||||||
|
return (
|
||||||
|
<div className="mt-1 space-y-0.5">
|
||||||
|
<Progress value={hasProgress ? pct : undefined} className="h-1" />
|
||||||
|
<div className="text-[10px] text-muted-foreground truncate">
|
||||||
|
{hasProgress
|
||||||
|
? `${formatBytes(dl.current ?? 0)} / ${formatBytes(dl.total!)} (${pct.toFixed(0)}%)`
|
||||||
|
: dl?.filename || 'Connecting...'}
|
||||||
|
</div>
|
||||||
|
</div>
|
||||||
|
);
|
||||||
|
})()}
|
||||||
</div>
|
</div>
|
||||||
|
|
||||||
{/* Right side info */}
|
{/* Right side info */}
|
||||||
<div className="shrink-0 flex items-center gap-2">
|
<div className="shrink-0 flex items-center gap-2">
|
||||||
{isDownloading && (
|
|
||||||
<span className="text-xs text-muted-foreground">Downloading...</span>
|
|
||||||
)}
|
|
||||||
{hasError && (
|
{hasError && (
|
||||||
<Badge variant="destructive" className="text-[10px] h-5">
|
<Badge variant="destructive" className="text-[10px] h-5">
|
||||||
Error
|
Error
|
||||||
@@ -510,12 +549,6 @@ export function ModelManagement() {
|
|||||||
Downloaded
|
Downloaded
|
||||||
</Badge>
|
</Badge>
|
||||||
)}
|
)}
|
||||||
{selectedState?.isDownloading && (
|
|
||||||
<Badge variant="outline" className="text-xs">
|
|
||||||
<Loader2 className="h-3 w-3 mr-1 animate-spin" />
|
|
||||||
Downloading
|
|
||||||
</Badge>
|
|
||||||
)}
|
|
||||||
{selectedState?.hasError && (
|
{selectedState?.hasError && (
|
||||||
<Badge variant="destructive" className="text-xs">
|
<Badge variant="destructive" className="text-xs">
|
||||||
<CircleX className="h-3 w-3 mr-1" />
|
<CircleX className="h-3 w-3 mr-1" />
|
||||||
@@ -633,10 +666,25 @@ export function ModelManagement() {
|
|||||||
</>
|
</>
|
||||||
) : selectedState?.isDownloading ? (
|
) : selectedState?.isDownloading ? (
|
||||||
<>
|
<>
|
||||||
<Button size="sm" variant="outline" disabled className="flex-1">
|
<div className="flex-1 space-y-2">
|
||||||
<Loader2 className="h-4 w-4 mr-2 animate-spin" />
|
{(() => {
|
||||||
Downloading...
|
const dl = freshSelectedModel
|
||||||
</Button>
|
? downloadProgressMap.get(freshSelectedModel.model_name)
|
||||||
|
: undefined;
|
||||||
|
const pct = dl?.progress ?? 0;
|
||||||
|
const hasProgress = dl && dl.total && dl.total > 0;
|
||||||
|
return (
|
||||||
|
<>
|
||||||
|
<Progress value={hasProgress ? pct : undefined} className="h-2" />
|
||||||
|
<div className="text-xs text-muted-foreground">
|
||||||
|
{hasProgress
|
||||||
|
? `${formatBytes(dl.current ?? 0)} / ${formatBytes(dl.total!)} (${pct.toFixed(1)}%)`
|
||||||
|
: dl?.filename || 'Connecting to HuggingFace...'}
|
||||||
|
</div>
|
||||||
|
</>
|
||||||
|
);
|
||||||
|
})()}
|
||||||
|
</div>
|
||||||
<Button
|
<Button
|
||||||
size="sm"
|
size="sm"
|
||||||
onClick={() => handleCancel(freshSelectedModel.model_name)}
|
onClick={() => handleCancel(freshSelectedModel.model_name)}
|
||||||
|
|||||||
@@ -155,6 +155,10 @@ export interface ActiveDownloadTask {
|
|||||||
status: string;
|
status: string;
|
||||||
started_at: string;
|
started_at: string;
|
||||||
error?: string;
|
error?: string;
|
||||||
|
progress?: number; // 0-100 percentage
|
||||||
|
current?: number; // bytes downloaded
|
||||||
|
total?: number; // total bytes
|
||||||
|
filename?: string; // current file being downloaded
|
||||||
}
|
}
|
||||||
|
|
||||||
export interface ActiveGenerationTask {
|
export interface ActiveGenerationTask {
|
||||||
|
|||||||
@@ -103,19 +103,27 @@ class ChatterboxTTSBackend:
|
|||||||
|
|
||||||
def _load_model_sync(self):
|
def _load_model_sync(self):
|
||||||
"""Synchronous model loading."""
|
"""Synchronous model loading."""
|
||||||
|
from ..utils.hf_progress import HFProgressTracker, create_hf_progress_callback
|
||||||
|
|
||||||
progress_manager = get_progress_manager()
|
progress_manager = get_progress_manager()
|
||||||
task_manager = get_task_manager()
|
task_manager = get_task_manager()
|
||||||
model_name = "chatterbox-tts"
|
model_name = "chatterbox-tts"
|
||||||
|
|
||||||
is_cached = self._is_model_cached()
|
is_cached = self._is_model_cached()
|
||||||
|
|
||||||
|
# Set up HF progress tracking (intercepts tqdm for file-level progress)
|
||||||
|
progress_callback = create_hf_progress_callback(model_name, progress_manager)
|
||||||
|
tracker = HFProgressTracker(progress_callback, filter_non_downloads=is_cached)
|
||||||
|
tracker_context = tracker.patch_download()
|
||||||
|
tracker_context.__enter__()
|
||||||
|
|
||||||
if not is_cached:
|
if not is_cached:
|
||||||
task_manager.start_download(model_name)
|
task_manager.start_download(model_name)
|
||||||
progress_manager.update_progress(
|
progress_manager.update_progress(
|
||||||
model_name=model_name,
|
model_name=model_name,
|
||||||
current=0,
|
current=0,
|
||||||
total=0,
|
total=0,
|
||||||
filename="Downloading Chatterbox model...",
|
filename="Connecting to HuggingFace...",
|
||||||
status="downloading",
|
status="downloading",
|
||||||
)
|
)
|
||||||
|
|
||||||
@@ -131,25 +139,28 @@ class ChatterboxTTSBackend:
|
|||||||
# Monkey-patch torch.load for CPU loading. The model's .pt files
|
# Monkey-patch torch.load for CPU loading. The model's .pt files
|
||||||
# were saved on CUDA; from_pretrained() doesn't pass map_location
|
# were saved on CUDA; from_pretrained() doesn't pass map_location
|
||||||
# so loading on CPU fails without this.
|
# so loading on CPU fails without this.
|
||||||
if device == "cpu":
|
try:
|
||||||
_orig_torch_load = torch.load
|
if device == "cpu":
|
||||||
|
_orig_torch_load = torch.load
|
||||||
|
|
||||||
def _patched_load(*args, **kwargs):
|
def _patched_load(*args, **kwargs):
|
||||||
kwargs.setdefault("map_location", "cpu")
|
kwargs.setdefault("map_location", "cpu")
|
||||||
return _orig_torch_load(*args, **kwargs)
|
return _orig_torch_load(*args, **kwargs)
|
||||||
|
|
||||||
with ChatterboxTTSBackend._load_lock:
|
with ChatterboxTTSBackend._load_lock:
|
||||||
torch.load = _patched_load
|
torch.load = _patched_load
|
||||||
try:
|
try:
|
||||||
self.model = ChatterboxMultilingualTTS.from_pretrained(
|
self.model = ChatterboxMultilingualTTS.from_pretrained(
|
||||||
device=device,
|
device=device,
|
||||||
)
|
)
|
||||||
finally:
|
finally:
|
||||||
torch.load = _orig_torch_load
|
torch.load = _orig_torch_load
|
||||||
else:
|
else:
|
||||||
self.model = ChatterboxMultilingualTTS.from_pretrained(
|
self.model = ChatterboxMultilingualTTS.from_pretrained(
|
||||||
device=device,
|
device=device,
|
||||||
)
|
)
|
||||||
|
finally:
|
||||||
|
tracker_context.__exit__(None, None, None)
|
||||||
|
|
||||||
# Fix: transformers >= 4.36 defaults LlamaModel to sdpa attention
|
# Fix: transformers >= 4.36 defaults LlamaModel to sdpa attention
|
||||||
# which doesn't support output_attentions=True (needed by
|
# which doesn't support output_attentions=True (needed by
|
||||||
|
|||||||
@@ -94,19 +94,27 @@ class LuxTTSBackend:
|
|||||||
|
|
||||||
def _load_model_sync(self):
|
def _load_model_sync(self):
|
||||||
"""Synchronous model loading."""
|
"""Synchronous model loading."""
|
||||||
|
from ..utils.hf_progress import HFProgressTracker, create_hf_progress_callback
|
||||||
|
|
||||||
progress_manager = get_progress_manager()
|
progress_manager = get_progress_manager()
|
||||||
task_manager = get_task_manager()
|
task_manager = get_task_manager()
|
||||||
model_name = "luxtts"
|
model_name = "luxtts"
|
||||||
|
|
||||||
is_cached = self._is_model_cached()
|
is_cached = self._is_model_cached()
|
||||||
|
|
||||||
|
# Set up HF progress tracking (intercepts tqdm for file-level progress)
|
||||||
|
progress_callback = create_hf_progress_callback(model_name, progress_manager)
|
||||||
|
tracker = HFProgressTracker(progress_callback, filter_non_downloads=is_cached)
|
||||||
|
tracker_context = tracker.patch_download()
|
||||||
|
tracker_context.__enter__()
|
||||||
|
|
||||||
if not is_cached:
|
if not is_cached:
|
||||||
task_manager.start_download(model_name)
|
task_manager.start_download(model_name)
|
||||||
progress_manager.update_progress(
|
progress_manager.update_progress(
|
||||||
model_name=model_name,
|
model_name=model_name,
|
||||||
current=0,
|
current=0,
|
||||||
total=0,
|
total=0,
|
||||||
filename="Downloading LuxTTS model...",
|
filename="Connecting to HuggingFace...",
|
||||||
status="downloading",
|
status="downloading",
|
||||||
)
|
)
|
||||||
|
|
||||||
@@ -117,19 +125,22 @@ class LuxTTSBackend:
|
|||||||
logger.info(f"Loading LuxTTS on {device}...")
|
logger.info(f"Loading LuxTTS on {device}...")
|
||||||
|
|
||||||
# LuxTTS constructor downloads model and loads everything
|
# LuxTTS constructor downloads model and loads everything
|
||||||
if device == "cpu":
|
try:
|
||||||
import os
|
if device == "cpu":
|
||||||
threads = os.cpu_count() or 4
|
import os
|
||||||
self.model = LuxTTS(
|
threads = os.cpu_count() or 4
|
||||||
model_path=LUXTTS_HF_REPO,
|
self.model = LuxTTS(
|
||||||
device="cpu",
|
model_path=LUXTTS_HF_REPO,
|
||||||
threads=min(threads, 8),
|
device="cpu",
|
||||||
)
|
threads=min(threads, 8),
|
||||||
else:
|
)
|
||||||
self.model = LuxTTS(
|
else:
|
||||||
model_path=LUXTTS_HF_REPO,
|
self.model = LuxTTS(
|
||||||
device=device,
|
model_path=LUXTTS_HF_REPO,
|
||||||
)
|
device=device,
|
||||||
|
)
|
||||||
|
finally:
|
||||||
|
tracker_context.__exit__(None, None, None)
|
||||||
|
|
||||||
if not is_cached:
|
if not is_cached:
|
||||||
progress_manager.mark_complete(model_name)
|
progress_manager.mark_complete(model_name)
|
||||||
|
|||||||
@@ -1912,11 +1912,22 @@ async def get_active_tasks():
|
|||||||
pm_data = progress_manager._progress.get(model_name)
|
pm_data = progress_manager._progress.get(model_name)
|
||||||
if pm_data:
|
if pm_data:
|
||||||
error = pm_data.get("error")
|
error = pm_data.get("error")
|
||||||
|
# Include progress data if available
|
||||||
|
prog = progress or {}
|
||||||
|
if not prog:
|
||||||
|
with progress_manager._lock:
|
||||||
|
pm_data = progress_manager._progress.get(model_name)
|
||||||
|
if pm_data:
|
||||||
|
prog = pm_data
|
||||||
active_downloads.append(models.ActiveDownloadTask(
|
active_downloads.append(models.ActiveDownloadTask(
|
||||||
model_name=model_name,
|
model_name=model_name,
|
||||||
status=task.status,
|
status=task.status,
|
||||||
started_at=task.started_at,
|
started_at=task.started_at,
|
||||||
error=error,
|
error=error,
|
||||||
|
progress=prog.get("progress"),
|
||||||
|
current=prog.get("current"),
|
||||||
|
total=prog.get("total"),
|
||||||
|
filename=prog.get("filename"),
|
||||||
))
|
))
|
||||||
elif progress:
|
elif progress:
|
||||||
# Progress exists but no task - create from progress data
|
# Progress exists but no task - create from progress data
|
||||||
@@ -1934,6 +1945,10 @@ async def get_active_tasks():
|
|||||||
status=progress.get("status", "downloading"),
|
status=progress.get("status", "downloading"),
|
||||||
started_at=started_at,
|
started_at=started_at,
|
||||||
error=progress.get("error"),
|
error=progress.get("error"),
|
||||||
|
progress=progress.get("progress"),
|
||||||
|
current=progress.get("current"),
|
||||||
|
total=progress.get("total"),
|
||||||
|
filename=progress.get("filename"),
|
||||||
))
|
))
|
||||||
|
|
||||||
# Get active generations
|
# Get active generations
|
||||||
|
|||||||
@@ -158,6 +158,10 @@ class ActiveDownloadTask(BaseModel):
|
|||||||
status: str
|
status: str
|
||||||
started_at: datetime
|
started_at: datetime
|
||||||
error: Optional[str] = None
|
error: Optional[str] = None
|
||||||
|
progress: Optional[float] = None # 0-100 percentage
|
||||||
|
current: Optional[int] = None # bytes downloaded
|
||||||
|
total: Optional[int] = None # total bytes
|
||||||
|
filename: Optional[str] = None # current file being downloaded
|
||||||
|
|
||||||
|
|
||||||
class ActiveGenerationTask(BaseModel):
|
class ActiveGenerationTask(BaseModel):
|
||||||
|
|||||||
Reference in New Issue
Block a user