diff --git a/pkg/apis/llm/llm_spec.go b/pkg/apis/llm/llm_spec.go index dd42e79774..d4deede93c 100644 --- a/pkg/apis/llm/llm_spec.go +++ b/pkg/apis/llm/llm_spec.go @@ -58,7 +58,13 @@ func (s *LLMSpecOllama) IsZero() bool { // LLMSpecVllm holds type-specific fields for vllm SKUs (includes PreferredModel). type LLMSpecVllm struct { - PreferredModel string `json:"preferred_model"` + PreferredModel string `json:"preferred_model"` + CustomizedArgs []*VllmCustomizedArg `json:"customized_args,omitempty"` +} + +type VllmCustomizedArg struct { + Key string `json:"key"` + Value string `json:"value"` } func (s *LLMSpecVllm) String() string { @@ -69,7 +75,7 @@ func (s *LLMSpecVllm) IsZero() bool { if s == nil { return true } - return s.PreferredModel == "" + return s.PreferredModel == "" && len(s.CustomizedArgs) == 0 } // LLMSpecDify holds type-specific fields for Dify SKUs (multiple image ids + customized envs). diff --git a/pkg/apis/llm/vllm_const.go b/pkg/apis/llm/vllm_const.go index 6b02f216ae..82454d86ac 100644 --- a/pkg/apis/llm/vllm_const.go +++ b/pkg/apis/llm/vllm_const.go @@ -8,81 +8,10 @@ const ( LLM_VLLM_EXEC_PATH = "python3 -m vllm.entrypoints.openai.api_server" LLM_VLLM_HF_ENDPOINT = "https://hf-mirror.com" - - // Directory constants LLM_VLLM_CACHE_DIR = "/root/.cache/huggingface" LLM_VLLM_BASE_PATH = "/data/models" LLM_VLLM_MODELS_PATH = "/data/models/huggingface" - // Health check - LLM_VLLM_HEALTH_CHECK_TIMEOUT = 120 * time.Second // 2 minutes - LLM_VLLM_HEALTH_CHECK_INTERVAL = 10 * time.Second // 10 seconds - - // Default vLLM memory params when Python estimation fails (conservative to avoid OOM) - LLM_VLLM_DEFAULT_GPU_MEMORY_UTIL = 0.9 - LLM_VLLM_DEFAULT_MAX_MODEL_LEN = 2048 - LLM_VLLM_DEFAULT_MAX_NUM_SEQS = 1 - - // Prefixes for parsing resolveModelAndParams output line (KEY=value) - LLM_VLLM_RESOLVE_OUTPUT_PREFIX_GPU_UTIL = "GPU_MEMORY_UTIL=" - LLM_VLLM_RESOLVE_OUTPUT_PREFIX_MAX_LEN = "MAX_MODEL_LEN=" - LLM_VLLM_RESOLVE_OUTPUT_PREFIX_MAX_NUM_SEQ = "MAX_NUM_SEQS=" -) - -const ( - - // vllmEstimateParamsScript is a Python script run inside the container to estimate - // --gpu-memory-utilization, --max-model-len, and --max-num-seqs from GPU memory and model config. - // Args: sys.argv[1]=model path, sys.argv[2]=tensor_parallel_size. - // Prints one line: GPU_MEMORY_UTIL=0.9 MAX_MODEL_LEN=2624 MAX_NUM_SEQS=1 (eval-safe). - LLM_VLLM_ESTIMATE_PARAMS_SCRIPT = ` -import sys, json, os -model_path = sys.argv[1] if len(sys.argv) > 1 else "" -tp = int(sys.argv[2]) if len(sys.argv) > 2 else 1 -if not model_path or not os.path.isdir(model_path): - sys.exit(1) -config_path = os.path.join(model_path, "config.json") -if not os.path.isfile(config_path): - sys.exit(1) -with open(config_path) as f: - config = json.load(f) -def get_nested(d, *keys): - for k in keys: - d = d.get(k) if isinstance(d, dict) else None - if d is None: - return None - return d -num_layers = config.get("num_hidden_layers") or config.get("n_layer") or get_nested(config, "text_config", "num_hidden_layers") or 0 -num_heads = config.get("num_attention_heads") or config.get("n_head") or get_nested(config, "text_config", "num_attention_heads") or 0 -num_kv_heads = config.get("num_key_value_heads") or get_nested(config, "text_config", "num_key_value_heads") or num_heads -hidden_size = config.get("hidden_size") or config.get("n_embd") or get_nested(config, "text_config", "hidden_size") or 0 -head_dim = hidden_size // num_heads if num_heads else (hidden_size // 64) -max_pos = config.get("max_position_embeddings") or config.get("n_positions") or get_nested(config, "text_config", "max_position_embeddings") or 4096 -if not num_layers or not num_kv_heads or not head_dim: - sys.exit(1) -try: - import torch - total_mem = sum(torch.cuda.get_device_properties(i).total_memory for i in range(torch.cuda.device_count())) -except Exception: - sys.exit(1) -gpu_util = 0.9 -activation_overhead = 2 * (1024**3) -num_params = config.get("num_parameters") or config.get("num_params") or get_nested(config, "text_config", "num_parameters") -if num_params is not None: - model_bytes = num_params * 2 -else: - model_bytes = 12 * (1024**3) -available_kv = total_mem * gpu_util - model_bytes - activation_overhead -if available_kv <= 0: - available_kv = total_mem * 0.5 -kv_per_token = num_layers * 2 * num_kv_heads * head_dim * 2 -max_num_seqs = 4 -max_model_len = int(available_kv / (kv_per_token * max_num_seqs)) -max_model_len = max(1, min(max_model_len, max_pos)) -if max_model_len < 256: - max_num_seqs = 1 - max_model_len = int(available_kv / (kv_per_token * max_num_seqs)) - max_model_len = max(1, min(max_model_len, max_pos)) -print("GPU_MEMORY_UTIL=%s MAX_MODEL_LEN=%d MAX_NUM_SEQS=%d" % (gpu_util, max_model_len, max_num_seqs)) -` + LLM_VLLM_HEALTH_CHECK_TIMEOUT = 120 * time.Second + LLM_VLLM_HEALTH_CHECK_INTERVAL = 10 * time.Second ) diff --git a/pkg/llm/drivers/llm_container/vllm.go b/pkg/llm/drivers/llm_container/vllm.go index 6d26dd7d50..858c7f97d2 100644 --- a/pkg/llm/drivers/llm_container/vllm.go +++ b/pkg/llm/drivers/llm_container/vllm.go @@ -12,6 +12,7 @@ import ( "strconv" "strings" "time" + "unicode" "yunion.io/x/log" "yunion.io/x/pkg/errors" @@ -40,6 +41,134 @@ func escapeShellSingleQuoted(s string) string { return strings.ReplaceAll(s, "'", "'\\''") } +func shellQuoteSingle(s string) string { + return "'" + escapeShellSingleQuoted(s) + "'" +} + +var protectedVLLMArgKeys = map[string]struct{}{ + "model": {}, + "served-model-name": {}, + "port": {}, + "tensor-parallel-size": {}, +} + +func validateVLLMArgKey(key string) error { + if key == "" { + return errors.Error("vllm arg key is empty") + } + if strings.HasPrefix(key, "--") { + return errors.Errorf("invalid vllm arg key %q: do not include leading --", key) + } + for _, r := range key { + if unicode.IsLetter(r) || unicode.IsDigit(r) || r == '-' || r == '_' { + continue + } + return errors.Errorf("invalid vllm arg key %q", key) + } + if _, ok := protectedVLLMArgKeys[key]; ok { + return errors.Errorf("vllm arg key %q is protected", key) + } + return nil +} + +func normalizeVLLMCustomizedArgs(args []*api.VllmCustomizedArg) ([]*api.VllmCustomizedArg, error) { + if len(args) == 0 { + return nil, nil + } + out := make([]*api.VllmCustomizedArg, 0, len(args)) + indexByKey := make(map[string]int, len(args)) + for _, arg := range args { + if arg == nil { + continue + } + key := strings.TrimSpace(arg.Key) + if err := validateVLLMArgKey(key); err != nil { + return nil, err + } + next := &api.VllmCustomizedArg{ + Key: key, + Value: arg.Value, + } + if idx, ok := indexByKey[key]; ok { + out[idx] = next + continue + } + indexByKey[key] = len(out) + out = append(out, next) + } + if len(out) == 0 { + return nil, nil + } + return out, nil +} + +func mergeVLLMCustomizedArgs(base, overrides []*api.VllmCustomizedArg) ([]*api.VllmCustomizedArg, error) { + out := make([]*api.VllmCustomizedArg, 0, len(base)+len(overrides)) + indexByKey := make(map[string]int, len(base)+len(overrides)) + appendNormalized := func(items []*api.VllmCustomizedArg) error { + normalized, err := normalizeVLLMCustomizedArgs(items) + if err != nil { + return err + } + for _, arg := range normalized { + if idx, ok := indexByKey[arg.Key]; ok { + out[idx] = arg + continue + } + indexByKey[arg.Key] = len(out) + out = append(out, arg) + } + return nil + } + if err := appendNormalized(base); err != nil { + return nil, err + } + if err := appendNormalized(overrides); err != nil { + return nil, err + } + if len(out) == 0 { + return nil, nil + } + return out, nil +} + +func buildVLLMServeFlags(modelPath string, tensorParallelSize, defaultSwapSpaceGiB int, effSpec *api.LLMSpecVllm) []string { + modelQuoted := shellQuoteSingle(modelPath) + flags := []string{ + fmt.Sprintf("--model %s", modelQuoted), + fmt.Sprintf(`--served-model-name "$(basename %s)"`, modelQuoted), + fmt.Sprintf("--port %d", api.LLM_VLLM_DEFAULT_PORT), + fmt.Sprintf("--tensor-parallel-size %d", tensorParallelSize), + fmt.Sprintf("--swap-space %d", defaultSwapSpaceGiB), + } + if effSpec == nil || len(effSpec.CustomizedArgs) == 0 { + return flags + } + + normalizedArgs, err := normalizeVLLMCustomizedArgs(effSpec.CustomizedArgs) + if err != nil { + log.Errorf("normalize vllm customized args: %v", err) + return flags + } + for _, arg := range normalizedArgs { + flagName := "--" + arg.Key + if arg.Key == "swap-space" { + if arg.Value == "" { + flags[4] = flagName + } else { + flags[4] = fmt.Sprintf("%s %s", flagName, shellQuoteSingle(arg.Value)) + } + continue + } + if arg.Value == "" { + flags = append(flags, flagName) + continue + } + flags = append(flags, fmt.Sprintf("%s %s", flagName, shellQuoteSingle(arg.Value))) + } + return flags +} + func (v *vllm) GetSpec(sku *models.SLLMSku) interface{} { if sku == nil || sku.LLMType != string(api.LLM_CONTAINER_VLLM) || sku.LLMSpec == nil || sku.LLMSpec.Vllm == nil { return nil @@ -52,18 +181,39 @@ func (v *vllm) GetEffectiveSpec(llm *models.SLLM, sku *models.SLLMSku) interface if s := v.GetSpec(sku); s != nil { skuSpec = s.(*api.LLMSpecVllm) } + var llmSpec *api.LLMSpecVllm if llm != nil && llm.LLMSpec != nil && llm.LLMSpec.Vllm != nil { - if llm.LLMSpec.Vllm.PreferredModel != "" { - out := *llm.LLMSpec.Vllm - return &out - } - // llm explicitly present but empty -> fall back to sku default + llmSpec = llm.LLMSpec.Vllm } + if skuSpec == nil && llmSpec == nil { + return nil + } + out := &api.LLMSpecVllm{} if skuSpec != nil { - out := *skuSpec - return &out + out.PreferredModel = skuSpec.PreferredModel + out.CustomizedArgs = skuSpec.CustomizedArgs } - return nil + if llmSpec != nil { + if llmSpec.PreferredModel != "" { + out.PreferredModel = llmSpec.PreferredModel + } + } + mergedArgs, err := mergeVLLMCustomizedArgs(out.CustomizedArgs, nil) + if err != nil { + log.Errorf("normalize sku vllm customized args: %v", err) + out.CustomizedArgs = nil + } else { + out.CustomizedArgs = mergedArgs + } + if llmSpec != nil { + mergedArgs, err = mergeVLLMCustomizedArgs(out.CustomizedArgs, llmSpec.CustomizedArgs) + if err != nil { + log.Errorf("merge vllm customized args: %v", err) + } else { + out.CustomizedArgs = mergedArgs + } + } + return out } func (v *vllm) ValidateLLMSkuCreateData(ctx context.Context, userCred mcclient.TokenCredential, input *api.LLMSkuCreateInput) (*api.LLMSkuCreateInput, error) { @@ -113,14 +263,31 @@ func (v *vllm) ValidateLLMCreateSpec(ctx context.Context, userCred mcclient.Toke if input == nil { return nil, nil } - preferred := "" - if input.Vllm != nil { - preferred = input.Vllm.PreferredModel + if input.Vllm == nil { + input.Vllm = &api.LLMSpecVllm{} } + + preferred := input.Vllm.PreferredModel if preferred == "" && sku != nil && sku.LLMSpec != nil && sku.LLMSpec.Vllm != nil { preferred = sku.LLMSpec.Vllm.PreferredModel } - return &api.LLMSpec{Vllm: &api.LLMSpecVllm{PreferredModel: preferred}}, nil + + spec := &api.LLMSpecVllm{} + if sku != nil && sku.LLMSpec != nil && sku.LLMSpec.Vllm != nil { + base := *sku.LLMSpec.Vllm + spec = &base + } + // Apply create overrides + if preferred != "" { + spec.PreferredModel = preferred + } + mergedArgs, err := mergeVLLMCustomizedArgs(spec.CustomizedArgs, input.Vllm.CustomizedArgs) + if err != nil { + return nil, err + } + spec.CustomizedArgs = mergedArgs + + return &api.LLMSpec{Vllm: spec}, nil } // ValidateLLMUpdateSpec implements ILLMContainerDriver. Merges preferred_model with current LLM spec; only overwrite when non-empty. @@ -128,15 +295,23 @@ func (v *vllm) ValidateLLMUpdateSpec(ctx context.Context, userCred mcclient.Toke if input == nil || input.Vllm == nil { return input, nil } - current := "" + base := &api.LLMSpecVllm{} if llm != nil && llm.LLMSpec != nil && llm.LLMSpec.Vllm != nil { - current = llm.LLMSpec.Vllm.PreferredModel + b := *llm.LLMSpec.Vllm + base = &b } - preferred := input.Vllm.PreferredModel - if preferred == "" { - preferred = current + + // preferred_model: only overwrite when non-empty + if input.Vllm.PreferredModel != "" { + base.PreferredModel = input.Vllm.PreferredModel } - return &api.LLMSpec{Vllm: &api.LLMSpecVllm{PreferredModel: preferred}}, nil + mergedArgs, err := mergeVLLMCustomizedArgs(base.CustomizedArgs, input.Vllm.CustomizedArgs) + if err != nil { + return nil, err + } + base.CustomizedArgs = mergedArgs + + return &api.LLMSpec{Vllm: base}, nil } func (v *vllm) GetContainerSpec(ctx context.Context, llm *models.SLLM, image *models.SLLMImage, sku *models.SLLMSku, props []string, devices []computeapi.SIsolatedDevice, diskId string) *computeapi.PodContainerCreateInput { @@ -242,6 +417,52 @@ func (v *vllm) GetLLMAccessUrlInfo(ctx context.Context, userCred mcclient.TokenC return models.GetLLMAccessUrlInfo(ctx, userCred, llm, input, "http", api.LLM_VLLM_DEFAULT_PORT) } +func buildVLLMHealthCheckURL(networkType, llmIP, hostAccessIP string, accessInfo *models.SAccessInfo) (string, error) { + if networkType == string(computeapi.NETWORK_TYPE_GUEST) { + if len(llmIP) == 0 { + return "", errors.Error("LLM IP is empty for guest network") + } + return fmt.Sprintf("http://%s:%d/health", llmIP, api.LLM_VLLM_DEFAULT_PORT), nil + } + if len(llmIP) > 0 { + return fmt.Sprintf("http://%s:%d/health", llmIP, api.LLM_VLLM_DEFAULT_PORT), nil + } + if len(hostAccessIP) == 0 { + return "", errors.Error("host access IP is empty") + } + port := api.LLM_VLLM_DEFAULT_PORT + if accessInfo != nil && accessInfo.AccessPort > 0 { + port = accessInfo.AccessPort + } + return fmt.Sprintf("http://%s:%d/health", hostAccessIP, port), nil +} + +// resolveModelPath resolves the model directory inside the container. +// It prefers preferredPath when it exists; otherwise it picks the first directory under models path. +// Returns (empty, nil) when no model is found. +func (v *vllm) resolveModelPath(ctx context.Context, containerId string, preferredPath string) (string, error) { + preferredQuoted := shellQuoteSingle(preferredPath) + cmd := fmt.Sprintf( + `mkdir -p %s; + preferred=%s; + if [ -n "$preferred" ] && [ -d "$preferred" ]; then model="$preferred"; else model=$(ls -d %s/* 2>/dev/null | head -n 1); fi; + if [ -z "$model" ]; then echo "NO_MODEL"; exit 0; fi; + printf '%%s\n' "$model"`, + api.LLM_VLLM_MODELS_PATH, + preferredQuoted, + api.LLM_VLLM_MODELS_PATH, + ) + out, err := exec(ctx, containerId, cmd, 30) + if err != nil { + return "", errors.Wrap(err, "exec resolve model path") + } + out = strings.TrimSpace(out) + if out == "NO_MODEL" || out == "" { + return "", nil + } + return out, nil +} + // StartLLM starts the vLLM server inside the container via exec, then waits for the health endpoint to be ready. func (v *vllm) StartLLM(ctx context.Context, userCred mcclient.TokenCredential, llm *models.SLLM) error { lc, err := llm.GetLLMContainer() @@ -261,35 +482,26 @@ func (v *vllm) StartLLM(ctx context.Context, userCred mcclient.TokenCredential, swapSpaceGiB = 1 } + effSpec := (*api.LLMSpecVllm)(nil) preferredPath := "" if eff := v.GetEffectiveSpec(llm, sku); eff != nil { - if preferred := eff.(*api.LLMSpecVllm).PreferredModel; preferred != "" { + effSpec = eff.(*api.LLMSpecVllm) + if preferred := effSpec.PreferredModel; preferred != "" { preferredPath = path.Join(api.LLM_VLLM_MODELS_PATH, preferred) } } - resolved, err := v.resolveModelAndParams(ctx, lc.CmpId, preferredPath, tensorParallelSize) + modelPath, err := v.resolveModelPath(ctx, lc.CmpId, preferredPath) if err != nil { return err } - if resolved == nil { + if modelPath == "" { return nil // no model } - modelEscaped := escapeShellSingleQuoted(resolved.ModelPath) startCmd := fmt.Sprintf( - `nohup %s --model '%s' --served-model-name "$(basename '%s')" --port %d \ - --tensor-parallel-size %d --swap-space %d --enable-prefix-caching \ - --gpu-memory-utilization %s --max-model-len %d --max-num-seqs %d \ - > /tmp/vllm.log 2>&1 &`, + "nohup %s %s > /tmp/vllm.log 2>&1 &", api.LLM_VLLM_EXEC_PATH, - modelEscaped, - modelEscaped, - api.LLM_VLLM_DEFAULT_PORT, - tensorParallelSize, - swapSpaceGiB, - resolved.GpuUtil, - resolved.MaxModelLen, - resolved.MaxNumSeqs, + strings.Join(buildVLLMServeFlags(modelPath, tensorParallelSize, swapSpaceGiB, effSpec), " "), ) _, err = exec(ctx, lc.CmpId, startCmd, 30) if err != nil { @@ -303,7 +515,20 @@ func (v *vllm) StartLLM(ctx context.Context, userCred mcclient.TokenCredential, if err != nil { return errors.Wrap(err, "get llm url for health check") } - healthURL := fmt.Sprintf("http://%s:%d/health", input.ServerIp, api.LLM_VLLM_DEFAULT_PORT) + var accessInfo *models.SAccessInfo + for i := range input.AccessInfos { + if input.AccessInfos[i].ListenPort == api.LLM_VLLM_DEFAULT_PORT { + accessInfo = &input.AccessInfos[i] + break + } + } + if accessInfo == nil && len(input.AccessInfos) > 0 { + accessInfo = &input.AccessInfos[0] + } + healthURL, err := buildVLLMHealthCheckURL(llm.NetworkType, llm.LLMIp, input.HostInternalIp, accessInfo) + if err != nil { + return errors.Wrap(err, "build health check url") + } deadline := time.Now().Add(api.LLM_VLLM_HEALTH_CHECK_TIMEOUT) for time.Now().Before(deadline) { req, err := http.NewRequestWithContext(ctx, http.MethodGet, healthURL, nil) @@ -492,7 +717,7 @@ func isNonEmptyFile(p string) bool { func (v *vllm) DownloadModel(ctx context.Context, userCred mcclient.TokenCredential, llm *models.SLLM, tmpDir string, modelName string, modelTag string) (string, []string, error) { // Download HF model on host into tmpDir for instant-model import. - // We place files under tmpDir/huggingface// so that the archive contains relative paths. + // We place files under tmpDir/huggingface/ so that the archive contains relative paths. if strings.TrimSpace(tmpDir) == "" { return "", nil, errors.Error("tmpDir is empty") } @@ -500,13 +725,14 @@ func (v *vllm) DownloadModel(ctx context.Context, userCred mcclient.TokenCredent return "", nil, errors.Error("modelName is empty") } - localDir := filepath.Join(tmpDir, "huggingface", filepath.FromSlash(modelName)) + modelBase := filepath.Base(modelName) + localDir := filepath.Join(tmpDir, "huggingface", modelBase) if err := os.MkdirAll(localDir, 0755); err != nil { return "", nil, errors.Wrap(err, "mkdir local model dir") } // If already downloaded, short-circuit (directory exists and non-empty). if entries, err := os.ReadDir(localDir); err == nil && len(entries) > 0 { - targetDir := path.Join(api.LLM_VLLM_MODELS_PATH, modelName) + targetDir := path.Join(api.LLM_VLLM_MODELS_PATH, modelBase) log.Infof("Model %s already exists in import dir %s", modelName, localDir) return modelName, []string{targetDir}, nil } @@ -553,74 +779,6 @@ func (v *vllm) DownloadModel(ctx context.Context, userCred mcclient.TokenCredent } } - targetDir := path.Join(api.LLM_VLLM_MODELS_PATH, modelName) + targetDir := path.Join(api.LLM_VLLM_MODELS_PATH, modelBase) return modelName, []string{targetDir}, nil } - -// vllmResolveResult is the result of resolving model path and estimating vLLM memory params in the container. -type vllmResolveResult struct { - ModelPath string - GpuUtil string - MaxModelLen int - MaxNumSeqs int -} - -// resolveModelAndParams runs one exec in the container to resolve the model path and estimate -// --gpu-memory-utilization, --max-model-len, --max-num-seqs. Returns (nil, nil) when no model is found. -func (v *vllm) resolveModelAndParams(ctx context.Context, containerId string, preferredPath string, tensorParallelSize int) (*vllmResolveResult, error) { - preferredEscaped := escapeShellSingleQuoted(preferredPath) - escapedScript := escapeShellSingleQuoted(strings.TrimSpace(api.LLM_VLLM_ESTIMATE_PARAMS_SCRIPT)) - defaultGpuUtil := strconv.FormatFloat(float64(api.LLM_VLLM_DEFAULT_GPU_MEMORY_UTIL), 'f', -1, 64) - cmd := fmt.Sprintf( - `mkdir -p %s; - preferred='%s'; - if [ -n "$preferred" ] && [ -d "$preferred" ]; then model="$preferred"; else model=$(ls -d %s/* 2>/dev/null | head -n 1); fi; - if [ -z "$model" ]; then echo "NO_MODEL"; exit 0; fi; - tp=%d; - vllm_out=$(python3 -c '%s' "$model" "$tp" 2>/dev/null) || true; - GPU_MEMORY_UTIL=%s; MAX_MODEL_LEN=%d; MAX_NUM_SEQS=%d; - [ -n "$vllm_out" ] && eval "$vllm_out"; - printf '%%s\n' "$model"; - printf 'GPU_MEMORY_UTIL=%%s MAX_MODEL_LEN=%%s MAX_NUM_SEQS=%%s\n' "$GPU_MEMORY_UTIL" "$MAX_MODEL_LEN" "$MAX_NUM_SEQS"`, - api.LLM_VLLM_MODELS_PATH, - preferredEscaped, - api.LLM_VLLM_MODELS_PATH, - tensorParallelSize, - escapedScript, - defaultGpuUtil, - api.LLM_VLLM_DEFAULT_MAX_MODEL_LEN, - api.LLM_VLLM_DEFAULT_MAX_NUM_SEQS, - ) - out, err := exec(ctx, containerId, cmd, 30) - if err != nil { - return nil, errors.Wrapf(err, "exec resolve model and params") - } - out = strings.TrimSpace(out) - if out == "NO_MODEL" { - return nil, nil - } - lines := strings.SplitN(out, "\n", 2) - if len(lines) < 2 { - return nil, errors.Errorf("vLLM resolve output missing params line: %s", out) - } - res := &vllmResolveResult{ - ModelPath: strings.TrimSpace(lines[0]), - GpuUtil: defaultGpuUtil, - MaxModelLen: api.LLM_VLLM_DEFAULT_MAX_MODEL_LEN, - MaxNumSeqs: api.LLM_VLLM_DEFAULT_MAX_NUM_SEQS, - } - for _, f := range strings.Fields(lines[1]) { - if val, ok := strings.CutPrefix(f, api.LLM_VLLM_RESOLVE_OUTPUT_PREFIX_GPU_UTIL); ok { - res.GpuUtil = val - } else if val, ok := strings.CutPrefix(f, api.LLM_VLLM_RESOLVE_OUTPUT_PREFIX_MAX_LEN); ok { - if n, e := strconv.Atoi(val); e == nil && n > 0 { - res.MaxModelLen = n - } - } else if val, ok := strings.CutPrefix(f, api.LLM_VLLM_RESOLVE_OUTPUT_PREFIX_MAX_NUM_SEQ); ok { - if n, e := strconv.Atoi(val); e == nil && n > 0 { - res.MaxNumSeqs = n - } - } - } - return res, nil -} diff --git a/pkg/mcclient/options/llm/llm.go b/pkg/mcclient/options/llm/llm.go index cd5b99a31b..7d12b42b3d 100644 --- a/pkg/mcclient/options/llm/llm.go +++ b/pkg/mcclient/options/llm/llm.go @@ -69,8 +69,9 @@ type LLMBaseCreateOptions struct { type LLMCreateOptions struct { LLMBaseCreateOptions - LLM_SKU_ID string `help:"llm sku id or name" json:"llm_sku_id"` - PreferredModel string `help:"vLLM preferred model dir name under models path (e.g. Qwen/Qwen2-7B)" json:"-"` + LLM_SKU_ID string `help:"llm sku id or name" json:"llm_sku_id"` + PreferredModel string `help:"vLLM preferred model dir name under models path (e.g. Qwen/Qwen2-7B)" json:"-"` + VllmArg []string `help:"vLLM args in format key=value; use key= for flags without values" json:"-"` } func (o *LLMCreateOptions) Params() (jsonutils.JSONObject, error) { @@ -88,12 +89,12 @@ func (o *LLMCreateOptions) Params() (jsonutils.JSONObject, error) { params.Add(jsonutils.Marshal(nets), "nets") } - if o.PreferredModel != "" { - spec := &api.LLMSpec{ - Ollama: nil, - Vllm: &api.LLMSpecVllm{PreferredModel: o.PreferredModel}, - Dify: nil, - } + vllmSpec, err := newVLLMSpecFromArgs(o.PreferredModel, o.VllmArg) + if err != nil { + return nil, err + } + if vllmSpec != nil { + spec := &api.LLMSpec{Ollama: nil, Vllm: vllmSpec, Dify: nil} params.Set("llm_spec", jsonutils.Marshal(spec)) } @@ -107,7 +108,8 @@ func (o *LLMCreateOptions) GetCountParam() int { type LLMUpdateOptions struct { options.BaseIdOptions - PreferredModel string `help:"vLLM preferred model dir name under models path (e.g. Qwen/Qwen2-7B); takes effect after pod recreate" json:"-"` + PreferredModel string `help:"vLLM preferred model dir name under models path (e.g. Qwen/Qwen2-7B); takes effect after pod recreate" json:"-"` + VllmArg []string `help:"vLLM args in format key=value; use key= for flags without values" json:"-"` } func (o *LLMUpdateOptions) GetId() string { @@ -119,12 +121,12 @@ func (o *LLMUpdateOptions) Params() (jsonutils.JSONObject, error) { if err != nil { return nil, err } - if o.PreferredModel != "" { - spec := &api.LLMSpec{ - Ollama: nil, - Vllm: &api.LLMSpecVllm{PreferredModel: o.PreferredModel}, - Dify: nil, - } + vllmSpec, err := newVLLMSpecFromArgs(o.PreferredModel, o.VllmArg) + if err != nil { + return nil, err + } + if vllmSpec != nil { + spec := &api.LLMSpec{Ollama: nil, Vllm: vllmSpec, Dify: nil} dict.Set("llm_spec", jsonutils.Marshal(spec)) } return dict, nil diff --git a/pkg/mcclient/options/llm/llm_sku.go b/pkg/mcclient/options/llm/llm_sku.go index 0a2637d906..33d5bf18d0 100644 --- a/pkg/mcclient/options/llm/llm_sku.go +++ b/pkg/mcclient/options/llm/llm_sku.go @@ -33,7 +33,8 @@ type LLMSkuCreateOptions struct { LLM_IMAGE_ID string `json:"llm_image_id"` LLM_TYPE string `json:"llm_type" choices:"ollama|vllm|comfyui"` - PreferredModel string `help:"preferred model (vllm only), sets llm_spec.vllm.preferred_model" json:"-"` + PreferredModel string `help:"preferred model (vllm only), sets llm_spec.vllm.preferred_model" json:"-"` + VllmArg []string `help:"vLLM args in format key=value; use key= for flags without values" json:"-"` } func (o *LLMSkuCreateOptions) Params() (jsonutils.JSONObject, error) { @@ -44,13 +45,15 @@ func (o *LLMSkuCreateOptions) Params() (jsonutils.JSONObject, error) { return nil, err } fetchMountedModels(o.MountedModels, dict) - if o.LLM_TYPE == string(api.LLM_CONTAINER_VLLM) && len(o.PreferredModel) > 0 { + vllmSpec, err := newVLLMSpecFromArgs(o.PreferredModel, o.VllmArg) + if err != nil { + return nil, err + } + if o.LLM_TYPE == string(api.LLM_CONTAINER_VLLM) && vllmSpec != nil { spec := &api.LLMSpec{ Ollama: nil, - Vllm: &api.LLMSpecVllm{ - PreferredModel: o.PreferredModel, - }, - Dify: nil, + Vllm: vllmSpec, + Dify: nil, } dict.Set("llm_spec", jsonutils.Marshal(spec)) } @@ -77,7 +80,8 @@ type LLMSkuUpdateOptions struct { // For ollama/vllm; backend merges into LLMSpec. Use dify-sku update for dify type. LlmImageId string `json:"llm_image_id"` - PreferredModel string `help:"preferred model (vllm only), sets llm_spec.vllm.preferred_model" json:"-"` + PreferredModel string `help:"preferred model (vllm only), sets llm_spec.vllm.preferred_model" json:"-"` + VllmArg []string `help:"vLLM args in format key=value; use key= for flags without values" json:"-"` } func (o *LLMSkuUpdateOptions) GetId() string { @@ -92,13 +96,15 @@ func (o *LLMSkuUpdateOptions) Params() (jsonutils.JSONObject, error) { return nil, err } fetchMountedModels(o.MountedModels, dict) - if len(o.PreferredModel) > 0 { + vllmSpec, err := newVLLMSpecFromArgs(o.PreferredModel, o.VllmArg) + if err != nil { + return nil, err + } + if vllmSpec != nil { spec := &api.LLMSpec{ Ollama: nil, - Vllm: &api.LLMSpecVllm{ - PreferredModel: o.PreferredModel, - }, - Dify: nil, + Vllm: vllmSpec, + Dify: nil, } dict.Set("llm_spec", jsonutils.Marshal(spec)) } diff --git a/pkg/mcclient/options/llm/vllm_args.go b/pkg/mcclient/options/llm/vllm_args.go new file mode 100644 index 0000000000..966c686284 --- /dev/null +++ b/pkg/mcclient/options/llm/vllm_args.go @@ -0,0 +1,44 @@ +package llm + +import ( + "fmt" + "strings" + + api "yunion.io/x/onecloud/pkg/apis/llm" +) + +func parseVLLMCustomizedArgs(items []string) ([]*api.VllmCustomizedArg, error) { + if len(items) == 0 { + return nil, nil + } + out := make([]*api.VllmCustomizedArg, 0, len(items)) + for _, item := range items { + idx := strings.Index(item, "=") + if idx <= 0 { + return nil, fmt.Errorf("invalid vllm arg %q, expected key=value", item) + } + key := strings.TrimSpace(item[:idx]) + if key == "" { + return nil, fmt.Errorf("invalid vllm arg %q, empty key", item) + } + out = append(out, &api.VllmCustomizedArg{ + Key: key, + Value: item[idx+1:], + }) + } + return out, nil +} + +func newVLLMSpecFromArgs(preferredModel string, items []string) (*api.LLMSpecVllm, error) { + customizedArgs, err := parseVLLMCustomizedArgs(items) + if err != nil { + return nil, err + } + if preferredModel == "" && len(customizedArgs) == 0 { + return nil, nil + } + return &api.LLMSpecVllm{ + PreferredModel: preferredModel, + CustomizedArgs: customizedArgs, + }, nil +}