From 894eaaa86fb16119cfca78e4963814dda8669912 Mon Sep 17 00:00:00 2001 From: Andrey Cheptsov Date: Wed, 15 Jul 2026 13:31:34 +0200 Subject: [PATCH] Show host GPU driver on fleet instances Detect the installed accelerator driver (NVIDIA, AMD, Tenstorrent) on the host via dstack-shim, report it through the existing healthcheck response, store it in JobProvisioningData, and surface it on the Instance API model and as a DRIVER column in `dstack fleet -v`. The field stays None when the driver cannot be detected. Co-Authored-By: Claude Fable 5 --- runner/docs/shim.openapi.yaml | 15 ++ runner/internal/shim/api/api_test.go | 6 + runner/internal/shim/api/handlers.go | 9 +- runner/internal/shim/api/handlers_test.go | 26 ++++ runner/internal/shim/api/schemas.go | 4 + runner/internal/shim/api/server.go | 2 + runner/internal/shim/docker.go | 6 + runner/internal/shim/host/gpu.go | 134 ++++++++++++++---- runner/internal/shim/host/gpu_test.go | 102 +++++++++++++ src/dstack/_internal/cli/utils/fleet.py | 3 + .../core/backends/kubernetes/compute.py | 4 + .../_internal/core/backends/vastai/compute.py | 1 + src/dstack/_internal/core/models/instances.py | 19 ++- src/dstack/_internal/core/models/runs.py | 7 + .../pipeline_tasks/instances/check.py | 6 + .../pipeline_tasks/instances/common.py | 20 +++ .../pipeline_tasks/instances/ssh_deploy.py | 1 + .../_internal/server/schemas/instances.py | 2 + src/dstack/_internal/server/schemas/runner.py | 5 + .../_internal/server/services/instances.py | 1 + .../server/services/runner/client.py | 11 +- .../test_instances/test_check.py | 90 +++++++++++- .../test_instances/test_ssh_deploy.py | 45 +++++- .../_internal/server/routers/test_fleets.py | 4 + .../server/services/runner/test_client.py | 43 ++++++ .../server/services/test_instances.py | 5 +- 26 files changed, 534 insertions(+), 37 deletions(-) diff --git a/runner/docs/shim.openapi.yaml b/runner/docs/shim.openapi.yaml index e375e4e9d3..55ed252602 100644 --- a/runner/docs/shim.openapi.yaml +++ b/runner/docs/shim.openapi.yaml @@ -479,6 +479,21 @@ components: type: string examples: - 0.18.34 + gpu_vendor: + description: > + (since [0.20.28](https://github.com/dstackai/dstack/releases/tag/0.20.28)) + Host GPU vendor. Omitted on hosts without GPUs. + type: string + examples: + - nvidia + gpu_driver_version: + description: > + (since [0.20.28](https://github.com/dstackai/dstack/releases/tag/0.20.28)) + Host GPU driver version. Omitted on hosts without GPUs + or if detection failed. + type: string + examples: + - 570.86.15 required: - service - version diff --git a/runner/internal/shim/api/api_test.go b/runner/internal/shim/api/api_test.go index b6879187af..777c8e67a6 100644 --- a/runner/internal/shim/api/api_test.go +++ b/runner/internal/shim/api/api_test.go @@ -5,10 +5,12 @@ import ( "sync" "github.com/dstackai/dstack/runner/internal/shim" + "github.com/dstackai/dstack/runner/internal/shim/host" ) type DummyRunner struct { tasks map[string]bool + gpus []host.GpuInfo mu sync.Mutex } @@ -46,6 +48,10 @@ func (ds *DummyRunner) Resources(context.Context) shim.Resources { return shim.Resources{} } +func (ds *DummyRunner) Gpus(context.Context) []host.GpuInfo { + return ds.gpus +} + func NewDummyRunner() *DummyRunner { return &DummyRunner{ tasks: map[string]bool{}, diff --git a/runner/internal/shim/api/handlers.go b/runner/internal/shim/api/handlers.go index b3382d0f26..9cbbbb9aaf 100644 --- a/runner/internal/shim/api/handlers.go +++ b/runner/internal/shim/api/handlers.go @@ -16,10 +16,15 @@ func (s *ShimServer) HealthcheckHandler(w http.ResponseWriter, r *http.Request) s.mu.RLock() defer s.mu.RUnlock() - return &HealthcheckResponse{ + response := &HealthcheckResponse{ Service: "dstack-shim", Version: s.version, - }, nil + } + if gpus := s.runner.Gpus(r.Context()); len(gpus) > 0 { + response.GpuVendor = string(gpus[0].Vendor) + response.GpuDriverVersion = gpus[0].DriverVersion + } + return response, nil } func (s *ShimServer) ShutdownHandler(w http.ResponseWriter, r *http.Request) (interface{}, error) { diff --git a/runner/internal/shim/api/handlers_test.go b/runner/internal/shim/api/handlers_test.go index bb19ebbf1b..011b985bc5 100644 --- a/runner/internal/shim/api/handlers_test.go +++ b/runner/internal/shim/api/handlers_test.go @@ -7,6 +7,8 @@ import ( "testing" commonapi "github.com/dstackai/dstack/runner/internal/common/api" + "github.com/dstackai/dstack/runner/internal/common/gpu" + "github.com/dstackai/dstack/runner/internal/shim/host" ) func TestHealthcheck(t *testing.T) { @@ -29,6 +31,30 @@ func TestHealthcheck(t *testing.T) { } } +func TestHealthcheckWithGpus(t *testing.T) { + request := httptest.NewRequest("GET", "/api/healthcheck", nil) + responseRecorder := httptest.NewRecorder() + + runner := NewDummyRunner() + runner.gpus = []host.GpuInfo{ + {Vendor: gpu.GpuVendorNvidia, Name: "T4", Vram: 16384, DriverVersion: "570.86.15"}, + } + server := NewShimServer(context.Background(), ":12346", "0.0.1.dev2", runner, nil, nil, nil, nil) + + f := commonapi.JSONResponseHandler(server.HealthcheckHandler) + f(responseRecorder, request) + + if responseRecorder.Code != 200 { + t.Errorf("Want status '%d', got '%d'", 200, responseRecorder.Code) + } + + expected := `{"service":"dstack-shim","version":"0.0.1.dev2","gpu_vendor":"nvidia","gpu_driver_version":"570.86.15"}` + + if strings.TrimSpace(responseRecorder.Body.String()) != expected { + t.Errorf("Want '%s', got '%s'", expected, responseRecorder.Body.String()) + } +} + func TestTaskSubmit(t *testing.T) { server := NewShimServer(context.Background(), ":12340", "0.0.1.dev2", NewDummyRunner(), nil, nil, nil, nil) requestBody := `{ diff --git a/runner/internal/shim/api/schemas.go b/runner/internal/shim/api/schemas.go index 0e96028a5b..ce3dccd7b9 100644 --- a/runner/internal/shim/api/schemas.go +++ b/runner/internal/shim/api/schemas.go @@ -9,6 +9,10 @@ import ( type HealthcheckResponse struct { Service string `json:"service"` Version string `json:"version"` + // Optional host GPU driver info; empty on hosts without GPUs or if + // detection failed. Old servers ignore these fields. + GpuVendor string `json:"gpu_vendor,omitempty"` + GpuDriverVersion string `json:"gpu_driver_version,omitempty"` } type ShutdownRequest struct { diff --git a/runner/internal/shim/api/server.go b/runner/internal/shim/api/server.go index 9008aa2efe..28780295b5 100644 --- a/runner/internal/shim/api/server.go +++ b/runner/internal/shim/api/server.go @@ -13,6 +13,7 @@ import ( "github.com/dstackai/dstack/runner/internal/shim" "github.com/dstackai/dstack/runner/internal/shim/components" "github.com/dstackai/dstack/runner/internal/shim/dcgm" + "github.com/dstackai/dstack/runner/internal/shim/host" ) type TaskRunner interface { @@ -22,6 +23,7 @@ type TaskRunner interface { Remove(ctx context.Context, taskID string) error Resources(context.Context) shim.Resources + Gpus(context.Context) []host.GpuInfo TaskList() []*shim.TaskListItem TaskInfo(taskID string) shim.TaskInfo } diff --git a/runner/internal/shim/docker.go b/runner/internal/shim/docker.go index b149516a6c..7fbda1458d 100644 --- a/runner/internal/shim/docker.go +++ b/runner/internal/shim/docker.go @@ -324,6 +324,12 @@ func (d *DockerRunner) Resources(ctx context.Context) Resources { } } +// Gpus returns the GPUs detected at startup without collecting other host +// resources, making it suitable for frequently called paths. +func (d *DockerRunner) Gpus(ctx context.Context) []host.GpuInfo { + return d.gpus +} + func (d *DockerRunner) TaskList() []*TaskListItem { tasks := d.tasks.List() result := make([]*TaskListItem, 0, len(tasks)) diff --git a/runner/internal/shim/host/gpu.go b/runner/internal/shim/host/gpu.go index eff57f2e00..fd51d1fc8a 100644 --- a/runner/internal/shim/host/gpu.go +++ b/runner/internal/shim/host/gpu.go @@ -7,6 +7,7 @@ import ( "errors" "fmt" "io" + "os" "path/filepath" "strconv" "strings" @@ -41,6 +42,10 @@ type GpuInfo struct { // AMD: empty string // Intel: accelerator index: ("0", "1", ...), as reported by `hl-smi -Q index` Index string + // Version of the installed host driver, e.g., "570.86.15" (NVIDIA), + // "6.10.5" (AMD amdgpu), "2.0.0" (Tenstorrent TT-KMD). + // Empty string if detection failed. All GPUs on a host share the same driver. + DriverVersion string } func GetGpuInfo(ctx context.Context) []GpuInfo { @@ -59,12 +64,23 @@ func GetGpuInfo(ctx context.Context) []GpuInfo { return []GpuInfo{} } +// normalizeDriverVersion filters out placeholder values SMI tools emit when a +// query field is not available, e.g., "N/A" or "[Not Supported]". +func normalizeDriverVersion(value string) string { + value = strings.TrimSpace(value) + switch strings.ToUpper(value) { + case "N/A", "[N/A]", "UNKNOWN", "[UNKNOWN]", "[NOT SUPPORTED]", "[NOT AVAILABLE]": + return "" + } + return value +} + func getNvidiaGpuInfo(ctx context.Context) []GpuInfo { gpus := []GpuInfo{} cmd := execute.ExecTask{ Command: "nvidia-smi", - Args: []string{"--query-gpu=name,memory.total,uuid", "--format=csv,noheader,nounits"}, + Args: []string{"--query-gpu=name,memory.total,uuid,driver_version", "--format=csv,noheader,nounits"}, StreamStdio: false, } res, err := cmd.Execute(ctx) @@ -90,8 +106,8 @@ func getNvidiaGpuInfo(ctx context.Context) []GpuInfo { log.Error(ctx, "cannot read csv", "err", err) return gpus } - if len(record) != 3 { - log.Error(ctx, "3 csv fields expected", "len", len(record)) + if len(record) != 4 { + log.Error(ctx, "4 csv fields expected", "len", len(record)) return gpus } vram, err := strconv.Atoi(strings.TrimSpace(record[1])) @@ -100,19 +116,43 @@ func getNvidiaGpuInfo(ctx context.Context) []GpuInfo { vram = 0 } gpus = append(gpus, GpuInfo{ - Vendor: gpu.GpuVendorNvidia, - Name: strings.TrimSpace(record[0]), - Vram: vram, - ID: strings.TrimSpace(record[2]), + Vendor: gpu.GpuVendorNvidia, + Name: strings.TrimSpace(record[0]), + Vram: vram, + ID: strings.TrimSpace(record[2]), + DriverVersion: normalizeDriverVersion(record[3]), }) } return gpus } type amdGpu struct { - Asic amdAsic `json:"asic"` - Vram amdVram `json:"vram"` - Bus amdBus `json:"bus"` + Asic amdAsic `json:"asic"` + Vram amdVram `json:"vram"` + Bus amdBus `json:"bus"` + Driver amdDriver `json:"driver"` +} + +// amdDriver is the `driver` section of `amd-smi static --driver`. +// Key names and value shapes differ between amd-smi versions, so it is parsed +// defensively: an unexpected format leaves Version empty instead of failing +// the whole GPU detection. Key matching is case-insensitive (encoding/json). +type amdDriver struct { + Version string +} + +func (d *amdDriver) UnmarshalJSON(data []byte) error { + var section struct { + Version string `json:"version"` + DriverVersion string `json:"driver_version"` + } + // The error is ignored deliberately: an unexpected shape leaves Version empty. + _ = json.Unmarshal(data, §ion) + d.Version = section.Version + if d.Version == "" { + d.Version = section.DriverVersion + } + return nil } // amd-smi >= 7.x wraps the array in {"gpu_data": [...]} @@ -151,38 +191,53 @@ func parseAmdSmiOutput(data []byte) ([]amdGpu, error) { return wrapped.GpuData, nil } -func getAmdGpuInfo(ctx context.Context) []GpuInfo { - gpus := []GpuInfo{} - +func execAmdSmiStatic(ctx context.Context, withDriver bool) (string, error) { ctx, cancel := context.WithTimeout(ctx, 2*time.Minute) defer cancel() + args := []string{ + "run", + "--rm", + "--device", "/dev/kfd", + "--device", "/dev/dri", + amdSmiImage, + "static", "--json", "--asic", "--vram", "--bus", + } + if withDriver { + args = append(args, "--driver") + } cmd := execute.ExecTask{ - Command: "docker", - Args: []string{ - "run", - "--rm", - "--device", "/dev/kfd", - "--device", "/dev/dri", - amdSmiImage, - "static", "--json", "--asic", "--vram", "--bus", - }, + Command: "docker", + Args: args, StreamStdio: false, } res, err := cmd.Execute(ctx) if err != nil { - log.Error(ctx, "failed to execute amd-smi", "err", err) - return gpus + return "", err } if res.ExitCode != 0 { - log.Error( - ctx, "failed to execute amd-smi", - "exitcode", res.ExitCode, "stdout", res.Stdout, "stderr", res.Stderr, + return "", fmt.Errorf( + "exitcode: %d, stdout: %s, stderr: %s", res.ExitCode, res.Stdout, res.Stderr, ) + } + return res.Stdout, nil +} + +func getAmdGpuInfo(ctx context.Context) []GpuInfo { + gpus := []GpuInfo{} + + stdout, err := execAmdSmiStatic(ctx, true) + if err != nil { + // Fall back for amd-smi versions without the --driver option. + log.Error(ctx, "failed to execute amd-smi with --driver, retrying without", "err", err) + stdout, err = execAmdSmiStatic(ctx, false) + } + if err != nil { + log.Error(ctx, "failed to execute amd-smi", "err", err) return gpus } - amdGpus, err := parseAmdSmiOutput([]byte(res.Stdout)) + amdGpus, err := parseAmdSmiOutput([]byte(stdout)) if err != nil { log.Error(ctx, "cannot read json", "err", err) return gpus @@ -198,6 +253,7 @@ func getAmdGpuInfo(ctx context.Context) []GpuInfo { Name: amdGpu.Asic.Name, Vram: amdGpu.Vram.Size.Value, RenderNodePath: renderNodePath, + DriverVersion: amdGpu.Driver.Version, }) } return gpus @@ -413,6 +469,20 @@ func getGpusFromTtSmiSnapshot(snapshot *ttSmiSnapshot) []GpuInfo { return gpus } +// tenstorrentDriverVersionPath is the TT-KMD version file; it is what tt-smi +// itself reads to report the driver version. It is a variable so tests can +// override it. +var tenstorrentDriverVersionPath = "/sys/module/tenstorrent/version" + +func getTenstorrentDriverVersion(ctx context.Context) string { + data, err := os.ReadFile(tenstorrentDriverVersionPath) + if err != nil { + log.Error(ctx, "failed to read tenstorrent driver version", "err", err) + return "" + } + return strings.TrimSpace(string(data)) +} + func getTenstorrentGpuInfo(ctx context.Context) []GpuInfo { gpus := []GpuInfo{} @@ -447,7 +517,13 @@ func getTenstorrentGpuInfo(ctx context.Context) []GpuInfo { return gpus } - return getGpusFromTtSmiSnapshot(ttSmiSnapshot) + gpus = getGpusFromTtSmiSnapshot(ttSmiSnapshot) + if driverVersion := getTenstorrentDriverVersion(ctx); driverVersion != "" { + for i := range gpus { + gpus[i].DriverVersion = driverVersion + } + } + return gpus } func getAmdRenderNodePath(bdf string) (string, error) { diff --git a/runner/internal/shim/host/gpu_test.go b/runner/internal/shim/host/gpu_test.go index 4110816003..d33b2c2119 100644 --- a/runner/internal/shim/host/gpu_test.go +++ b/runner/internal/shim/host/gpu_test.go @@ -15,6 +15,108 @@ func loadTestData(filename string) ([]byte, error) { return os.ReadFile(path) } +func TestParseAmdSmiOutputWithDriver(t *testing.T) { + tests := []struct { + name string + data string + wantName string + wantDriver string + }{ + { + name: "rocm 6.x flat array with driver", + data: `[{"gpu": 0, "asic": {"market_name": "MI300X"}, "vram": {"size": {"value": 196592, "unit": "MB"}},` + + ` "bus": {"bdf": "0000:05:00.0"}, "driver": {"name": "amdgpu", "version": "6.10.5"}}]`, + wantName: "MI300X", + wantDriver: "6.10.5", + }, + { + name: "version preferred over driver_version when both present", + data: `[{"gpu": 0, "asic": {"market_name": "MI300X"}, "vram": {"size": {"value": 196592}},` + + ` "bus": {"bdf": "0000:05:00.0"}, "driver": {"version": "6.10.5", "driver_version": "6.8.5"}}]`, + wantName: "MI300X", + wantDriver: "6.10.5", + }, + { + name: "rocm 7.x wrapped with uppercase driver keys", + data: `{"gpu_data": [{"gpu": 0, "asic": {"market_name": "MI300X"}, "vram": {"size": {"value": 196592}},` + + ` "bus": {"bdf": "0000:05:00.0"}, "driver": {"NAME": "amdgpu", "VERSION": "6.12.12"}}]}`, + wantName: "MI300X", + wantDriver: "6.12.12", + }, + { + name: "driver_version key variant", + data: `[{"gpu": 0, "asic": {"market_name": "MI300X"}, "vram": {"size": {"value": 196592}},` + + ` "bus": {"bdf": "0000:05:00.0"}, "driver": {"driver_name": "amdgpu", "driver_version": "6.8.5"}}]`, + wantName: "MI300X", + wantDriver: "6.8.5", + }, + { + name: "no driver section", + data: `[{"gpu": 0, "asic": {"market_name": "MI300X"}, "vram": {"size": {"value": 196592}},` + + ` "bus": {"bdf": "0000:05:00.0"}}]`, + wantName: "MI300X", + wantDriver: "", + }, + { + name: "unexpected driver section shape does not fail parsing", + data: `[{"gpu": 0, "asic": {"market_name": "MI300X"}, "vram": {"size": {"value": 196592}},` + + ` "bus": {"bdf": "0000:05:00.0"}, "driver": "amdgpu 6.8.5"}]`, + wantName: "MI300X", + wantDriver: "", + }, + } + + for _, tt := range tests { + t.Run(tt.name, func(t *testing.T) { + amdGpus, err := parseAmdSmiOutput([]byte(tt.data)) + if err != nil { + t.Fatalf("parseAmdSmiOutput() error = %v", err) + } + if len(amdGpus) != 1 { + t.Fatalf("parseAmdSmiOutput() returned %d GPUs, want 1", len(amdGpus)) + } + if amdGpus[0].Asic.Name != tt.wantName { + t.Errorf("name = %q, want %q", amdGpus[0].Asic.Name, tt.wantName) + } + if amdGpus[0].Driver.Version != tt.wantDriver { + t.Errorf("driver version = %q, want %q", amdGpus[0].Driver.Version, tt.wantDriver) + } + }) + } +} + +func TestNormalizeDriverVersion(t *testing.T) { + for input, want := range map[string]string{ + " 570.86.15 ": "570.86.15", + "N/A": "", + "[Not Supported]": "", + "Unknown": "", + } { + if got := normalizeDriverVersion(input); got != want { + t.Errorf("normalizeDriverVersion(%q) = %q, want %q", input, got, want) + } + } +} + +func TestGetTenstorrentDriverVersion(t *testing.T) { + versionFile := filepath.Join(t.TempDir(), "version") + if err := os.WriteFile(versionFile, []byte("2.0.0\n"), 0o644); err != nil { + t.Fatalf("failed to write version file: %v", err) + } + origPath := tenstorrentDriverVersionPath + tenstorrentDriverVersionPath = versionFile + defer func() { tenstorrentDriverVersionPath = origPath }() + + if got := getTenstorrentDriverVersion(t.Context()); got != "2.0.0" { + t.Errorf("getTenstorrentDriverVersion() = %q, want %q", got, "2.0.0") + } + + tenstorrentDriverVersionPath = filepath.Join(t.TempDir(), "nonexistent") + if got := getTenstorrentDriverVersion(t.Context()); got != "" { + t.Errorf("getTenstorrentDriverVersion() = %q, want empty string", got) + } +} + func TestUnmarshalTtSmiSnapshot(t *testing.T) { tests := []struct { name string diff --git a/src/dstack/_internal/cli/utils/fleet.py b/src/dstack/_internal/cli/utils/fleet.py index ccb2400857..acf6290296 100644 --- a/src/dstack/_internal/cli/utils/fleet.py +++ b/src/dstack/_internal/cli/utils/fleet.py @@ -37,6 +37,8 @@ def get_fleets_table( table.add_column("GPU") table.add_column("SPOT") table.add_column("BACKEND") + if verbose: + table.add_column("DRIVER") table.add_column("PRICE") table.add_column("STATUS", no_wrap=True) table.add_column("CREATED", no_wrap=True) @@ -123,6 +125,7 @@ def get_fleets_table( "RESOURCES": _format_instance_resources(instance), "GPU": _format_instance_gpu(instance), "BACKEND": backend_with_region, + "DRIVER": instance.gpu_driver.version if instance.gpu_driver else "-", "PRICE": instance_price, "SPOT": instance_spot, "STATUS": _format_instance_status(instance), diff --git a/src/dstack/_internal/core/backends/kubernetes/compute.py b/src/dstack/_internal/core/backends/kubernetes/compute.py index b69e546a30..f925fc5019 100644 --- a/src/dstack/_internal/core/backends/kubernetes/compute.py +++ b/src/dstack/_internal/core/backends/kubernetes/compute.py @@ -394,6 +394,10 @@ def update_provisioning_data( provisioning_data.hostname = get_or_error(service_spec.cluster_ip) pod_spec = get_or_error(pod.spec) node = api.read_node(name=get_or_error(pod_spec.node_name)) + # TODO: Set provisioning_data.gpu_driver from the node labels: + # nvidia.com/cuda.driver-version.full set by GPU Feature Discovery + # (or nvidia.com/cuda.driver.{major,minor,rev} set by older GFD versions), + # amd.com/gpu.driver-version set by the AMD GPU Operator. instance_offer = get_instance_offer_from_node(node=node, region=cluster.region) if instance_offer is not None: resource_requirements = get_or_error(pod_spec.containers[0].resources) diff --git a/src/dstack/_internal/core/backends/vastai/compute.py b/src/dstack/_internal/core/backends/vastai/compute.py index fc1a957078..1424720c0b 100644 --- a/src/dstack/_internal/core/backends/vastai/compute.py +++ b/src/dstack/_internal/core/backends/vastai/compute.py @@ -169,6 +169,7 @@ def update_provisioning_data( provisioning_data.ssh_port = int( resp["ports"][f"{DSTACK_RUNNER_SSH_PORT}/tcp"][0]["HostPort"] ) + # TODO: Set provisioning_data.gpu_driver from resp["driver_version"] if ( resp["actual_status"] == "created" and ": OCI runtime create failed:" in resp["status_msg"] diff --git a/src/dstack/_internal/core/models/instances.py b/src/dstack/_internal/core/models/instances.py index dfce209c32..de27b5021c 100644 --- a/src/dstack/_internal/core/models/instances.py +++ b/src/dstack/_internal/core/models/instances.py @@ -4,7 +4,7 @@ from uuid import UUID import gpuhunt -from pydantic import Field, root_validator +from pydantic import Field, root_validator, validator from dstack._internal.core.models.backends.base import BackendType from dstack._internal.core.models.common import ( @@ -46,6 +46,21 @@ def validate_name_and_vendor(cls, values): return values +class GpuDriverInfo(CoreModel): + vendor: Optional[gpuhunt.AcceleratorVendor] = None + version: str + + @validator("vendor", pre=True) + def _cast_vendor(cls, v: Any) -> Optional[gpuhunt.AcceleratorVendor]: + if v is None or isinstance(v, gpuhunt.AcceleratorVendor): + return v + try: + return gpuhunt.AcceleratorVendor.cast(v) + except ValueError: + # Tolerate vendors unknown to this server/client version + return None + + class Disk(CoreModel): size_mib: int @@ -324,3 +339,5 @@ class Instance(CoreModel): price: Optional[float] = None total_blocks: Optional[int] = None busy_blocks: int = 0 + gpu_driver: Optional[GpuDriverInfo] = None + """`gpu_driver` is the accelerator driver installed on the host, when known.""" diff --git a/src/dstack/_internal/core/models/runs.py b/src/dstack/_internal/core/models/runs.py index 04f4c326d8..393454c95e 100644 --- a/src/dstack/_internal/core/models/runs.py +++ b/src/dstack/_internal/core/models/runs.py @@ -30,6 +30,7 @@ ) from dstack._internal.core.models.files import FileArchiveMapping from dstack._internal.core.models.instances import ( + GpuDriverInfo, InstanceOfferWithAvailability, InstanceType, SSHConnectionParams, @@ -336,6 +337,12 @@ class JobProvisioningData(CoreModel): ssh_proxy: Optional[SSHConnectionParams] = None backend_data: Optional[str] = None """`backend_data` stores backend-specific data in JSON.""" + gpu_driver: Optional[GpuDriverInfo] = None + """`gpu_driver` is the accelerator driver installed on the host, when known. + Detected via the shim for VM-based backends and SSH fleets; taken from the + provider API or node labels for some container-based backends. May be set + after provisioning. + """ def get_base_backend(self) -> BackendType: if self.base_backend is not None: diff --git a/src/dstack/_internal/server/background/pipeline_tasks/instances/check.py b/src/dstack/_internal/server/background/pipeline_tasks/instances/check.py index 486c83dbf6..c3d16b0767 100644 --- a/src/dstack/_internal/server/background/pipeline_tasks/instances/check.py +++ b/src/dstack/_internal/server/background/pipeline_tasks/instances/check.py @@ -32,6 +32,7 @@ can_terminate_fleet_instances_on_idle_duration, get_instance_idle_duration, get_provisioning_deadline, + set_gpu_driver_update, set_health_update, set_status_update, set_unreachable_update, @@ -181,6 +182,11 @@ async def check_instance(instance_model: InstanceModel) -> ProcessResult: if instance_check.reachable: result.instance_update_map["termination_deadline"] = None + set_gpu_driver_update( + update_map=result.instance_update_map, + job_provisioning_data=job_provisioning_data, + gpu_driver=instance_check.gpu_driver, + ) if instance_model.status == InstanceStatus.PROVISIONING: set_status_update( update_map=result.instance_update_map, diff --git a/src/dstack/_internal/server/background/pipeline_tasks/instances/common.py b/src/dstack/_internal/server/background/pipeline_tasks/instances/common.py index a386960478..55cc68f510 100644 --- a/src/dstack/_internal/server/background/pipeline_tasks/instances/common.py +++ b/src/dstack/_internal/server/background/pipeline_tasks/instances/common.py @@ -11,6 +11,7 @@ from dstack._internal.core.models.backends.base import BackendType from dstack._internal.core.models.health import HealthStatus from dstack._internal.core.models.instances import ( + GpuDriverInfo, InstanceStatus, InstanceTerminationReason, SSHKey, @@ -176,3 +177,22 @@ def set_unreachable_update( return False update_map["unreachable"] = unreachable return True + + +def set_gpu_driver_update( + update_map: InstanceUpdateMap, + job_provisioning_data: JobProvisioningData, + gpu_driver: Optional[GpuDriverInfo], +) -> bool: + """ + Stores the shim-reported GPU driver in the instance provisioning data. + Also fills it for instances provisioned before the server upgrade. + """ + if gpu_driver is None: + return False + current = job_provisioning_data.gpu_driver + if current is not None and current.dict() == gpu_driver.dict(): + return False + job_provisioning_data.gpu_driver = gpu_driver + update_map["job_provisioning_data"] = job_provisioning_data.json() + return True diff --git a/src/dstack/_internal/server/background/pipeline_tasks/instances/ssh_deploy.py b/src/dstack/_internal/server/background/pipeline_tasks/instances/ssh_deploy.py index b4e3e1122a..dbf5ce398c 100644 --- a/src/dstack/_internal/server/background/pipeline_tasks/instances/ssh_deploy.py +++ b/src/dstack/_internal/server/background/pipeline_tasks/instances/ssh_deploy.py @@ -177,6 +177,7 @@ async def add_ssh_instance(instance_model: InstanceModel) -> ProcessResult: dockerized=True, backend_data=None, ssh_proxy=remote_details.ssh_proxy, + gpu_driver=health.gpu_driver, ) instance_offer = InstanceOfferWithAvailability( backend=BackendType.REMOTE, diff --git a/src/dstack/_internal/server/schemas/instances.py b/src/dstack/_internal/server/schemas/instances.py index 8f87935b92..8d0850983b 100644 --- a/src/dstack/_internal/server/schemas/instances.py +++ b/src/dstack/_internal/server/schemas/instances.py @@ -4,6 +4,7 @@ from dstack._internal.core.models.common import CoreModel from dstack._internal.core.models.health import HealthCheck, HealthStatus +from dstack._internal.core.models.instances import GpuDriverInfo from dstack._internal.server.schemas.runner import InstanceHealthResponse @@ -26,6 +27,7 @@ class InstanceCheck(CoreModel): reachable: bool message: Optional[str] = None health_response: Optional[InstanceHealthResponse] = None + gpu_driver: Optional[GpuDriverInfo] = None def get_health_status(self) -> HealthStatus: if self.health_response is None: diff --git a/src/dstack/_internal/server/schemas/runner.py b/src/dstack/_internal/server/schemas/runner.py index c1ad0407d0..c22bb9cb16 100644 --- a/src/dstack/_internal/server/schemas/runner.py +++ b/src/dstack/_internal/server/schemas/runner.py @@ -125,6 +125,11 @@ class SubmitBody(CoreModel): class HealthcheckResponse(CoreModel): service: str version: str + gpu_vendor: Optional[str] = None + """`gpu_vendor` is not set by old shims and on hosts without GPUs.""" + gpu_driver_version: Optional[str] = None + """`gpu_driver_version` is not set by old shims, on hosts without GPUs, + and when driver detection fails.""" class InstanceHealthResponse(CoreModel): diff --git a/src/dstack/_internal/server/services/instances.py b/src/dstack/_internal/server/services/instances.py index 913d3c9f44..71cd3aac2d 100644 --- a/src/dstack/_internal/server/services/instances.py +++ b/src/dstack/_internal/server/services/instances.py @@ -259,6 +259,7 @@ def instance_model_to_instance(instance_model: InstanceModel) -> Instance: instance.instance_type = jpd.instance_type instance.hostname = jpd.hostname instance.availability_zone = jpd.availability_zone + instance.gpu_driver = jpd.gpu_driver return instance diff --git a/src/dstack/_internal/server/services/runner/client.py b/src/dstack/_internal/server/services/runner/client.py index 7ccc2b1af7..5f2d6a5923 100644 --- a/src/dstack/_internal/server/services/runner/client.py +++ b/src/dstack/_internal/server/services/runner/client.py @@ -14,6 +14,7 @@ from dstack._internal.core.errors import DstackError from dstack._internal.core.models.common import CoreModel, NetworkMode from dstack._internal.core.models.envs import Env +from dstack._internal.core.models.instances import GpuDriverInfo from dstack._internal.core.models.repos.remote import RemoteRepoCreds from dstack._internal.core.models.resources import Memory from dstack._internal.core.models.runs import ClusterInfo, Job, Run @@ -680,8 +681,16 @@ def healthcheck_response_to_instance_check( and instance_health_response.dcgm.incidents ): message = instance_health_response.dcgm.incidents[0].error_message + gpu_driver = None + if response.gpu_driver_version: + gpu_driver = GpuDriverInfo.parse_obj( + {"vendor": response.gpu_vendor, "version": response.gpu_driver_version} + ) return InstanceCheck( - reachable=True, health_response=instance_health_response, message=message + reachable=True, + health_response=instance_health_response, + message=message, + gpu_driver=gpu_driver, ) return InstanceCheck( reachable=False, diff --git a/src/tests/_internal/server/background/pipeline_tasks/test_instances/test_check.py b/src/tests/_internal/server/background/pipeline_tasks/test_instances/test_check.py index 33e57df016..6cb851e05a 100644 --- a/src/tests/_internal/server/background/pipeline_tasks/test_instances/test_check.py +++ b/src/tests/_internal/server/background/pipeline_tasks/test_instances/test_check.py @@ -4,16 +4,25 @@ import pytest import pytest_asyncio +from gpuhunt import AcceleratorVendor from sqlalchemy import select from sqlalchemy.ext.asyncio import AsyncSession from dstack._internal.core.models.fleets import FleetNodesSpec from dstack._internal.core.models.health import HealthStatus -from dstack._internal.core.models.instances import InstanceStatus, InstanceTerminationReason +from dstack._internal.core.models.instances import ( + GpuDriverInfo, + InstanceStatus, + InstanceTerminationReason, +) from dstack._internal.core.models.profiles import TerminationPolicy -from dstack._internal.core.models.runs import JobStatus +from dstack._internal.core.models.runs import JobProvisioningData, JobStatus from dstack._internal.server.background.pipeline_tasks.instances import InstanceWorker from dstack._internal.server.background.pipeline_tasks.instances import check as instances_check +from dstack._internal.server.background.pipeline_tasks.instances.common import ( + InstanceUpdateMap, + set_gpu_driver_update, +) from dstack._internal.server.models import InstanceHealthCheckModel, InstanceModel from dstack._internal.server.schemas.health.dcgm import DCGMHealthResponse, DCGMHealthResult from dstack._internal.server.schemas.instances import InstanceCheck @@ -36,6 +45,7 @@ create_user, get_fleet_configuration, get_fleet_spec, + get_job_provisioning_data, get_remote_connection_info, list_events, ) @@ -146,6 +156,41 @@ async def test_check_shim_transitions_provisioning_on_busy( assert instance.termination_deadline is None assert job.instance == instance + async def test_check_shim_stores_gpu_driver( + self, + test_db, + session: AsyncSession, + worker: InstanceWorker, + monkeypatch: pytest.MonkeyPatch, + ): + project = await create_project(session=session) + instance = await create_instance( + session=session, + project=project, + status=InstanceStatus.IDLE, + ) + await session.commit() + + monkeypatch.setattr( + instances_check, + "_check_instance_inner", + Mock( + return_value=InstanceCheck( + reachable=True, + gpu_driver=GpuDriverInfo(vendor=AcceleratorVendor.NVIDIA, version="570.86.15"), + ) + ), + ) + await process_instance(session, worker, instance) + + await session.refresh(instance) + + assert instance.job_provisioning_data is not None + jpd = JobProvisioningData.__response__.parse_raw(instance.job_provisioning_data) + assert jpd.gpu_driver is not None + assert jpd.gpu_driver.vendor == AcceleratorVendor.NVIDIA + assert jpd.gpu_driver.version == "570.86.15" + async def test_check_shim_start_termination_deadline( self, test_db, @@ -942,3 +987,44 @@ async def test_outdated_but_shim_installation_requested( shim_client_mock.get_components.assert_called_once() shim_client_mock.shutdown.assert_not_called() + + +class TestSetGpuDriverUpdate: + def test_noop_without_driver(self): + jpd = get_job_provisioning_data(dockerized=True) + update_map = InstanceUpdateMap() + assert not set_gpu_driver_update( + update_map=update_map, + job_provisioning_data=jpd, + gpu_driver=None, + ) + assert update_map == {} + + def test_noop_when_driver_unchanged(self): + jpd = get_job_provisioning_data(dockerized=True) + jpd.gpu_driver = GpuDriverInfo(vendor=AcceleratorVendor.NVIDIA, version="570.86.15") + update_map = InstanceUpdateMap() + assert not set_gpu_driver_update( + update_map=update_map, + job_provisioning_data=jpd, + gpu_driver=GpuDriverInfo(vendor=AcceleratorVendor.NVIDIA, version="570.86.15"), + ) + assert update_map == {} + + @pytest.mark.parametrize("current_version", [None, "550.90.07"]) + def test_sets_new_or_changed_driver(self, current_version): + jpd = get_job_provisioning_data(dockerized=True) + if current_version is not None: + jpd.gpu_driver = GpuDriverInfo( + vendor=AcceleratorVendor.NVIDIA, version=current_version + ) + update_map = InstanceUpdateMap() + assert set_gpu_driver_update( + update_map=update_map, + job_provisioning_data=jpd, + gpu_driver=GpuDriverInfo(vendor=AcceleratorVendor.NVIDIA, version="570.86.15"), + ) + assert "job_provisioning_data" in update_map + parsed = JobProvisioningData.__response__.parse_raw(update_map["job_provisioning_data"]) + assert parsed.gpu_driver is not None + assert parsed.gpu_driver.version == "570.86.15" diff --git a/src/tests/_internal/server/background/pipeline_tasks/test_instances/test_ssh_deploy.py b/src/tests/_internal/server/background/pipeline_tasks/test_instances/test_ssh_deploy.py index c103458ed4..6613569531 100644 --- a/src/tests/_internal/server/background/pipeline_tasks/test_instances/test_ssh_deploy.py +++ b/src/tests/_internal/server/background/pipeline_tasks/test_instances/test_ssh_deploy.py @@ -3,15 +3,23 @@ from unittest.mock import Mock import pytest +from gpuhunt import AcceleratorVendor from sqlalchemy.ext.asyncio import AsyncSession +from dstack._internal.core.backends.base.compute import GoArchType from dstack._internal.core.errors import SSHProvisioningError from dstack._internal.core.models.backends.base import BackendType -from dstack._internal.core.models.instances import InstanceStatus, InstanceTerminationReason +from dstack._internal.core.models.instances import ( + GpuDriverInfo, + InstanceStatus, + InstanceTerminationReason, +) +from dstack._internal.core.models.runs import JobProvisioningData from dstack._internal.server.background.pipeline_tasks.instances import InstanceWorker from dstack._internal.server.background.pipeline_tasks.instances import ( ssh_deploy as instances_ssh_deploy, ) +from dstack._internal.server.schemas.instances import InstanceCheck from dstack._internal.server.testing.common import ( create_instance, create_project, @@ -98,6 +106,41 @@ async def test_adds_ssh_instance( assert instance.busy_blocks == 0 deploy_instance_mock.assert_called_once() + async def test_adds_ssh_instance_with_gpu_driver( + self, + test_db, + session: AsyncSession, + worker: InstanceWorker, + host_info: dict, + deploy_instance_mock: Mock, + ): + deploy_instance_mock.return_value = ( + InstanceCheck( + reachable=True, + gpu_driver=GpuDriverInfo(vendor=AcceleratorVendor.NVIDIA, version="570.86.15"), + ), + host_info, + GoArchType.AMD64, + ) + project = await create_project(session=session) + instance = await create_instance( + session=session, + project=project, + status=InstanceStatus.PENDING, + created_at=get_current_datetime(), + remote_connection_info=get_remote_connection_info(), + ) + await session.commit() + + await process_instance(session, worker, instance) + + await session.refresh(instance) + assert instance.status == InstanceStatus.IDLE + assert instance.job_provisioning_data is not None + jpd = JobProvisioningData.__response__.parse_raw(instance.job_provisioning_data) + assert jpd.gpu_driver is not None + assert jpd.gpu_driver.version == "570.86.15" + async def test_retries_ssh_instance_if_provisioning_fails( self, test_db, diff --git a/src/tests/_internal/server/routers/test_fleets.py b/src/tests/_internal/server/routers/test_fleets.py index 04d0145dfe..b7a48b54df 100644 --- a/src/tests/_internal/server/routers/test_fleets.py +++ b/src/tests/_internal/server/routers/test_fleets.py @@ -1006,6 +1006,7 @@ async def test_creates_fleet(self, test_db, session: AsyncSession, client: Async "price": None, "total_blocks": 1, "busy_blocks": 0, + "gpu_driver": None, } ], } @@ -1138,6 +1139,7 @@ async def test_creates_ssh_fleet(self, test_db, session: AsyncSession, client: A "price": 0.0, "total_blocks": 1, "busy_blocks": 0, + "gpu_driver": None, } ], } @@ -1358,6 +1360,7 @@ async def test_updates_ssh_fleet(self, test_db, session: AsyncSession, client: A "price": 0.0, "total_blocks": 1, "busy_blocks": 0, + "gpu_driver": None, }, { "id": SomeUUID4Str(), @@ -1393,6 +1396,7 @@ async def test_updates_ssh_fleet(self, test_db, session: AsyncSession, client: A "price": 0.0, "total_blocks": 1, "busy_blocks": 0, + "gpu_driver": None, }, ], } diff --git a/src/tests/_internal/server/services/runner/test_client.py b/src/tests/_internal/server/services/runner/test_client.py index 588c231a19..844d21ce58 100644 --- a/src/tests/_internal/server/services/runner/test_client.py +++ b/src/tests/_internal/server/services/runner/test_client.py @@ -4,10 +4,12 @@ import pytest import requests_mock +from gpuhunt import AcceleratorVendor from dstack._internal.core.consts import DSTACK_SHIM_HTTP_PORT from dstack._internal.core.models.backends.base import BackendType from dstack._internal.core.models.common import NetworkMode +from dstack._internal.core.models.instances import GpuDriverInfo from dstack._internal.core.models.resources import Memory from dstack._internal.core.models.volumes import ( InstanceMountPoint, @@ -28,6 +30,7 @@ ShimClient, ShimHTTPError, _parse_version, + healthcheck_response_to_instance_check, ) from dstack._internal.server.testing.common import get_volume, get_volume_configuration @@ -528,3 +531,43 @@ def test_valid_major_only(self, value: str): @pytest.mark.parametrize("value", ["", "foo", "1.12.3-next.20241231"]) def test_invalid(self, value: str): assert _parse_version(value) is None + + +class TestHealthcheckResponseToInstanceCheck: + def test_reachable_without_gpu_driver(self): + response = HealthcheckResponse(service="dstack-shim", version="0.19.0") + check = healthcheck_response_to_instance_check(response) + assert check.reachable + assert check.gpu_driver is None + + def test_reachable_with_gpu_driver(self): + response = HealthcheckResponse( + service="dstack-shim", + version="0.19.0", + gpu_vendor="nvidia", + gpu_driver_version="570.86.15", + ) + check = healthcheck_response_to_instance_check(response) + assert check.reachable + assert check.gpu_driver == GpuDriverInfo( + vendor=AcceleratorVendor.NVIDIA, version="570.86.15" + ) + + def test_reachable_with_unknown_gpu_vendor(self): + response = HealthcheckResponse( + service="dstack-shim", + version="0.19.0", + gpu_vendor="quantumx", + gpu_driver_version="1.2.3", + ) + check = healthcheck_response_to_instance_check(response) + assert check.reachable + assert check.gpu_driver is not None + assert check.gpu_driver.vendor is None + assert check.gpu_driver.version == "1.2.3" + + def test_unexpected_service(self): + response = HealthcheckResponse(service="not-dstack-shim", version="0.19.0") + check = healthcheck_response_to_instance_check(response) + assert not check.reachable + assert check.gpu_driver is None diff --git a/src/tests/_internal/server/services/test_instances.py b/src/tests/_internal/server/services/test_instances.py index cba11c67ec..7a9606d5f9 100644 --- a/src/tests/_internal/server/services/test_instances.py +++ b/src/tests/_internal/server/services/test_instances.py @@ -2,12 +2,14 @@ from unittest.mock import Mock, call import pytest +from gpuhunt import AcceleratorVendor from sqlalchemy.ext.asyncio import AsyncSession import dstack._internal.server.services.instances as instances_services from dstack._internal.core.models.backends.base import BackendType from dstack._internal.core.models.health import HealthStatus from dstack._internal.core.models.instances import ( + GpuDriverInfo, Instance, InstanceStatus, InstanceTerminationReason, @@ -502,6 +504,7 @@ async def test_converts_instance(self, test_db, session: AsyncSession): price=1.0, total_blocks=1, busy_blocks=0, + gpu_driver=GpuDriverInfo(vendor=AcceleratorVendor.NVIDIA, version="570.86.15"), ) im = InstanceModel( id=instance_id, @@ -512,7 +515,7 @@ async def test_converts_instance(self, test_db, session: AsyncSession): unreachable=False, health=HealthStatus.WARNING, project=project, - job_provisioning_data='{"ssh_proxy":null, "backend":"aws","hostname":"hostname_test","region":"eu-west","price":1.0,"username":"user1","ssh_port":12345,"dockerized":false,"instance_id":"test_instance","instance_type": {"name": "instance", "resources": {"cpus": 1, "memory_mib": 512, "gpus": [], "spot": false, "disk": {"size_mib": 102400}, "description":""}}}', + job_provisioning_data='{"ssh_proxy":null, "backend":"aws","hostname":"hostname_test","region":"eu-west","price":1.0,"username":"user1","ssh_port":12345,"dockerized":false,"instance_id":"test_instance","gpu_driver":{"vendor":"nvidia","version":"570.86.15"},"instance_type": {"name": "instance", "resources": {"cpus": 1, "memory_mib": 512, "gpus": [], "spot": false, "disk": {"size_mib": 102400}, "description":""}}}', offer='{"price":1.0, "backend":"aws", "region":"eu-west-1", "availability":"available","instance": {"name": "instance", "resources": {"cpus": 1, "memory_mib": 512, "gpus": [], "spot": false, "disk": {"size_mib": 102400}, "description":""}}}', total_blocks=1, busy_blocks=0,