diff --git a/core/cli/models.go b/core/cli/models.go index e2a832b0c877..c6961c402ecd 100644 --- a/core/cli/models.go +++ b/core/cli/models.go @@ -28,6 +28,8 @@ type ModelsCMDFlags struct { Color string `env:"COLOR" hidden:""` NoColor string `env:"NO_COLOR" hidden:""` HFToken string `env:"HF_TOKEN" hidden:""` + + ArtifactDownloadConcurrency int `env:"LOCALAI_ARTIFACT_DOWNLOAD_CONCURRENCY" help:"How many files of a model artifact to download at once. 1 (the default) downloads sequentially. Raising it helps artifacts split into many files on a fast link, at the cost of more concurrent load on the models volume" group:"storage" default:"1"` } type ModelsList struct { @@ -87,6 +89,7 @@ func (mi *ModelsInstall) Run(ctx *cliContext.Context) error { artifactMaterializer := modelartifacts.NewDefaultManager( modelartifacts.WithHuggingFaceToken(mi.HFToken), + modelartifacts.WithDownloadConcurrency(mi.ArtifactDownloadConcurrency), ) galleryService := galleryop.NewGalleryService(&config.ApplicationConfig{ SystemState: systemState, diff --git a/core/cli/run.go b/core/cli/run.go index 749dd16f0dfc..c70693186d0b 100644 --- a/core/cli/run.go +++ b/core/cli/run.go @@ -41,6 +41,7 @@ type RunCMD struct { BackendsPath string `env:"LOCALAI_BACKENDS_PATH,BACKENDS_PATH" type:"path" default:"${basepath}/backends" help:"Path containing backends used for inferencing" group:"backends"` BackendsSystemPath string `env:"LOCALAI_BACKENDS_SYSTEM_PATH,BACKEND_SYSTEM_PATH" type:"path" default:"/var/lib/local-ai/backends" help:"Path containing system backends used for inferencing" group:"backends"` ModelsPath string `env:"LOCALAI_MODELS_PATH,MODELS_PATH" type:"path" default:"${basepath}/models" help:"Path containing models used for inferencing" group:"storage"` + ArtifactDownloadConcurrency int `env:"LOCALAI_ARTIFACT_DOWNLOAD_CONCURRENCY" help:"How many files of a model artifact to download at once. 1 (the default) downloads sequentially. Raising it helps artifacts split into many files on a fast link, at the cost of more concurrent load on the models volume" group:"storage" default:"1"` GeneratedContentPath string `env:"LOCALAI_GENERATED_CONTENT_PATH,GENERATED_CONTENT_PATH" type:"path" default:"${generatedcontentpath}" help:"Location for generated content (e.g. images, audio, videos)" group:"storage"` UploadPath string `env:"LOCALAI_UPLOAD_PATH,UPLOAD_PATH" type:"path" default:"${uploadpath}" help:"Path to store uploads from files api" group:"storage"` DataPath string `env:"LOCALAI_DATA_PATH" type:"path" default:"${basepath}/data" help:"Path for persistent data (collectiondb, agent state, tasks, jobs). Separates mutable data from configuration" group:"storage"` @@ -278,8 +279,10 @@ func (r *RunCMD) Run(ctx *cliContext.Context) error { opts := []config.AppOption{ config.WithContext(context.Background()), + config.WithArtifactDownloadConcurrency(r.ArtifactDownloadConcurrency), config.WithModelArtifactMaterializer(modelartifacts.NewDefaultManager( modelartifacts.WithHuggingFaceToken(r.HFToken), + modelartifacts.WithDownloadConcurrency(r.ArtifactDownloadConcurrency), )), config.WithModelPreloadDisplay(r.Color, r.NoColor != ""), config.WithConfigFile(r.ModelsConfigFile), diff --git a/core/config/application_artifact_materializer_test.go b/core/config/application_artifact_materializer_test.go index 13a9d856b8da..4e0eb927fe60 100644 --- a/core/config/application_artifact_materializer_test.go +++ b/core/config/application_artifact_materializer_test.go @@ -16,6 +16,15 @@ func (*applicationArtifactMaterializer) Ensure(context.Context, string, modelart return modelartifacts.Result{}, nil } +type configurableApplicationArtifactMaterializer struct { + applicationArtifactMaterializer + concurrency int +} + +func (m *configurableApplicationArtifactMaterializer) SetDownloadConcurrency(concurrency int) { + m.concurrency = concurrency +} + var _ = Describe("ApplicationConfig model artifact materializer", func() { It("provides a default materializer", func() { Expect(NewApplicationConfig().ModelArtifactMaterializer).NotTo(BeNil()) @@ -31,4 +40,15 @@ var _ = Describe("ApplicationConfig model artifact materializer", func() { Expect(field.Tag.Get("json")).To(Equal("-")) Expect(field.Tag.Get("yaml")).To(Equal("-")) }) + + It("applies runtime download concurrency to configurable materializers", func() { + materializer := &configurableApplicationArtifactMaterializer{} + appConfig := NewApplicationConfig(WithModelArtifactMaterializer(materializer)) + concurrency := 4 + + appConfig.ApplyRuntimeSettings(&RuntimeSettings{ArtifactDownloadConcurrency: &concurrency}) + + Expect(appConfig.ArtifactDownloadConcurrency).To(Equal(4)) + Expect(materializer.concurrency).To(Equal(4)) + }) }) diff --git a/core/config/application_config.go b/core/config/application_config.go index 7ea81b2b4ed7..b7ec9a316f36 100644 --- a/core/config/application_config.go +++ b/core/config/application_config.go @@ -35,6 +35,7 @@ type ApplicationConfig struct { // network interfaces (e.g. eth0), filtering out docker0/veth noise. WebRTCICEInterfaces []string UploadLimitMB, Threads, ContextSize int + ArtifactDownloadConcurrency int F16 bool Debug bool EnableTracing bool @@ -58,12 +59,12 @@ type ApplicationConfig struct { // gzip is skipped. 0 keeps middleware.DefaultCompressionMinLength. HTTPCompressionMinLength int PreloadJSONModels string - PreloadModelsFromPath string - CORSAllowOrigins string - ApiKeys []string - P2PToken string - P2PNetworkID string - Federated bool + PreloadModelsFromPath string + CORSAllowOrigins string + ApiKeys []string + P2PToken string + P2PNetworkID string + Federated bool // ExternalBaseURL is the externally visible base URL of this instance // (scheme+host[:port]), set via LOCALAI_BASE_URL. When non-empty it is @@ -276,11 +277,12 @@ func NewApplicationConfig(o ...AppOption) *ApplicationConfig { // force-enables it). It's a small in-memory ring buffer; the Settings // toggle can still turn it off (a persisted false wins - see // loadRuntimeSettingsFromFile). - EnableBackendLogging: true, - AgentJobRetentionDays: 30, // Default: 30 days - LRUEvictionMaxRetries: 30, // Default: 30 retries - LRUEvictionRetryInterval: 1 * time.Second, // Default: 1 second - ModelLoadFailureCooldown: 10 * time.Second, // Default: 10s base cooldown after a failed load + EnableBackendLogging: true, + ArtifactDownloadConcurrency: modelartifacts.DefaultDownloadConcurrency, + AgentJobRetentionDays: 30, // Default: 30 days + LRUEvictionMaxRetries: 30, // Default: 30 retries + LRUEvictionRetryInterval: 1 * time.Second, // Default: 1 second + ModelLoadFailureCooldown: 10 * time.Second, // Default: 10s base cooldown after a failed load // WatchDogInterval is intentionally left at the zero value here. // The startup loader applies a persisted runtime_settings.json value // only when the interval is still 0 (its "not set by env var" @@ -685,6 +687,15 @@ func WithModelArtifactMaterializer(materializer ArtifactMaterializer) AppOption } } +func WithArtifactDownloadConcurrency(concurrency int) AppOption { + return func(o *ApplicationConfig) { + if concurrency < 1 { + concurrency = modelartifacts.DefaultDownloadConcurrency + } + o.ArtifactDownloadConcurrency = concurrency + } +} + // WithModelPreloadDisplay configures terminal rendering for model preload output. func WithModelPreloadDisplay(renderMode string, disableColor bool) AppOption { return func(o *ApplicationConfig) { @@ -1190,6 +1201,11 @@ func (o *ApplicationConfig) ApplyRuntimeSettings(settings *RuntimeSettings) (req xsysinfo.SetDefaultVRAMBudget(b) } } + if settings.ArtifactDownloadConcurrency != nil { + if configurable, ok := o.ModelArtifactMaterializer.(interface{ SetDownloadConcurrency(int) }); ok { + configurable.SetDownloadConcurrency(o.ArtifactDownloadConcurrency) + } + } // Note: ApiKeys need env-merge handling (MergeAPIKeys) - done by the // caller, because the env-provided keys live on the startup config. return requireRestart diff --git a/core/config/runtime_settings.go b/core/config/runtime_settings.go index 7c29e63375d9..6e4381d8c97d 100644 --- a/core/config/runtime_settings.go +++ b/core/config/runtime_settings.go @@ -33,15 +33,16 @@ type RuntimeSettings struct { LRUEvictionRetryInterval *string `json:"lru_eviction_retry_interval,omitempty"` // Interval between retries when waiting for busy models (e.g., 1s, 2s) (default: 1s) // Performance settings - Threads *int `json:"threads,omitempty"` - ContextSize *int `json:"context_size,omitempty"` - VRAMBudget *string `json:"vram_budget,omitempty"` // Cap VRAM for allocation ("80%" or "12GB"; "" = no cap) - F16 *bool `json:"f16,omitempty"` - Debug *bool `json:"debug,omitempty"` - EnableTracing *bool `json:"enable_tracing,omitempty"` - TracingMaxItems *int `json:"tracing_max_items,omitempty"` - TracingMaxBodyBytes *int `json:"tracing_max_body_bytes,omitempty"` // Per-body cap in bytes; 0 disables the cap - EnableBackendLogging *bool `json:"enable_backend_logging,omitempty"` + Threads *int `json:"threads,omitempty"` + ContextSize *int `json:"context_size,omitempty"` + ArtifactDownloadConcurrency *int `json:"artifact_download_concurrency,omitempty"` + VRAMBudget *string `json:"vram_budget,omitempty"` // Cap VRAM for allocation ("80%" or "12GB"; "" = no cap) + F16 *bool `json:"f16,omitempty"` + Debug *bool `json:"debug,omitempty"` + EnableTracing *bool `json:"enable_tracing,omitempty"` + TracingMaxItems *int `json:"tracing_max_items,omitempty"` + TracingMaxBodyBytes *int `json:"tracing_max_body_bytes,omitempty"` // Per-body cap in bytes; 0 disables the cap + EnableBackendLogging *bool `json:"enable_backend_logging,omitempty"` // Security/CORS settings CORS *bool `json:"cors,omitempty"` diff --git a/core/config/runtime_settings_registry.go b/core/config/runtime_settings_registry.go index f06e14c91d52..1678c08b0543 100644 --- a/core/config/runtime_settings_registry.go +++ b/core/config/runtime_settings_registry.go @@ -227,6 +227,15 @@ var runtimeSettingsFields = []fieldSpec{ func(s *RuntimeSettings) **int { return &s.ContextSize }, func(o *ApplicationConfig) int { return o.ContextSize }, func(o *ApplicationConfig, v int) { o.ContextSize = v }), + field("artifact_download_concurrency", + func(s *RuntimeSettings) **int { return &s.ArtifactDownloadConcurrency }, + func(o *ApplicationConfig) int { return o.ArtifactDownloadConcurrency }, + func(o *ApplicationConfig, v int) { + if v < 1 { + v = 1 + } + o.ArtifactDownloadConcurrency = v + }), // VRAM budget: the cap string ("80%"/"12GB"/"" = uncapped). The live // side effect (xsysinfo.SetDefaultVRAMBudget) is post-processing in the // apply loop, not here - the row only owns the config member, matching diff --git a/core/config/runtime_settings_registry_internal_test.go b/core/config/runtime_settings_registry_internal_test.go index 7124eb85bf7b..ef0af7781cef 100644 --- a/core/config/runtime_settings_registry_internal_test.go +++ b/core/config/runtime_settings_registry_internal_test.go @@ -71,6 +71,7 @@ var _ = Describe("runtime settings registry", func() { src.LRUEvictionRetryInterval = 3 * time.Second src.Threads = 7 src.ContextSize = 8192 + src.ArtifactDownloadConcurrency = 6 src.VRAMBudget = "12GiB" src.F16 = true src.Debug = true diff --git a/core/config/runtime_settings_startup.go b/core/config/runtime_settings_startup.go index 9eda102f61c6..f9da8c6c3b24 100644 --- a/core/config/runtime_settings_startup.go +++ b/core/config/runtime_settings_startup.go @@ -93,4 +93,9 @@ func (o *ApplicationConfig) ApplyRuntimeSettingsAtStartup(settings *RuntimeSetti xsysinfo.SetDefaultVRAMBudget(b) } } + if settings.ArtifactDownloadConcurrency != nil { + if configurable, ok := o.ModelArtifactMaterializer.(interface{ SetDownloadConcurrency(int) }); ok { + configurable.SetDownloadConcurrency(o.ArtifactDownloadConcurrency) + } + } } diff --git a/core/http/react-ui/e2e/settings-backend-logging.spec.js b/core/http/react-ui/e2e/settings-backend-logging.spec.js index 4b03ba143f88..7b5459d8973e 100644 --- a/core/http/react-ui/e2e/settings-backend-logging.spec.js +++ b/core/http/react-ui/e2e/settings-backend-logging.spec.js @@ -11,6 +11,13 @@ test.describe('Settings - Backend Logging', () => { await expect(page.locator('text=Enable Backend Logging')).toBeVisible() }) + test('artifact download concurrency is configurable', async ({ page }) => { + const input = page.getByLabel('Artifact Download Concurrency') + await expect(input).toBeVisible() + await input.fill('4') + await expect(input).toHaveValue('4') + }) + test('backend logging toggle can be toggled', async ({ page }) => { // Find the checkbox associated with backend logging const section = page.locator('div', { has: page.locator('text=Enable Backend Logging') }) diff --git a/core/http/react-ui/src/pages/Settings.jsx b/core/http/react-ui/src/pages/Settings.jsx index 5299454f6492..d6a813e35bdc 100644 --- a/core/http/react-ui/src/pages/Settings.jsx +++ b/core/http/react-ui/src/pages/Settings.jsx @@ -420,6 +420,9 @@ export default function Settings() { update('context_size', parseInt(e.target.value) || 0)} placeholder="2048" /> + + update('artifact_download_concurrency', Math.max(1, parseInt(e.target.value) || 1))} /> + update('vram_budget', e.target.value)} placeholder="e.g. 80% or 12GB" /> diff --git a/docs/content/features/runtime-settings.md b/docs/content/features/runtime-settings.md index e310bb06d590..4c70f43e2e4e 100644 --- a/docs/content/features/runtime-settings.md +++ b/docs/content/features/runtime-settings.md @@ -49,6 +49,7 @@ You can configure these settings via the web UI or through environment variables - **Threads**: Number of threads used for parallel computation (recommended: number of physical cores) - **Context Size**: Default context size for models (default: `512`) +- **Artifact Download Concurrency**: Maximum number of artifact files downloaded at once. `1` downloads sequentially (default: `1`) - **F16**: Enable GPU acceleration using 16-bit floating point - **VRAM Budget**: Cap on VRAM used for model allocation (for example `80%` or `12GB`; empty means no cap). See [VRAM Management]({{%relref "advanced/vram-management" %}}) @@ -138,6 +139,7 @@ The `runtime_settings.json` file follows this structure: "lru_eviction_retry_interval": "1s", "threads": 8, "context_size": 2048, + "artifact_download_concurrency": 4, "f16": false, "debug": false, "cors": true, @@ -221,4 +223,3 @@ If P2P is not starting: 2. Check network connectivity 3. Ensure the P2P network ID matches across nodes (if using federated mode) 4. Review logs for P2P-related errors - diff --git a/docs/content/reference/cli-reference.md b/docs/content/reference/cli-reference.md index d0b82ed6c3f9..5c47b6e36e49 100644 --- a/docs/content/reference/cli-reference.md +++ b/docs/content/reference/cli-reference.md @@ -28,6 +28,7 @@ Complete reference for all LocalAI command-line interface (CLI) parameters and e | `--localai-config-dir` | `BASEPATH/configuration` | Directory for dynamic loading of certain configuration files (currently runtime_settings.json, api_keys.json, and external_backends.json). See [Runtime Settings]({{%relref "features/runtime-settings" %}}) for web-based configuration. | `$LOCALAI_CONFIG_DIR` | | `--localai-config-dir-poll-interval` | | Time duration to poll the LocalAI Config Dir if your system has broken fsnotify events (example: `1m`) | `$LOCALAI_CONFIG_DIR_POLL_INTERVAL` | | `--models-config-file` | | YAML file containing a list of model backend configs (alias: `--config-file`) | `$LOCALAI_MODELS_CONFIG_FILE`, `$CONFIG_FILE` | +| `--artifact-download-concurrency` | `1` | How many files of a model artifact to download at once. `1` downloads sequentially. Raising it helps artifacts split into many files on a fast link, at the cost of more concurrent load on the models volume. Whole files only — a single file is never split, so resume and per-file checksum verification are unaffected | `$LOCALAI_ARTIFACT_DOWNLOAD_CONCURRENCY` | ## Backend Flags diff --git a/pkg/downloader/download_plan.go b/pkg/downloader/download_plan.go index bc007ed005c5..4f091383911c 100644 --- a/pkg/downloader/download_plan.go +++ b/pkg/downloader/download_plan.go @@ -2,8 +2,10 @@ package downloader import ( "context" + "sync" "github.com/mudler/xlog" + "golang.org/x/sync/errgroup" ) // FileTask describes one download operation and an optional post-download @@ -23,23 +25,72 @@ type FileTask struct { // The helper centralizes the shared download path so callers only provide // source/destination metadata and any post-download hook they need. func DownloadFilesWithContext(ctx context.Context, tasks []FileTask, status func(string, string, string, float64), opts ...DownloadOption) error { + return DownloadFilesWithConcurrency(ctx, tasks, status, 1, opts...) +} + +// DownloadFilesWithConcurrency runs up to concurrency downloads at once. A +// concurrency of one or less keeps the original sequential path, so callers that +// have not opted in are byte-for-byte unaffected: tasks still run in slice order +// and the first failure still returns before any later task starts. +// +// Only whole files run in parallel. A single file is never split, so the +// .partial resume machinery and the per-file SHA check in downloadTaskWithRetry +// keep working untouched. +// +// The status callback is serialized, because it belongs to the caller and the +// sequential path gave it an implicit guarantee of never being entered twice at +// once. AfterDownload is deliberately *not* serialized: it does the per-file +// verify-and-promote work that parallelism is meant to overlap, so hooks must be +// safe to run concurrently with each other. +func DownloadFilesWithConcurrency(ctx context.Context, tasks []FileTask, status func(string, string, string, float64), concurrency int, opts ...DownloadOption) error { + if concurrency < 1 { + concurrency = 1 + } + + if status != nil && concurrency > 1 { + var statusMutex sync.Mutex + unsynchronized := status + status = func(fileName, current, total string, percent float64) { + statusMutex.Lock() + defer statusMutex.Unlock() + unsynchronized(fileName, current, total, percent) + } + } + + // errgroup.WithContext cancels the derived context on the first error, which + // is what stops in-flight transfers instead of letting them run to + // completion, and Wait reports that first error rather than the + // context.Canceled the siblings observe. + group, groupCtx := errgroup.WithContext(ctx) + group.SetLimit(concurrency) + for i := range tasks { task := tasks[i] - if err := ctx.Err(); err != nil { - return err + if err := groupCtx.Err(); err != nil { + break } - taskOpts := append([]DownloadOption{}, opts...) - taskOpts = append(taskOpts, task.Options...) - if err := downloadTaskWithRetry(ctx, task, status, taskOpts); err != nil { - return err - } - if task.AfterDownload != nil { - if err := task.AfterDownload(task.Destination); err != nil { + group.Go(func() error { + if err := groupCtx.Err(); err != nil { return err } - } + taskOpts := append([]DownloadOption{}, opts...) + taskOpts = append(taskOpts, task.Options...) + if err := downloadTaskWithRetry(groupCtx, task, status, taskOpts); err != nil { + return err + } + if task.AfterDownload != nil { + return task.AfterDownload(task.Destination) + } + return nil + }) + } + + if err := group.Wait(); err != nil { + return err } - return nil + // A caller-cancelled context with no task in flight leaves the group clean, + // so report the cancellation the sequential loop would have reported. + return ctx.Err() } // downloadTaskWithRetry fetches one file, retrying transient failures. Without diff --git a/pkg/downloader/plan_concurrency_test.go b/pkg/downloader/plan_concurrency_test.go new file mode 100644 index 000000000000..ad7d35d06c45 --- /dev/null +++ b/pkg/downloader/plan_concurrency_test.go @@ -0,0 +1,151 @@ +package downloader_test + +import ( + "context" + "fmt" + "net/http" + "net/http/httptest" + "path/filepath" + "sync/atomic" + "time" + + . "github.com/onsi/ginkgo/v2" + . "github.com/onsi/gomega" + + "github.com/mudler/LocalAI/pkg/downloader" +) + +var _ = Describe("DownloadFilesWithConcurrency", func() { + // slowServer holds every request open until it has seen `hold` of them at + // once, or the client gives up. A sequential executor can never satisfy a + // hold above one, so this doubles as proof that parallelism really happens + // rather than just being configured. + slowServer := func(delay time.Duration) (*httptest.Server, *int32) { + var inFlight int32 + var peak int32 + server := httptest.NewServer(http.HandlerFunc(func(w http.ResponseWriter, r *http.Request) { + current := atomic.AddInt32(&inFlight, 1) + for { + observed := atomic.LoadInt32(&peak) + if current <= observed || atomic.CompareAndSwapInt32(&peak, observed, current) { + break + } + } + time.Sleep(delay) + atomic.AddInt32(&inFlight, -1) + _, _ = w.Write([]byte("payload")) + })) + return server, &peak + } + + tasksFor := func(server *httptest.Server, dir string, count int) []downloader.FileTask { + tasks := make([]downloader.FileTask, 0, count) + for i := 0; i < count; i++ { + tasks = append(tasks, downloader.FileTask{ + URI: downloader.URI(fmt.Sprintf("%s/file-%d", server.URL, i)), + Destination: filepath.Join(dir, fmt.Sprintf("file-%d.bin", i)), + FileIndex: i, + TotalFiles: count, + }) + } + return tasks + } + + It("overlaps transfers up to the limit and no further", func() { + server, peak := slowServer(60 * time.Millisecond) + DeferCleanup(server.Close) + + tasks := tasksFor(server, GinkgoT().TempDir(), 8) + err := downloader.DownloadFilesWithConcurrency(context.Background(), tasks, nil, 3) + + Expect(err).NotTo(HaveOccurred()) + Expect(*peak).To(BeNumerically(">", 1), "downloads never overlapped, so the limit was not applied") + Expect(*peak).To(BeNumerically("<=", 3), "more transfers ran at once than the configured limit") + }) + + It("keeps a concurrency of one strictly sequential", func() { + server, peak := slowServer(10 * time.Millisecond) + DeferCleanup(server.Close) + + tasks := tasksFor(server, GinkgoT().TempDir(), 5) + err := downloader.DownloadFilesWithConcurrency(context.Background(), tasks, nil, 1) + + Expect(err).NotTo(HaveOccurred()) + Expect(*peak).To(Equal(int32(1)), "a limit of one must never overlap transfers") + }) + + It("treats a non-positive concurrency as sequential", func() { + server, peak := slowServer(10 * time.Millisecond) + DeferCleanup(server.Close) + + tasks := tasksFor(server, GinkgoT().TempDir(), 4) + err := downloader.DownloadFilesWithConcurrency(context.Background(), tasks, nil, 0) + + Expect(err).NotTo(HaveOccurred()) + Expect(*peak).To(Equal(int32(1))) + }) + + It("reports the first hook error and stops starting new work", func() { + server, _ := slowServer(0) + DeferCleanup(server.Close) + + var started int32 + tasks := tasksFor(server, GinkgoT().TempDir(), 24) + for i := range tasks { + index := i + tasks[i].AfterDownload = func(string) error { + atomic.AddInt32(&started, 1) + if index == 0 { + return fmt.Errorf("verification failed for shard %d", index) + } + time.Sleep(20 * time.Millisecond) + return nil + } + } + + err := downloader.DownloadFilesWithConcurrency(context.Background(), tasks, nil, 2) + + Expect(err).To(MatchError(ContainSubstring("verification failed for shard 0"))) + Expect(atomic.LoadInt32(&started)).To(BeNumerically("<", int32(len(tasks))), + "the executor kept starting work after a failure instead of cancelling") + }) + + It("returns the caller's cancellation rather than running the plan", func() { + server, _ := slowServer(0) + DeferCleanup(server.Close) + + ctx, cancel := context.WithCancel(context.Background()) + cancel() + + var ran int32 + tasks := tasksFor(server, GinkgoT().TempDir(), 3) + for i := range tasks { + tasks[i].AfterDownload = func(string) error { + atomic.AddInt32(&ran, 1) + return nil + } + } + + err := downloader.DownloadFilesWithConcurrency(ctx, tasks, nil, 4) + + Expect(err).To(MatchError(context.Canceled)) + Expect(atomic.LoadInt32(&ran)).To(BeZero()) + }) + + It("serializes the status callback so callers need no locking of their own", func() { + server, _ := slowServer(5 * time.Millisecond) + DeferCleanup(server.Close) + + // A deliberately unsynchronized counter: if the executor let two + // callbacks in at once, -race would flag this write. + unguarded := 0 + + tasks := tasksFor(server, GinkgoT().TempDir(), 6) + err := downloader.DownloadFilesWithConcurrency(context.Background(), tasks, func(string, string, string, float64) { + unguarded++ + }, 4) + + Expect(err).NotTo(HaveOccurred()) + Expect(unguarded).To(BeNumerically(">", 0)) + }) +}) diff --git a/pkg/modelartifacts/materializer.go b/pkg/modelartifacts/materializer.go index 8f1f031929fd..a3e2a9e8ece0 100644 --- a/pkg/modelartifacts/materializer.go +++ b/pkg/modelartifacts/materializer.go @@ -13,6 +13,7 @@ import ( "path" "path/filepath" "strings" + "sync/atomic" "syscall" "time" @@ -51,6 +52,12 @@ const ( // the backend download the same repo in-band. DefaultLockWait = 30 * time.Minute + // DefaultDownloadConcurrency keeps materialization sequential unless an + // operator opts in. Parallel transfers help a repo of many small shards on a + // fast link, but they multiply memory and disk pressure on the shared models + // volume, so the safe default is the behaviour this package already had. + DefaultDownloadConcurrency = 1 + initialLockRetryInterval = 100 * time.Millisecond maxLockRetryInterval = 5 * time.Second ) @@ -65,6 +72,9 @@ type Manager struct { // the process run that created it, and outliving that run is precisely what // it must not do. writerID string + // downloadConcurrency bounds how many of a snapshot's files transfer at + // once. One means the sequential behaviour this package shipped with. + downloadConcurrency atomic.Int64 } type ManagerOption func(*Manager) @@ -100,6 +110,25 @@ func WithLockWait(wait time.Duration) ManagerOption { } } +// WithDownloadConcurrency bounds how many of a snapshot's files are fetched at +// once. Values below one mean sequential, which is the default: a shared models +// volume is often the bottleneck rather than the network, so raising this is a +// deployment decision rather than something to assume. +func WithDownloadConcurrency(concurrency int) ManagerOption { + return func(manager *Manager) { + manager.SetDownloadConcurrency(concurrency) + } +} + +// SetDownloadConcurrency updates the limit used by future file download +// batches. Values below one select the safe sequential default. +func (m *Manager) SetDownloadConcurrency(concurrency int) { + if concurrency < 1 { + concurrency = DefaultDownloadConcurrency + } + m.downloadConcurrency.Store(int64(concurrency)) +} + func NewManager(resolver SnapshotResolver, options ...ManagerOption) *Manager { manager := &Manager{ resolver: resolver, @@ -107,6 +136,7 @@ func NewManager(resolver SnapshotResolver, options ...ManagerOption) *Manager { lockWait: DefaultLockWait, writerID: newWriterID(), } + manager.SetDownloadConcurrency(DefaultDownloadConcurrency) for _, option := range options { option(manager) } @@ -376,7 +406,11 @@ func (m *Manager) materializeLocked(ctx context.Context, modelsPath string, spec // read this manifest, and getting its order or contents wrong would make a // corrupt tree look valid. manifest := Manifest{Version: ManifestVersion, Artifact: spec, Files: make([]ManifestFile, len(snapshot.Files))} - completedBytes := int64(0) + // completedBytes is atomic because WithDownloadConcurrency lets several + // AfterDownload hooks add to it while other files' progress callbacks read + // it. Each hook still writes its own manifest.Files slot, so the manifest + // stays in snapshot order no matter which file finishes first. + completedBytes := new(atomic.Int64) skippedFiles := 0 skippedBytes := int64(0) tasks := make([]downloader.FileTask, 0, len(snapshot.Files)) @@ -398,7 +432,7 @@ func (m *Manager) materializeLocked(ctx context.Context, modelsPath string, spec snapshotAbs := filepath.Join(layout.Partial, filepath.FromSlash(snapshotRel)) if entry, ok := reuseMaterializedFile(snapshotAbs, file); ok { manifest.Files[taskIndex] = entry - completedBytes += file.Size + completedBytes.Add(file.Size) skippedFiles++ skippedBytes += file.Size continue @@ -419,7 +453,7 @@ func (m *Manager) materializeLocked(ctx context.Context, modelsPath string, spec Phase: PhaseDownloading, Artifact: spec.Name, File: file.Path, - CurrentBytes: completedBytes + event.Written, + CurrentBytes: completedBytes.Load() + event.Written, TotalBytes: totalBytes, CompletedFiles: taskIndex, TotalFiles: len(snapshot.Files), @@ -431,7 +465,7 @@ func (m *Manager) materializeLocked(ctx context.Context, modelsPath string, spec Phase: PhaseVerifying, Artifact: spec.Name, File: file.Path, - CurrentBytes: completedBytes + file.Size, + CurrentBytes: completedBytes.Load() + file.Size, TotalBytes: totalBytes, CompletedFiles: taskIndex, TotalFiles: len(snapshot.Files), @@ -454,7 +488,7 @@ func (m *Manager) materializeLocked(ctx context.Context, modelsPath string, spec return err } manifest.Files[taskIndex] = entry - completedBytes += file.Size + completedBytes.Add(file.Size) return nil }, } @@ -471,7 +505,7 @@ func (m *Manager) materializeLocked(ctx context.Context, modelsPath string, spec "remaining_files", len(tasks), "total_files", len(snapshot.Files)) } - if err := downloader.DownloadFilesWithContext(ctx, tasks, nil); err != nil { + if err := downloader.DownloadFilesWithConcurrency(ctx, tasks, nil, int(m.downloadConcurrency.Load())); err != nil { return Result{}, err } if err := root.RemoveAll(".downloads"); err != nil { diff --git a/pkg/modelartifacts/materializer_concurrency_test.go b/pkg/modelartifacts/materializer_concurrency_test.go new file mode 100644 index 000000000000..a89f67f4d3c3 --- /dev/null +++ b/pkg/modelartifacts/materializer_concurrency_test.go @@ -0,0 +1,180 @@ +package modelartifacts_test + +import ( + "context" + "crypto/sha256" + "encoding/hex" + "fmt" + "net/http" + "net/http/httptest" + "os" + "path/filepath" + "sync/atomic" + "time" + + . "github.com/onsi/ginkgo/v2" + . "github.com/onsi/gomega" + + hfapi "github.com/mudler/LocalAI/pkg/huggingface-api" + "github.com/mudler/LocalAI/pkg/modelartifacts" +) + +var _ = Describe("artifact materialization with bounded download concurrency", func() { + // shardedSnapshot serves `count` distinct files and reports the peak number + // of simultaneous requests, so a test can tell configured concurrency from + // actual concurrency. + shardedSnapshot := func(count int, delay time.Duration) (hfapi.Snapshot, *httptest.Server, *int32) { + bodies := make(map[string][]byte, count) + files := make([]hfapi.SnapshotFile, 0, count) + + var inFlight, peak int32 + server := httptest.NewServer(http.HandlerFunc(func(w http.ResponseWriter, r *http.Request) { + current := atomic.AddInt32(&inFlight, 1) + for { + observed := atomic.LoadInt32(&peak) + if current <= observed || atomic.CompareAndSwapInt32(&peak, observed, current) { + break + } + } + time.Sleep(delay) + atomic.AddInt32(&inFlight, -1) + _, _ = w.Write(bodies[r.URL.Path]) + })) + + for i := 0; i < count; i++ { + // Later shards are served first-come, so give them descending delays + // as well: completion order ends up unrelated to snapshot order, + // which is exactly what the manifest must survive. + body := []byte(fmt.Sprintf("shard-%02d-bytes", i)) + urlPath := fmt.Sprintf("/shard-%02d", i) + bodies[urlPath] = body + sum := sha256.Sum256(body) + files = append(files, hfapi.SnapshotFile{ + Path: fmt.Sprintf("shards/model-%02d.safetensors", i), + Size: int64(len(body)), + LFSOID: hex.EncodeToString(sum[:]), + URL: server.URL + urlPath, + }) + } + + return hfapi.Snapshot{ + Endpoint: "https://huggingface.co", Repo: "owner/sharded", + RequestedRevision: "main", ResolvedRevision: "0123456789abcdef0123456789abcdef01234567", + Files: files, + }, server, &peak + } + + spec := modelartifacts.Spec{Source: modelartifacts.Source{Type: "huggingface", Repo: "owner/sharded"}} + + It("records the manifest in snapshot order regardless of completion order", func() { + snapshot, server, peak := shardedSnapshot(12, 40*time.Millisecond) + DeferCleanup(server.Close) + + manager := modelartifacts.NewManager(&fakeSnapshotResolver{snapshot: snapshot}, + modelartifacts.WithDownloadConcurrency(4)) + modelsPath := GinkgoT().TempDir() + + result, err := manager.Ensure(context.Background(), modelsPath, spec) + Expect(err).NotTo(HaveOccurred()) + Expect(*peak).To(BeNumerically(">", 1), "files never overlapped, so this proves nothing about ordering") + Expect(*peak).To(BeNumerically("<=", 4)) + + Expect(result.Manifest.Files).To(HaveLen(len(snapshot.Files))) + for i, file := range result.Manifest.Files { + Expect(file.Path).To(Equal(snapshot.Files[i].Path), + "manifest entry %d is out of snapshot order", i) + Expect(file.SHA256).To(HaveLen(64)) + } + + // Every shard must also be on disk, not merely recorded. + for _, file := range snapshot.Files { + onDisk := filepath.Join(modelsPath, filepath.FromSlash(result.RelativePath), filepath.FromSlash(file.Path)) + info, statErr := os.Stat(onDisk) + Expect(statErr).NotTo(HaveOccurred()) + Expect(info.Size()).To(Equal(file.Size)) + } + }) + + It("produces the same manifest sequentially and concurrently", func() { + sequentialSnapshot, sequentialServer, _ := shardedSnapshot(8, 0) + DeferCleanup(sequentialServer.Close) + sequential, err := modelartifacts.NewManager(&fakeSnapshotResolver{snapshot: sequentialSnapshot}). + Ensure(context.Background(), GinkgoT().TempDir(), spec) + Expect(err).NotTo(HaveOccurred()) + + concurrentSnapshot, concurrentServer, _ := shardedSnapshot(8, 0) + DeferCleanup(concurrentServer.Close) + concurrent, err := modelartifacts.NewManager(&fakeSnapshotResolver{snapshot: concurrentSnapshot}, + modelartifacts.WithDownloadConcurrency(8)). + Ensure(context.Background(), GinkgoT().TempDir(), spec) + Expect(err).NotTo(HaveOccurred()) + + Expect(concurrent.Manifest.Files).To(Equal(sequential.Manifest.Files)) + }) + + It("applies live concurrency updates to subsequent materializations", func() { + snapshot, server, peak := shardedSnapshot(8, 40*time.Millisecond) + DeferCleanup(server.Close) + manager := modelartifacts.NewManager(&fakeSnapshotResolver{snapshot: snapshot}) + manager.SetDownloadConcurrency(4) + + _, err := manager.Ensure(context.Background(), GinkgoT().TempDir(), spec) + + Expect(err).NotTo(HaveOccurred()) + Expect(*peak).To(BeNumerically(">", 1)) + Expect(*peak).To(BeNumerically("<=", 4)) + }) + + It("still resumes past files an interrupted pass already completed", func() { + snapshot, server, _ := shardedSnapshot(6, 0) + DeferCleanup(server.Close) + + var requests atomic.Int32 + counting := httptest.NewServer(http.HandlerFunc(func(w http.ResponseWriter, r *http.Request) { + requests.Add(1) + server.Config.Handler.ServeHTTP(w, r) + })) + DeferCleanup(counting.Close) + for i := range snapshot.Files { + snapshot.Files[i].URL = counting.URL + snapshot.Files[i].URL[len(server.URL):] + } + + manager := modelartifacts.NewManager(&fakeSnapshotResolver{snapshot: snapshot}, + modelartifacts.WithDownloadConcurrency(3)) + modelsPath := GinkgoT().TempDir() + + first, err := manager.Ensure(context.Background(), modelsPath, spec) + Expect(err).NotTo(HaveOccurred()) + Expect(requests.Load()).To(Equal(int32(len(snapshot.Files)))) + + // A committed artifact is served from cache without touching the network. + second, err := manager.Ensure(context.Background(), modelsPath, first.Spec) + Expect(err).NotTo(HaveOccurred()) + Expect(second.CacheHit).To(BeTrue()) + Expect(requests.Load()).To(Equal(int32(len(snapshot.Files)))) + }) + + It("fails the whole materialization when a shard cannot be verified", func() { + snapshot, server, _ := shardedSnapshot(6, 0) + DeferCleanup(server.Close) + // Corrupt one shard's expected digest: the download succeeds, the + // per-file SHA check does not. + snapshot.Files[3].LFSOID = hex.EncodeToString(make([]byte, 32)) + + manager := modelartifacts.NewManager(&fakeSnapshotResolver{snapshot: snapshot}, + modelartifacts.WithDownloadConcurrency(3)) + modelsPath := GinkgoT().TempDir() + + _, err := manager.Ensure(context.Background(), modelsPath, spec) + Expect(err).To(HaveOccurred()) + + // Nothing may be published under the final path when a shard failed. + entries, readErr := os.ReadDir(filepath.Join(modelsPath, ".artifacts", "huggingface")) + if readErr == nil { + for _, entry := range entries { + _, statErr := os.Stat(filepath.Join(modelsPath, ".artifacts", "huggingface", entry.Name(), "manifest.json")) + Expect(statErr).To(HaveOccurred(), "a failed materialization published a manifest") + } + } + }) +})