diff --git a/.gitignore b/.gitignore index 222caa50b..a76454c79 100644 --- a/.gitignore +++ b/.gitignore @@ -158,3 +158,4 @@ ast_index_file.py test_cookbook/ /test*.py swanlog/ +tests/server/config/_generated_e2e.yaml diff --git a/src/twinkle/model/megatron/megatron.py b/src/twinkle/model/megatron/megatron.py index 89d0a171d..805dd4f80 100644 --- a/src/twinkle/model/megatron/megatron.py +++ b/src/twinkle/model/megatron/megatron.py @@ -650,8 +650,11 @@ def zero_grad(self, **kwargs): # For DDP-wrapped models, ALWAYS zero the gradient buffer # This is essential because Megatron's forward_backward_func uses # the buffer's state to track gradient accumulation - if self._is_model_ddp_wrapped() and hasattr(self.model, 'zero_grad_buffer'): - self.model.zero_grad_buffer() + # (self.model is a list of chunks; zero_grad_buffer lives on each chunk) + if self._is_model_ddp_wrapped(): + for model_chunk in self.model: + if hasattr(model_chunk, 'zero_grad_buffer'): + model_chunk.zero_grad_buffer() if not optimizer_config.do_grad_sync(kwargs.pop('gradient_accumulation_steps', None)): return @@ -988,6 +991,32 @@ def load(self, name: str, output_dir: Optional[str] = None, **kwargs): if dist.is_initialized(): dist.barrier() + @remote_function(dispatch='all') + def reload_initial_weights(self, **kwargs): + """Reload the base model weights from the original model path. + + Used by full-parameter server deployments after a tenant releases the + (exclusive) model, so the next tenant starts from clean pretrained + weights instead of the previous tenant's trained weights. + """ + bridge = self.strategy.bridge + bridge.load_weights( + self.strategy.unwrap_model(self.model), + self._model_path, + peft_format=False, + ) + # Drop any leftover gradients from the previous tenant so they cannot + # leak into the next tenant's first optimizer step. + if self._is_model_ddp_wrapped(): + for model_chunk in self.model: + if hasattr(model_chunk, 'zero_grad_buffer'): + model_chunk.zero_grad_buffer() + for _model in self.strategy.unwrap_model(self.model): + for param in _model.parameters(): + param.grad = None + if dist.is_initialized(): + dist.barrier() + @remote_function(dispatch='all') def resume_from_checkpoint(self, checkpoint_dir, *, resume_only_model=False, **kwargs): adapter_name = kwargs.pop('adapter_name', self._get_default_group()) diff --git a/src/twinkle/model/multi_lora.py b/src/twinkle/model/multi_lora.py index 43cd6108d..f310306c6 100644 --- a/src/twinkle/model/multi_lora.py +++ b/src/twinkle/model/multi_lora.py @@ -479,9 +479,12 @@ def patch(self, target_modules='all-linear', *args, **kwargs): - module_device = getattr(module, 'device', None) + # ``module`` may be a list of model chunks (Megatron); probe the first + # chunk for the device in that case. + first_module = module[0] if isinstance(module, (list, tuple)) else module + module_device = getattr(first_module, 'device', None) if module_device is None: - module_device = next(module.parameters())[1].device + module_device = next(first_module.parameters()).device low_cpu_mem_usage = module_device.type == 'meta' for i in range(self.max_loras): diff --git a/src/twinkle/model/multi_lora_target_parameters.py b/src/twinkle/model/multi_lora_target_parameters.py index 0eb33eea7..8955917da 100644 --- a/src/twinkle/model/multi_lora_target_parameters.py +++ b/src/twinkle/model/multi_lora_target_parameters.py @@ -369,14 +369,21 @@ def named_slot_parameters(self, tenant_adapter_name: str) -> Iterator[tuple[str, yield from wrapper.named_slot_parameters(slot_name) def get_state_dict(self, tenant_adapter_name: str) -> dict[str, torch.Tensor]: - slot_name = self.tenant_to_slot[tenant_adapter_name] + # Tenants without target_parameters never acquire a slot here; return + # an empty dict so plain-LoRA save paths do not fail. + slot_name = self.tenant_to_slot.get(tenant_adapter_name) + if slot_name is None: + return {} state_dict = {} for wrapper in self.wrappers: state_dict.update(wrapper.get_state_dict(slot_name)) return state_dict def set_state_dict(self, tenant_adapter_name: str, state_dict: dict[str, torch.Tensor]) -> set[str]: - slot_name = self.tenant_to_slot[tenant_adapter_name] + # Same tolerance as get_state_dict for plain-LoRA load paths. + slot_name = self.tenant_to_slot.get(tenant_adapter_name) + if slot_name is None: + return set() consumed_keys = set() for wrapper in self.wrappers: consumed_keys.update(wrapper.set_state_dict(slot_name, state_dict)) diff --git a/src/twinkle/model/transformers/strategy/accelerate.py b/src/twinkle/model/transformers/strategy/accelerate.py index 3bf627e9d..f74689d11 100644 --- a/src/twinkle/model/transformers/strategy/accelerate.py +++ b/src/twinkle/model/transformers/strategy/accelerate.py @@ -203,6 +203,20 @@ def get_full_state_dict(self, model) -> dict: del local return state_dict + def load_full_state_dict(self, model, state_dict) -> None: + """Load a full (non-sharded) state dict into the model in-place. + + Used by full-parameter training to (re)load base weights, e.g. after a + tenant releases an exclusive full-parameter deployment. + """ + fsdp_plugin = self._get_fsdp_plugin() + if fsdp_plugin is not None and fsdp_plugin.fsdp_version == 2: + from torch.distributed.checkpoint.state_dict import set_model_state_dict + set_model_state_dict(model, state_dict, options=self._prepare_fsdp2_sd_options()) + return + unwrapped = self.unwrap_model(model) + unwrapped.load_state_dict(state_dict, strict=False) + def get_adapter_state_dict(self, model, adapter_name: str) -> dict: """Collect only LoRA adapter parameters.""" from twinkle.utils import torch_util diff --git a/src/twinkle/model/transformers/strategy/native_fsdp.py b/src/twinkle/model/transformers/strategy/native_fsdp.py index f57877bb0..925c84c82 100644 --- a/src/twinkle/model/transformers/strategy/native_fsdp.py +++ b/src/twinkle/model/transformers/strategy/native_fsdp.py @@ -345,6 +345,24 @@ def get_full_state_dict(self, model) -> dict: return state_dict + def load_full_state_dict(self, model, state_dict) -> None: + """Load a full (non-sharded) state dict into the (possibly sharded) model. + + Used by full-parameter training to (re)load base weights. Uses + ``set_model_state_dict`` when the model is distributed (FSDP2), else a + plain in-place ``load_state_dict``. + """ + if self.device_mesh is not None: + from torch.distributed.checkpoint.state_dict import StateDictOptions, set_model_state_dict + set_model_state_dict( + model, + state_dict, + options=StateDictOptions(full_state_dict=True, broadcast_from_rank0=True), + ) + return + unwrapped = self.unwrap_model(model) + unwrapped.load_state_dict(state_dict, strict=False) + def get_adapter_state_dict(self, model, adapter_name: str) -> dict: """Collect only LoRA adapter parameters, with EP-aware all-gather.""" unwrapped = self.unwrap_model(model) diff --git a/src/twinkle/model/transformers/transformers.py b/src/twinkle/model/transformers/transformers.py index 017a515b3..db853ebbf 100644 --- a/src/twinkle/model/transformers/transformers.py +++ b/src/twinkle/model/transformers/transformers.py @@ -146,6 +146,44 @@ def accumulate_metrics(self, is_training): DEFAULT_WEIGHT_DECAY = 0.01 +def _read_hf_state_dict(checkpoint_dir: str) -> Dict[str, torch.Tensor]: + """Read a full HuggingFace checkpoint directory into a CPU state dict. + + Supports both single-file and sharded ``safetensors`` layouts, falling back + to ``pytorch_model.bin`` variants. Returns tensors on CPU. + """ + import json + + state_dict: Dict[str, torch.Tensor] = {} + st_index = os.path.join(checkpoint_dir, 'model.safetensors.index.json') + st_single = os.path.join(checkpoint_dir, 'model.safetensors') + bin_index = os.path.join(checkpoint_dir, 'pytorch_model.bin.index.json') + bin_single = os.path.join(checkpoint_dir, 'pytorch_model.bin') + + if os.path.exists(st_index) or os.path.exists(st_single): + from safetensors.torch import load_file + if os.path.exists(st_index): + with open(st_index) as f: + weight_map = json.load(f)['weight_map'] + shards = sorted(set(weight_map.values())) + for shard in shards: + state_dict.update(load_file(os.path.join(checkpoint_dir, shard), device='cpu')) + else: + state_dict.update(load_file(st_single, device='cpu')) + elif os.path.exists(bin_index) or os.path.exists(bin_single): + if os.path.exists(bin_index): + with open(bin_index) as f: + weight_map = json.load(f)['weight_map'] + shards = sorted(set(weight_map.values())) + for shard in shards: + state_dict.update(torch.load(os.path.join(checkpoint_dir, shard), map_location='cpu', weights_only=True)) + else: + state_dict.update(torch.load(bin_single, map_location='cpu', weights_only=True)) + else: + raise FileNotFoundError(f'No safetensors/bin weights found in {checkpoint_dir}') + return state_dict + + @remote_class() class TransformersModel(TwinkleModel, PreTrainedModel, CheckpointEngineMixin): """The transformers model wrapper. @@ -856,7 +894,11 @@ def step(self, **kwargs): optim_params = kwargs.pop('optim_params', {}) if optim_params: - assert isinstance(optimizer, (AdamW, Adam)) + # After _lazy_wrap_model the optimizer may be wrapped (e.g. + # accelerate's AcceleratedOptimizer); check the inner instance. + inner_optimizer = getattr(optimizer, 'optimizer', optimizer) + assert isinstance(inner_optimizer, (AdamW, Adam)), \ + f'optim_params is only supported for Adam/AdamW, got {type(inner_optimizer).__name__}' for group in optimizer.param_groups: group['lr'] = optim_params['lr'] if group['weight_decay'] > 0.0 and optim_params.get('weight_decay', None) is not None: @@ -1137,11 +1179,31 @@ def load(self, name: str, output_dir: Optional[str] = None, **kwargs): adapter_weights = load_peft_weights(checkpoint_dir, device='cpu') self.strategy.load_peft_weights(model, adapter_weights, adapter_name) else: - raise NotImplementedError + # Full-parameter model: load a plain HF checkpoint in-place. + state_dict = _read_hf_state_dict(checkpoint_dir) + self.strategy.load_full_state_dict(self.model, state_dict) if load_optimizer: self._load_optimizer(checkpoint_dir, adapter_name=adapter_name) + @remote_function() + def reload_initial_weights(self, **kwargs): + """Reload the base model weights from ``self.model_id``. + + Used by full-parameter server deployments after a tenant releases the + (exclusive) model, so the next tenant starts from clean pretrained + weights instead of the previous tenant's trained weights. + """ + if not self.model_id: + logger.warning('reload_initial_weights skipped: model_id is not set (blank model).') + return + state_dict = _read_hf_state_dict(self.model_id) + self.strategy.load_full_state_dict(self.model, state_dict) + # Drop any leftover gradients from the previous tenant so they cannot + # leak into the next tenant's first optimizer step. + for param in self.model.parameters(): + param.grad = None + def _load_optimizer(self, checkpoint_dir, **kwargs): adapter_name = kwargs.pop('adapter_name', _default_adapter_name) strict = kwargs.pop('strict', False) diff --git a/src/twinkle/sampler/vllm_sampler/vllm_sampler.py b/src/twinkle/sampler/vllm_sampler/vllm_sampler.py index 3c7b2f686..f12d561ee 100644 --- a/src/twinkle/sampler/vllm_sampler/vllm_sampler.py +++ b/src/twinkle/sampler/vllm_sampler/vllm_sampler.py @@ -486,6 +486,52 @@ async def _receive_and_load(): self._run_in_loop(_receive_and_load()) + @remote_function(dispatch='all', collect='first', lazy_collect=False) + def load_full_weights_from_path(self, path: str) -> int: + """Load a full (non-LoRA) HF checkpoint into the engine's base model. + + Used by full-parameter training: the saved checkpoint is a plain HF + directory (no ``adapter_config.json``), so it replaces the sampler's + base weights instead of being loaded as a LoRA adapter. Idempotent: + repeated calls with the same resolved path are skipped. + + Returns: + 1 if weights were (re)loaded, 0 if the path was already loaded. + """ + import glob + import json + import os + + resolved = HubOperation.download_model(model_id_or_path=path) + if getattr(self, '_loaded_full_weights_path', None) == resolved: + return 0 + + from safetensors import safe_open + + def _weight_iter(): + index = os.path.join(resolved, 'model.safetensors.index.json') + if os.path.exists(index): + with open(index) as f: + shards = sorted(set(json.load(f)['weight_map'].values())) + files = [os.path.join(resolved, s) for s in shards] + else: + files = sorted(glob.glob(os.path.join(resolved, '*.safetensors'))) + for fp in files: + with safe_open(fp, framework='pt', device='cpu') as f: + for key in f.keys(): + yield key, f.get_tensor(key) + + async def _load(): + await self.engine.update_weights(_weight_iter(), peft_config=None, base_sync_done=False) + # A full base-model swap invalidates any previously synced LoRA. + self.engine.invalidate_synced_lora() + + logger.info(f'Loading full-parameter weights into sampler base model from {resolved}') + self._run_in_loop(_load()) + self._loaded_full_weights_path = resolved + self.reset_prefix_cache() + return 1 + @remote_function(dispatch='all', collect='first', lazy_collect=False) def shutdown(self): """Gracefully shutdown the vLLM engine and background event loop. diff --git a/src/twinkle/server/config/application_spec.py b/src/twinkle/server/config/application_spec.py index 51015245c..833ae2b49 100644 --- a/src/twinkle/server/config/application_spec.py +++ b/src/twinkle/server/config/application_spec.py @@ -54,6 +54,7 @@ class ModelArgs(_ArgsBase): device_group: dict[str, Any] device_mesh: dict[str, Any] backend: Literal['mock', 'transformers', 'megatron'] + train_mode: Literal['lora', 'full'] = 'lora' adapter_config: dict[str, Any] | None = None queue_config: TaskQueueConfig = Field(default_factory=TaskQueueConfig) max_loras: int = 5 @@ -169,3 +170,24 @@ def _coerce_args_to_schema(cls, data: Any) -> Any: # ``schema.model_validate`` rejects a non-dict itself with a clean # error, so no separate non-dict guard is needed here. return {**data, 'args': schema.model_validate(raw_args)} + + @model_validator(mode='after') + def _validate_full_mode_single_replica(self) -> ApplicationSpec: + """``train_mode: full`` requires exactly one replica. + + The exclusive-tenant lock lives in per-replica memory + (``ModelManagement._resource_records``), so more than one replica would + silently allow one full-parameter tenant per replica, each rewriting + its own copy of the base weights. Reject that at config-load time. + """ + if not (isinstance(self.args, ModelArgs) and self.args.train_mode == 'full'): + return self + for dep in self.deployments: + num_replicas = dep.get('num_replicas') + max_replicas = (dep.get('autoscaling_config') or {}).get('max_replicas') + if (num_replicas or 1) > 1 or (max_replicas or 1) > 1: + raise ValueError( + f"Application '{self.name}': train_mode='full' is an exclusive single-tenant " + 'mode and requires a single replica; set num_replicas/autoscaling_config.' + 'max_replicas to 1.') + return self diff --git a/src/twinkle/server/exceptions.py b/src/twinkle/server/exceptions.py index f57826b0b..dc58ad81d 100644 --- a/src/twinkle/server/exceptions.py +++ b/src/twinkle/server/exceptions.py @@ -51,3 +51,17 @@ class ConfigParseError(TwinkleServerError): class ResourceExhaustedError(TwinkleServerError): """Resource exhausted — queue full, insufficient memory, connection pool exhausted, etc.""" pass + + +class FullModeBusyError(TwinkleServerError): + """A full-parameter (exclusive) model deployment already has a holder. + + Full-parameter training rewrites the shared base-model weights, so a single + deployment can only host one training task at a time. + """ + + def __init__(self, current_holder: str) -> None: + self.current_holder = current_holder + super().__init__('This deployment runs in full-parameter (exclusive) mode and is already ' + f'held by another training task ({current_holder}). Only one full-parameter ' + 'training task is allowed at a time; retry after it is released.') diff --git a/src/twinkle/server/model/app.py b/src/twinkle/server/model/app.py index bae7b2a93..ddb4ddf7e 100644 --- a/src/twinkle/server/model/app.py +++ b/src/twinkle/server/model/app.py @@ -15,6 +15,7 @@ from twinkle import DeviceGroup from twinkle.server.common.router import StickyLoraRequestRouter from twinkle.server.deployment import LazyCleanupMixin, bind_deployment, build_deployment_app, init_twinkle_runtime +from twinkle.server.exceptions import FullModeBusyError from twinkle.server.state import ServerState, get_server_state from twinkle.server.utils import wrap_builder_with_device_group_env from twinkle.server.utils.backend_dispatch import BackendSelector @@ -27,22 +28,43 @@ logger = get_logger() +# ``FullModeBusyError`` lives in ``twinkle.server.exceptions``; re-exported here +# for backwards compatibility with callers importing it from this module. +__all__ = ['FullModeBusyError', 'ModelManagement', 'build_model_app'] + + +# Ctor kwargs consumed by the MultiLora wrappers' signatures but unknown to the +# plain (full-parameter) model classes, where **kwargs flows into HF +# ``from_pretrained``. Dropped when constructing a full-mode backend. +_LORA_ONLY_CTOR_KWARGS = ('max_loras', 'max_r', 'max_length', 'target_modules', 'lora_config') + def _make_mock_model(kw: dict[str, Any]) -> Any: from .backends.mock_model import TwinkleCompatMockModel + kw.pop('train_mode', None) # mock backend behaves the same for lora/full return TwinkleCompatMockModel(**kw) def _make_transformers_model(kw: dict[str, Any]) -> Any: + train_mode = kw.pop('train_mode', 'lora') + if train_mode == 'full': + from .backends.transformers_model import TwinkleCompatFullTransformersModel + for key in _LORA_ONLY_CTOR_KWARGS: + kw.pop(key, None) + return TwinkleCompatFullTransformersModel(**kw) from .backends.transformers_model import TwinkleCompatTransformersModel - return TwinkleCompatTransformersModel(**kw) def _make_megatron_model(kw: dict[str, Any]) -> Any: + train_mode = kw.pop('train_mode', 'lora') + if train_mode == 'full': + from .backends.megatron_model import TwinkleCompatFullMegatronModel + for key in _LORA_ONLY_CTOR_KWARGS: + kw.pop(key, None) + return TwinkleCompatFullMegatronModel(**kw) from .backends.megatron_model import TwinkleCompatMegatronModel - return TwinkleCompatMegatronModel(**kw) @@ -88,11 +110,17 @@ def __init__(self, self.replica_id = serve.get_replica_context().replica_id.unique_id self.max_loras = kwargs.get('max_loras', 5) self.base_model = model_id + # ``train_mode`` selects LoRA multi-tenant (default) or exclusive + # full-parameter training. Popped here so the underlying model + # constructor never receives it; re-injected into ctor_kwargs so the + # backend builder can pick the full vs LoRA wrapper class. + self.train_mode = kwargs.pop('train_mode', 'lora') ctor_kwargs: dict[str, Any] = { 'model_id': model_id, 'remote_group': self.device_group.name, 'instance_id': self.replica_id, + 'train_mode': self.train_mode, **kwargs, } if self.device_mesh is not None: @@ -165,10 +193,46 @@ def check_model_health(self) -> dict: async def _cleanup_adapter(self, adapter_name: str) -> None: if self.get_resource_info(adapter_name): self.clear_resource_state(adapter_name) - self.model.remove_adapter(adapter_name) + if self.train_mode == 'full': + # No PEFT adapter to remove; restore clean base weights so the + # next tenant does not inherit this tenant's trained weights. + self.model.reload_initial_weights() + else: + self.model.remove_adapter(adapter_name) self.unregister_resource(adapter_name) await self.state.unload_model(adapter_name) + @property + def is_full_mode(self) -> bool: + """True when this deployment runs exclusive full-parameter training.""" + return self.train_mode == 'full' + + def resolve_model_adapter_name(self, adapter_name: str | None) -> str | None: + """Map a per-tenant adapter name to the name the model expects. + + In full-parameter mode every request targets the empty-string default + optimizer group; in LoRA mode the per-tenant adapter name is used as-is. + """ + if self.is_full_mode: + return '' + return adapter_name + + def assert_full_mode_available(self, adapter_name: str | None = None) -> None: + """In full mode, ensure no other tenant currently holds this deployment. + + Call before registering any state so a rejected request leaves nothing + behind. ``adapter_name`` (when given) is excluded from the check so a + tenant re-issuing against its own resource is not blocked. + + Raises: + FullModeBusyError: another (non-expiring) tenant already holds it. + """ + if not self.is_full_mode: + return + for rid, info in self._resource_records.items(): + if rid != adapter_name and not info.get('expiring'): + raise FullModeBusyError(rid) + async def _on_adapter_expired(self, adapter_name: str) -> None: self.fail_pending_tasks_for_model(adapter_name, reason='Adapter expired') await self._cleanup_adapter(adapter_name) diff --git a/src/twinkle/server/model/backends/megatron_model.py b/src/twinkle/server/model/backends/megatron_model.py index 21b10ee73..f945f1efc 100644 --- a/src/twinkle/server/model/backends/megatron_model.py +++ b/src/twinkle/server/model/backends/megatron_model.py @@ -1,6 +1,13 @@ # Copyright (c) ModelScope Contributors. All rights reserved. """ -Megatron backend model for the unified model deployment. +Megatron backend models for the unified model deployment. + +Contains the shared Tinker/Twinkle compatibility mixin and two concrete +wrappers: +- TwinkleCompatMegatronModel: multi-LoRA training (default), wraps + MultiLoraMegatronModel. +- TwinkleCompatFullMegatronModel: exclusive full-parameter training, wraps the + plain MegatronModel (no PEFT/MultiLora). """ import torch from tinker import types @@ -9,18 +16,20 @@ from twinkle import remote_class, remote_function from twinkle.data_format import InputFeature, Trajectory from twinkle.infra import collect_tensor_dict -from twinkle.model.megatron import MultiLoraMegatronModel +from twinkle.model.megatron import MegatronModel, MultiLoraMegatronModel from twinkle.server.common.datum import datum_to_input_feature, extract_rl_features_for_loss from twinkle.server.model.backends.common import (TwinkleCompatModelBase, clean_metrics, collect_forward_backward_results, to_cpu_safe_output) from twinkle.utils.nccl_safe import nccl_safe_megatron -@remote_class(execute='all') -class TwinkleCompatMegatronModel(MultiLoraMegatronModel, TwinkleCompatModelBase): - """Compatibility wrapper around MultiLoraMegatronModel for Twinkle/Tinker. +class _MegatronTinkerCompatMixin(TwinkleCompatModelBase): + """Tinker/Twinkle-compat methods shared by LoRA and full-parameter wrappers. - Moved from tinker/common/megatron_model.py — logic unchanged. + Method bodies only call ``super()`` so they work against either + ``MultiLoraMegatronModel`` (LoRA) or ``MegatronModel`` (full) as the next + class in the MRO. For full-parameter training the ``adapter_name`` is the + empty-string default optimizer group. """ @remote_function(dispatch='slice_dp', collect=collect_forward_backward_results, sync=True) @@ -65,6 +74,9 @@ def tinker_step(self, *, adam_params: types.AdamParams, **kwargs): if hasattr(opt, 'chained_optimizers'): for chained_opt in opt.chained_optimizers: if hasattr(chained_opt, 'config'): + # ``config.*`` only seeds param_groups at optimizer-construction + # time; keep it in sync for clip_grad (read from config at step + # time) and for _get_lr fallbacks. chained_opt.config.lr = adam_params.learning_rate chained_opt.config.adam_eps = adam_params.eps chained_opt.config.adam_beta1 = adam_params.beta1 @@ -73,7 +85,16 @@ def tinker_step(self, *, adam_params: types.AdamParams, **kwargs): if adam_params.grad_clip_norm > 0: chained_opt.config.clip_grad = adam_params.grad_clip_norm - super().step(**kwargs) + # The inner Adam reads lr/eps/betas from ``param_groups`` at step time, so + # AdamParams must be applied there via optim_params (mirrors the + # transformers backend); updating ``config.lr`` alone has no effect. + optim_params = { + 'lr': adam_params.learning_rate, + 'eps': adam_params.eps, + 'betas': (adam_params.beta1, adam_params.beta2), + 'weight_decay': adam_params.weight_decay, + } + super().step(optim_params=optim_params, **kwargs) super().zero_grad(**kwargs) @remote_function(collect='last_pp_first', lazy_collect=False) @@ -131,3 +152,18 @@ def forward_backward(self, *, inputs: InputFeature | list[InputFeature] | Trajec def ping(self) -> bool: """Lightweight liveness probe for watchdog health checks.""" return True + + +@remote_class(execute='all') +class TwinkleCompatMegatronModel(_MegatronTinkerCompatMixin, MultiLoraMegatronModel): + """Compatibility wrapper around MultiLoraMegatronModel for Twinkle/Tinker.""" + + +@remote_class(execute='all') +class TwinkleCompatFullMegatronModel(_MegatronTinkerCompatMixin, MegatronModel): + """Full-parameter (non-LoRA) wrapper around the plain MegatronModel. + + Used by exclusive full-parameter server deployments. All training operations + target the empty-string default optimizer group created by + ``MegatronModel.__init__``. + """ diff --git a/src/twinkle/server/model/backends/mock_model.py b/src/twinkle/server/model/backends/mock_model.py index bc0fc06f2..7bf9866cc 100644 --- a/src/twinkle/server/model/backends/mock_model.py +++ b/src/twinkle/server/model/backends/mock_model.py @@ -192,6 +192,11 @@ def save(self, name: str = 'latest', output_dir: str | None = None, **kwargs: An def load(self, *args: Any, **kwargs: Any) -> None: return None + @remote_function() + def reload_initial_weights(self, *args: Any, **kwargs: Any) -> None: + """No-op reload for the mock full-parameter path.""" + return None + @remote_function() def resume_from_checkpoint(self, *args: Any, **kwargs: Any) -> dict[str, Any]: return {'status': 'ok', 'progress': {}} diff --git a/src/twinkle/server/model/backends/transformers_model.py b/src/twinkle/server/model/backends/transformers_model.py index e7677b619..1a382805b 100644 --- a/src/twinkle/server/model/backends/transformers_model.py +++ b/src/twinkle/server/model/backends/transformers_model.py @@ -2,9 +2,15 @@ """ Backend model implementations for the unified model deployment. -Contains one unified class: -- TwinkleCompatTransformersModel: handles both tinker (Datum-based I/O) via /tinker/* - endpoints and twinkle-native (InputFeature/Trajectory-based I/O) via /twinkle/* endpoints. +Contains the shared Tinker/Twinkle compatibility mixin and two concrete +wrappers: +- TwinkleCompatTransformersModel: multi-LoRA training (default), wraps + MultiLoraTransformersModel. +- TwinkleCompatFullTransformersModel: exclusive full-parameter training, wraps + the plain TransformersModel (no PEFT/MultiLora). + +Both handle tinker (Datum-based I/O) via /tinker/* endpoints and twinkle-native +(InputFeature/Trajectory-based I/O) via /twinkle/* endpoints. """ from tinker import types from typing import List, Union @@ -13,19 +19,20 @@ from twinkle.data_format import InputFeature, Trajectory from twinkle.infra import collect_tensor_dict from twinkle.model import MultiLoraTransformersModel +from twinkle.model.transformers import TransformersModel from twinkle.server.common.datum import datum_to_input_feature, extract_rl_features_for_loss from twinkle.server.model.backends.common import (TwinkleCompatModelBase, clean_metrics, collect_forward_backward_results, to_cpu_safe_output) from twinkle.utils.nccl_safe import nccl_safe -@remote_class() -class TwinkleCompatTransformersModel(MultiLoraTransformersModel, TwinkleCompatModelBase): - """Unified wrapper around MultiLoraTransformersModel. +class _TransformersTinkerCompatMixin(TwinkleCompatModelBase): + """Tinker/Twinkle-compat methods shared by LoRA and full-parameter wrappers. - Handles both: - - Tinker-compat I/O (Datum / TensorData) via /tinker/* endpoints. - - Twinkle-native I/O (InputFeature / Trajectory) via /twinkle/* endpoints. + Method bodies only call ``super()``, so they work against either + ``MultiLoraTransformersModel`` (LoRA) or ``TransformersModel`` (full) as the + next class in the MRO. For full-parameter training the ``adapter_name`` is + the empty-string default optimizer group. """ # ------------------------------------------------------------------ @@ -110,3 +117,18 @@ def forward_backward(self, *, inputs: InputFeature | list[InputFeature] | Trajec def ping(self) -> bool: """Lightweight liveness probe for watchdog health checks.""" return True + + +@remote_class() +class TwinkleCompatTransformersModel(_TransformersTinkerCompatMixin, MultiLoraTransformersModel): + """Unified multi-LoRA wrapper around MultiLoraTransformersModel.""" + + +@remote_class() +class TwinkleCompatFullTransformersModel(_TransformersTinkerCompatMixin, TransformersModel): + """Full-parameter (non-LoRA) wrapper around the plain TransformersModel. + + Used by exclusive full-parameter server deployments. All training operations + target the empty-string default optimizer group created by + ``TransformersModel.__init__``. + """ diff --git a/src/twinkle/server/model/tinker_handlers.py b/src/twinkle/server/model/tinker_handlers.py index 1ddc87489..3156fbefc 100644 --- a/src/twinkle/server/model/tinker_handlers.py +++ b/src/twinkle/server/model/tinker_handlers.py @@ -18,6 +18,7 @@ from .app import ModelManagement from twinkle.server.checkpoint import create_checkpoint_manager, create_training_run_manager +from twinkle.server.exceptions import FullModeBusyError from twinkle.server.utils import get_template_for_model from twinkle.utils.logger import get_logger @@ -41,17 +42,44 @@ async def create_model( async def _create_adapter(): _model_id = None + _adapter_name = None try: + # Validate lora_config against the deployment's train_mode up front. + if self.is_full_mode and body.lora_config: + return types.RequestFailedResponse( + error='This deployment runs in full-parameter (exclusive) mode; do not pass ' + 'lora_config. Use create_full_training_client (or omit lora_config).', + category=types.RequestErrorCategory.User, + ) + if (not self.is_full_mode) and (not body.lora_config): + return types.RequestFailedResponse( + error='This deployment runs in LoRA mode; lora_config is required. ' + 'Use create_lora_training_client.', + category=types.RequestErrorCategory.User, + ) + # Exclusive full-parameter training: reject early (before touching + # state) if another tenant already holds the deployment. + if self.is_full_mode: + self.assert_full_mode_available() + _model_id = await self.state.register_model( body.model_dump(), token=token, replica_id=self.replica_id, session_id=body.session_id) - if body.lora_config: + adapter_name = self.get_adapter_name(adapter_name=_model_id) + _adapter_name = adapter_name + model_adapter = self.resolve_model_adapter_name(adapter_name) + # Select template based on model type + template = get_template_for_model(self.base_model) + if self.is_full_mode: + self.register_resource(adapter_name, token, session_id=body.session_id) + self.model.set_template(template, adapter_name=model_adapter, model_id=self.base_model) + self.model.set_processor('InputProcessor', adapter_name=model_adapter) + self.model.set_optimizer('Adam', adapter_name=model_adapter) + self.set_resource_state(adapter_name, 'grad_ready', False) + else: # TODO: Make LoraConfig more flexible lora_cfg = LoraConfig(r=body.lora_config.rank, target_modules='all-linear') - adapter_name = self.get_adapter_name(adapter_name=_model_id) self.register_resource(adapter_name, token, session_id=body.session_id) self.model.add_adapter_to_model(adapter_name=adapter_name, config_or_dir=lora_cfg) - # Select template based on model type - template = get_template_for_model(self.base_model) self.model.set_template(template, adapter_name=adapter_name, model_id=self.base_model) self.model.set_processor('InputProcessor', adapter_name=adapter_name) self.model.set_optimizer('Adam', adapter_name=adapter_name) @@ -59,6 +87,9 @@ async def _create_adapter(): training_run_manager = create_training_run_manager(token, client_type='tinker') training_run_manager.save(_model_id, body) return types.CreateModelResponse(model_id=_model_id) + except FullModeBusyError as e: + # Nothing was registered yet (check runs before register_model). + return types.RequestFailedResponse(error=str(e), category=types.RequestErrorCategory.User) except Exception: if _model_id: adapter_name = self.get_adapter_name(adapter_name=_model_id) @@ -115,10 +146,11 @@ async def _do_forward(): try: adapter_name = self.get_adapter_name(adapter_name=body.model_id) self.assert_resource_exists(adapter_name) + model_adapter = self.resolve_model_adapter_name(adapter_name) datum_list = body.forward_input.data loss_fn_config = body.forward_input.loss_fn_config or {} output, loss = self.model.tinker_forward_only( - inputs=datum_list, adapter_name=adapter_name, **loss_fn_config) + inputs=datum_list, adapter_name=model_adapter, **loss_fn_config) return types.ForwardBackwardOutput( loss_fn_output_type='CrossEntropyLossReturn', loss_fn_outputs=output, @@ -156,11 +188,12 @@ async def _do_forward_backward(): try: adapter_name = self.get_adapter_name(adapter_name=body.model_id) self.assert_resource_exists(adapter_name) + model_adapter = self.resolve_model_adapter_name(adapter_name) datum_list = body.forward_backward_input.data loss_fn = body.forward_backward_input.loss_fn loss_fn_config = body.forward_backward_input.loss_fn_config or {} output, loss = self.model.tinker_forward_backward( - inputs=datum_list, adapter_name=adapter_name, loss_fn=loss_fn, **loss_fn_config) + inputs=datum_list, adapter_name=model_adapter, loss_fn=loss_fn, **loss_fn_config) output_type = ('ImportanceSamplingLossReturn' if loss_fn == 'importance_sampling' else 'CrossEntropyLossReturn') self.set_resource_state(adapter_name, 'grad_ready', True) @@ -203,12 +236,13 @@ async def _do_optim(): try: adapter_name = self.get_adapter_name(adapter_name=body.model_id) self.assert_resource_exists(adapter_name) + model_adapter = self.resolve_model_adapter_name(adapter_name) if not self.get_resource_state(adapter_name, 'grad_ready', False): raise RuntimeError(f'No accumulated gradients for adapter={adapter_name}; ' 'call forward_backward before optim_step') - self.model.tinker_step(adam_params=body.adam_params, adapter_name=adapter_name) + self.model.tinker_step(adam_params=body.adam_params, adapter_name=model_adapter) self.set_resource_state(adapter_name, 'grad_ready', False) - metrics = self.model.tinker_calculate_metric(is_training=True, adapter_name=adapter_name) + metrics = self.model.tinker_calculate_metric(is_training=True, adapter_name=model_adapter) return types.OptimStepResponse(metrics=metrics) except Exception: logger.error(traceback.format_exc()) @@ -231,11 +265,12 @@ async def _do_save(): try: adapter_name = self.get_adapter_name(adapter_name=body.model_id) self.assert_resource_exists(adapter_name) + model_adapter = self.resolve_model_adapter_name(adapter_name) checkpoint_manager = create_checkpoint_manager(token, client_type='tinker') checkpoint_name = checkpoint_manager.get_ckpt_name(body.path) save_dir = checkpoint_manager.get_save_dir(model_id=body.model_id, is_sampler=False) self.model.save( - name=checkpoint_name, output_dir=save_dir, adapter_name=adapter_name, save_optimizer=True) + name=checkpoint_name, output_dir=save_dir, adapter_name=model_adapter, save_optimizer=True) tinker_path = checkpoint_manager.save(body.model_id, name=checkpoint_name, is_sampler=False) return types.SaveWeightsResponse(path=tinker_path, type='save_weights') except Exception: @@ -265,7 +300,9 @@ async def _do_save_for_sampler(): # Must save the checkpoint in the twinkle format before calling model.save() tinker_path = checkpoint_manager.save(body.model_id, name=checkpoint_name, is_sampler=True) logger.info(f'Saving weights to {save_dir}') - self.model.save(name='latest', output_dir=save_dir, adapter_name=adapter_name, save_optimizer=False) + self.model.save( + name='latest', output_dir=save_dir, + adapter_name=self.resolve_model_adapter_name(adapter_name), save_optimizer=False) payload = body.model_dump() payload['model_path'] = tinker_path metadata = await self.state.get_model_metadata(body.model_id) or {} @@ -305,7 +342,8 @@ async def _do_load(): adapter_name = self.get_adapter_name(adapter_name=body.model_id) self.assert_resource_exists(adapter_name) self.model.tinker_load( - checkpoint_dir=body.path, load_optimizer=body.optimizer, adapter_name=adapter_name, token=token) + checkpoint_dir=body.path, load_optimizer=body.optimizer, + adapter_name=self.resolve_model_adapter_name(adapter_name), token=token) self.set_resource_state(adapter_name, 'grad_ready', False) return types.LoadWeightsResponse(path=body.path, type='load_weights') except Exception: diff --git a/src/twinkle/server/model/twinkle_handlers.py b/src/twinkle/server/model/twinkle_handlers.py index 2074f5e40..114a0e407 100644 --- a/src/twinkle/server/model/twinkle_handlers.py +++ b/src/twinkle/server/model/twinkle_handlers.py @@ -24,6 +24,7 @@ from twinkle.data_format import InputFeature, Trajectory from twinkle.server.checkpoint import (_resolve_client_save_dir, create_checkpoint_manager, create_training_run_manager, validate_user_path) +from twinkle.server.exceptions import FullModeBusyError from twinkle.server.utils.validation import get_session_id_from_request from twinkle.utils.logger import get_logger from twinkle_client.common.serialize import deserialize_object @@ -105,7 +106,7 @@ async def _task(): self.assert_resource_exists(adapter_name) extra_kwargs = body.model_extra or {} inputs = _parse_inputs(body.inputs) - ret = self.model.forward(inputs=inputs, adapter_name=adapter_name, **extra_kwargs) + ret = self.model.forward(inputs=inputs, adapter_name=self.resolve_model_adapter_name(adapter_name), **extra_kwargs) return {'result': ret} inputs_list = body.inputs if isinstance(body.inputs, list) else [body.inputs] @@ -135,7 +136,7 @@ async def _task(): self.assert_resource_exists(adapter_name) extra_kwargs = body.model_extra or {} inputs = _parse_inputs(body.inputs) - ret = self.model.forward_only(inputs=inputs, adapter_name=adapter_name, **extra_kwargs) + ret = self.model.forward_only(inputs=inputs, adapter_name=self.resolve_model_adapter_name(adapter_name), **extra_kwargs) return {'result': ret} inputs_list = body.inputs if isinstance(body.inputs, list) else [body.inputs] @@ -161,7 +162,7 @@ async def calculate_loss( async def _task(): self.assert_resource_exists(adapter_name) extra_kwargs = body.model_extra or {} - ret = self.model.calculate_loss(adapter_name=adapter_name, **extra_kwargs) + ret = self.model.calculate_loss(adapter_name=self.resolve_model_adapter_name(adapter_name), **extra_kwargs) return {'result': ret} return await run_task( @@ -175,7 +176,7 @@ async def backward(request: Request, body: types.AdapterRequest, self: ModelMana async def _task(): self.assert_resource_exists(adapter_name) extra_kwargs = body.model_extra or {} - self.model.backward(adapter_name=adapter_name, **extra_kwargs) + self.model.backward(adapter_name=self.resolve_model_adapter_name(adapter_name), **extra_kwargs) await run_task(self.schedule_task_and_wait(_task, model_id=adapter_name, token=token, task_type='backward')) @@ -203,7 +204,7 @@ async def _task(): for key in inputs: if isinstance(inputs[key], list) and isinstance(first_element(inputs[key]), (int, float)): inputs[key] = torch.tensor(inputs[key]) - ret = self.model.forward_backward(inputs=all_inputs, adapter_name=adapter_name, **extra_kwargs) + ret = self.model.forward_backward(inputs=all_inputs, adapter_name=self.resolve_model_adapter_name(adapter_name), **extra_kwargs) return {'result': ret} inputs_list = body.inputs if isinstance(body.inputs, list) else [body.inputs] @@ -232,7 +233,7 @@ async def clip_grad_norm( async def _task(): self.assert_resource_exists(adapter_name) extra_kwargs = body.model_extra or {} - ret = self.model.clip_grad_norm(adapter_name=adapter_name, **extra_kwargs) + ret = self.model.clip_grad_norm(adapter_name=self.resolve_model_adapter_name(adapter_name), **extra_kwargs) return {'result': str(ret)} return await run_task( @@ -246,7 +247,7 @@ async def step(request: Request, body: types.AdapterRequest, self: ModelManageme async def _task(): self.assert_resource_exists(adapter_name) extra_kwargs = body.model_extra or {} - self.model.step(adapter_name=adapter_name, **extra_kwargs) + self.model.step(adapter_name=self.resolve_model_adapter_name(adapter_name), **extra_kwargs) await run_task(self.schedule_task_and_wait(_task, model_id=adapter_name, token=token, task_type='step')) @@ -258,7 +259,7 @@ async def zero_grad(request: Request, body: types.AdapterRequest, self: ModelMan async def _task(): self.assert_resource_exists(adapter_name) extra_kwargs = body.model_extra or {} - self.model.zero_grad(adapter_name=adapter_name, **extra_kwargs) + self.model.zero_grad(adapter_name=self.resolve_model_adapter_name(adapter_name), **extra_kwargs) await run_task(self.schedule_task_and_wait(_task, model_id=adapter_name, token=token, task_type='zero_grad')) @@ -270,7 +271,7 @@ async def lr_step(request: Request, body: types.AdapterRequest, self: ModelManag async def _task(): self.assert_resource_exists(adapter_name) extra_kwargs = body.model_extra or {} - self.model.lr_step(adapter_name=adapter_name, **extra_kwargs) + self.model.lr_step(adapter_name=self.resolve_model_adapter_name(adapter_name), **extra_kwargs) await run_task(self.schedule_task_and_wait(_task, model_id=adapter_name, token=token, task_type='lr_step')) @@ -289,7 +290,7 @@ async def _task(): self.model.clip_grad_and_step( max_grad_norm=body.max_grad_norm, norm_type=body.norm_type, - adapter_name=adapter_name, + adapter_name=self.resolve_model_adapter_name(adapter_name), **extra_kwargs, ) @@ -308,7 +309,7 @@ async def get_train_configs( async def _task(): self.assert_resource_exists(adapter_name) extra_kwargs = body.model_extra or {} - ret = self.model.get_train_configs(adapter_name=adapter_name, **extra_kwargs) + ret = self.model.get_train_configs(adapter_name=self.resolve_model_adapter_name(adapter_name), **extra_kwargs) return {'result': ret} return await run_task( @@ -322,7 +323,7 @@ async def set_loss(request: Request, body: types.SetLossRequest, self: ModelMana async def _task(): self.assert_resource_exists(adapter_name) extra_kwargs = body.model_extra or {} - self.model.set_loss(body.loss_cls, adapter_name=adapter_name, **extra_kwargs) + self.model.set_loss(body.loss_cls, adapter_name=self.resolve_model_adapter_name(adapter_name), **extra_kwargs) await run_task(self.schedule_task_and_wait(_task, model_id=adapter_name, token=token, task_type='set_loss')) @@ -338,7 +339,7 @@ async def set_optimizer( async def _task(): self.assert_resource_exists(adapter_name) extra_kwargs = body.model_extra or {} - self.model.set_optimizer(body.optimizer_cls, adapter_name=adapter_name, **extra_kwargs) + self.model.set_optimizer(body.optimizer_cls, adapter_name=self.resolve_model_adapter_name(adapter_name), **extra_kwargs) await run_task( self.schedule_task_and_wait(_task, model_id=adapter_name, token=token, task_type='set_optimizer')) @@ -355,7 +356,7 @@ async def set_lr_scheduler( async def _task(): self.assert_resource_exists(adapter_name) extra_kwargs = body.model_extra or {} - self.model.set_lr_scheduler(body.scheduler_cls, adapter_name=adapter_name, **extra_kwargs) + self.model.set_lr_scheduler(body.scheduler_cls, adapter_name=self.resolve_model_adapter_name(adapter_name), **extra_kwargs) await run_task( self.schedule_task_and_wait(_task, model_id=adapter_name, token=token, task_type='set_lr_scheduler')) @@ -380,7 +381,7 @@ async def _task(): checkpoint_dir = self.model.save( name=model_save_name, output_dir=save_dir, - adapter_name=adapter_name, + adapter_name=self.resolve_model_adapter_name(adapter_name), save_optimizer=body.save_optimizer, **extra_kwargs) return {'twinkle_path': twinkle_path, 'checkpoint_dir': checkpoint_dir} @@ -400,7 +401,7 @@ async def _task(): self.model.load( name=resolved.checkpoint_name, output_dir=resolved.checkpoint_dir, - adapter_name=adapter_name, + adapter_name=self.resolve_model_adapter_name(adapter_name), load_optimizer=body.load_optimizer, token=token, **extra_kwargs) @@ -426,7 +427,7 @@ async def _task(): ret = self.model.resume_from_checkpoint( checkpoint_dir, resume_only_model=body.resume_only_model, - adapter_name=adapter_name, + adapter_name=self.resolve_model_adapter_name(adapter_name), ) return {'result': ret} @@ -511,6 +512,24 @@ async def _task(): extra_kwargs = body.model_extra or {} training_run_manager = create_training_run_manager(token, client_type='twinkle') + # Validate the supplied config against the deployment's train_mode. + if self.is_full_mode and config is not None: + raise HTTPException( + status_code=400, + detail='This deployment runs in full-parameter (exclusive) mode; pass config=None ' + '(do not send a LoraConfig).') + if (not self.is_full_mode) and config is None: + raise HTTPException( + status_code=400, + detail='This deployment runs in LoRA mode; a LoraConfig is required.') + + # In full mode ensure the exclusive deployment is free before touching state. + if self.is_full_mode: + try: + self.assert_full_mode_available(adapter_name) + except FullModeBusyError as e: + raise HTTPException(status_code=409, detail=str(e)) + lora_config = None if isinstance(config, LoraConfig): lora_config = types.LoraConfig(rank=config.r, train_unembed=False, train_mlp=True, train_attn=True) @@ -528,7 +547,11 @@ async def _task(): ) try: self.register_resource(adapter_name, token, session_id) - self.model.add_adapter_to_model(adapter_name, config, **extra_kwargs) + if self.is_full_mode: + # No PEFT adapter to add; the default optimizer group is used. + self.set_resource_state(adapter_name, 'grad_ready', False) + else: + self.model.add_adapter_to_model(adapter_name, config, **extra_kwargs) except Exception: self.unregister_resource(adapter_name) await self.state.unload_model(adapter_name) @@ -552,7 +575,7 @@ async def _task(): self.assert_resource_exists(adapter_name) extra_kwargs = body.model_extra or {} patch_cls = deserialize_object(body.patch_cls) - self.model.apply_patch(patch_cls, adapter_name=adapter_name, **extra_kwargs) + self.model.apply_patch(patch_cls, adapter_name=self.resolve_model_adapter_name(adapter_name), **extra_kwargs) await run_task(self.schedule_task_and_wait(_task, model_id=adapter_name, token=token, task_type='apply_patch')) @@ -569,7 +592,7 @@ async def _task(): self.assert_resource_exists(adapter_name) extra_kwargs = body.model_extra or {} metric_cls = deserialize_object(body.metric_cls) - self.model.add_metric(metric_cls, is_training=body.is_training, adapter_name=adapter_name, **extra_kwargs) + self.model.add_metric(metric_cls, is_training=body.is_training, adapter_name=self.resolve_model_adapter_name(adapter_name), **extra_kwargs) await run_task(self.schedule_task_and_wait(_task, model_id=adapter_name, token=token, task_type='add_metric')) @@ -585,7 +608,7 @@ async def set_template( async def _task(): self.assert_resource_exists(adapter_name) extra_kwargs = body.model_extra or {} - self.model.set_template(body.template_cls, adapter_name=adapter_name, **extra_kwargs) + self.model.set_template(body.template_cls, adapter_name=self.resolve_model_adapter_name(adapter_name), **extra_kwargs) await run_task(self.schedule_task_and_wait(_task, model_id=adapter_name, token=token, task_type='set_template')) @@ -601,7 +624,7 @@ async def set_processor( async def _task(): self.assert_resource_exists(adapter_name) extra_kwargs = body.model_extra or {} - self.model.set_processor(body.processor_cls, adapter_name=adapter_name, **extra_kwargs) + self.model.set_processor(body.processor_cls, adapter_name=self.resolve_model_adapter_name(adapter_name), **extra_kwargs) await run_task( self.schedule_task_and_wait(_task, model_id=adapter_name, token=token, task_type='set_processor')) @@ -618,7 +641,7 @@ async def calculate_metric( async def _task(): self.assert_resource_exists(adapter_name) extra_kwargs = body.model_extra or {} - ret = self.model.calculate_metric(is_training=body.is_training, adapter_name=adapter_name, **extra_kwargs) + ret = self.model.calculate_metric(is_training=body.is_training, adapter_name=self.resolve_model_adapter_name(adapter_name), **extra_kwargs) return {'result': ret} return await run_task( @@ -636,7 +659,7 @@ async def get_state_dict( async def _task(): self.assert_resource_exists(adapter_name) extra_kwargs = body.model_extra or {} - ret = self.model.get_state_dict(adapter_name=adapter_name, **extra_kwargs) + ret = self.model.get_state_dict(adapter_name=self.resolve_model_adapter_name(adapter_name), **extra_kwargs) return {'result': ret} return await run_task( diff --git a/src/twinkle/server/sampler/backends/mock_sampler.py b/src/twinkle/server/sampler/backends/mock_sampler.py index 300e5b0c3..735f52639 100644 --- a/src/twinkle/server/sampler/backends/mock_sampler.py +++ b/src/twinkle/server/sampler/backends/mock_sampler.py @@ -170,6 +170,11 @@ def sample_stream_to_queue(self, queue, inputs, sampling_params=None, adapter_na def apply_patch(self, patch_cls: Any, **kwargs: Any) -> None: return None + @remote_function() + def load_full_weights_from_path(self, path: str, *args: Any, **kwargs: Any) -> int: + """No-op full-weight load for the mock full-parameter sampling path.""" + return 0 + @remote_function() def set_template(self, template_cls: Any, **kwargs: Any) -> None: self.template = template_cls diff --git a/src/twinkle/server/sampler/tinker_handlers.py b/src/twinkle/server/sampler/tinker_handlers.py index 91fdf7813..de43b4b93 100644 --- a/src/twinkle/server/sampler/tinker_handlers.py +++ b/src/twinkle/server/sampler/tinker_handlers.py @@ -87,10 +87,21 @@ async def _do_sample(): stop=body.sampling_params.stop, ) + # A resolved checkpoint is either a LoRA adapter dir (has + # adapter_config.json) or a full-parameter HF checkpoint. Full + # checkpoints are loaded into the sampler base model instead of + # being passed as a LoRA adapter. + lora_path = None + if adapter_uri: + if os.path.exists(os.path.join(adapter_uri, 'adapter_config.json')): + lora_path = adapter_uri + else: + self.sampler.load_full_weights_from_path(adapter_uri) + responses = self.sampler.sample( inputs=[prompt_inputs] * body.num_samples, sampling_params=sampling_params, - adapter_path=adapter_uri, + adapter_path=lora_path, ) tinker_sequences = [] diff --git a/src/twinkle/server/sampler/twinkle_handlers.py b/src/twinkle/server/sampler/twinkle_handlers.py index 8a10f0b8d..0b07b4637 100644 --- a/src/twinkle/server/sampler/twinkle_handlers.py +++ b/src/twinkle/server/sampler/twinkle_handlers.py @@ -101,11 +101,18 @@ async def _task(): full_adapter_name = _get_twinkle_sampler_adapter_name(request, adapter_name) or '' if body.adapter_uri: + import os from twinkle.server.checkpoint import create_checkpoint_manager checkpoint_manager = create_checkpoint_manager(token, client_type='twinkle') - _, adapter_path = checkpoint_manager.parse_adapter_uri(body.adapter_uri) + _, resolved_uri = checkpoint_manager.parse_adapter_uri(body.adapter_uri) # Reset prefix cache only when new weights are loaded self.sampler.reset_prefix_cache() + # LoRA adapter dir (has adapter_config.json) vs full-parameter + # HF checkpoint. Full checkpoints replace the sampler base model. + if resolved_uri and os.path.exists(os.path.join(resolved_uri, 'adapter_config.json')): + adapter_path = resolved_uri + elif resolved_uri: + self.sampler.load_full_weights_from_path(resolved_uri) # Parse inputs inputs = body.inputs @@ -227,10 +234,15 @@ async def sample_stream( full_adapter_name = _get_twinkle_sampler_adapter_name(request, adapter_name) or '' if body.adapter_uri: + import os from twinkle.server.checkpoint import create_checkpoint_manager checkpoint_manager = create_checkpoint_manager(token, client_type='twinkle') - _, adapter_path = checkpoint_manager.parse_adapter_uri(body.adapter_uri) + _, resolved_uri = checkpoint_manager.parse_adapter_uri(body.adapter_uri) self.sampler.reset_prefix_cache() + if resolved_uri and os.path.exists(os.path.join(resolved_uri, 'adapter_config.json')): + adapter_path = resolved_uri + elif resolved_uri: + self.sampler.load_full_weights_from_path(resolved_uri) inputs = body.inputs if isinstance(inputs, list): diff --git a/src/twinkle_client/model/multi_lora_transformers.py b/src/twinkle_client/model/multi_lora_transformers.py index c628c353b..bf4ef54af 100644 --- a/src/twinkle_client/model/multi_lora_transformers.py +++ b/src/twinkle_client/model/multi_lora_transformers.py @@ -37,8 +37,13 @@ def __init__(self, model_id: str, **kwargs): ) response.raise_for_status() - def add_adapter_to_model(self, adapter_name: str, config: Dict[str, Any], **kwargs) -> None: - """Add a new adapter to the model.""" + def add_adapter_to_model(self, adapter_name: str, config: Optional[Dict[str, Any]] = None, **kwargs) -> None: + """Add a new adapter to the model. + + Pass a peft ``LoraConfig`` (or its dict form) for LoRA training against a + LoRA-mode deployment. Pass ``config=None`` for full-parameter training + against a ``train_mode: full`` deployment. + """ save_dir = kwargs.get('save_dir') if save_dir: kwargs['save_dir'] = Path(save_dir).expanduser().resolve().as_posix() diff --git a/src/twinkle_client/types/model.py b/src/twinkle_client/types/model.py index 10a60b947..83f489cde 100644 --- a/src/twinkle_client/types/model.py +++ b/src/twinkle_client/types/model.py @@ -109,7 +109,9 @@ class Config: class AddAdapterRequest(BaseModel): adapter_name: str - config: str + # ``config`` is None for full-parameter training (no LoRA adapter) and a + # serialized LoraConfig string for LoRA training. + config: Optional[str] = None save_dir: Optional[str] = None class Config: diff --git a/src/twinkle_client/utils/patch_tinker.py b/src/twinkle_client/utils/patch_tinker.py index 45b62d9ea..ed9abaa82 100644 --- a/src/twinkle_client/utils/patch_tinker.py +++ b/src/twinkle_client/utils/patch_tinker.py @@ -138,6 +138,62 @@ def _patched_service_client_init(self, user_metadata=None, **kwargs): return _patched_service_client_init +def _create_full_training_client_submit(self, base_model, seed=None, user_metadata=None): + """Submit a full-parameter (non-LoRA) training-client creation request. + + Mirrors the SDK's ``_create_lora_training_client_submit`` but sends + ``lora_config=None`` so the Twinkle server routes to a full-parameter + (``train_mode: full``) deployment. Returns the same ``TrainingClient`` so + the training loop (forward_backward / optim_step / save_weights / ...) is + identical to the LoRA path. + """ + from tinker.lib.public_interfaces import service_client as _sc + from tinker.lib.internal_client_holder import ClientConnectionPoolType + + session_id = self.holder.get_session_id() + model_seq_id = self.holder.get_training_client_id() + + async def _create_full_training_client_async(): + start_time = _sc.time.time() + with self.holder.aclient(ClientConnectionPoolType.TRAIN) as client: + request = _sc.types.CreateModelRequest( + session_id=session_id, + model_seq_id=model_seq_id, + base_model=base_model, + lora_config=None, + user_metadata=user_metadata, + ) + future = await client.models.create(request=request) + create_model_response = await _sc._APIFuture( + _sc.types.CreateModelResponse, + self.holder, + future, + request_start_time=start_time, + request_type='CreateModel', + queue_state_observer=_sc.QueueStateLogger(base_model, 'Model creation'), + ).result_async() + model_id = create_model_response.model_id + from tinker.lib.public_interfaces.training_client import TrainingClient + + training_client = TrainingClient(self.holder, model_seq_id=model_seq_id, model_id=model_id) + _sc.logger.info(f'Full-parameter TrainingClient initialized for model {model_id}') + return training_client + + return self.holder.run_coroutine_threadsafe(_create_full_training_client_async()) + + +def _create_full_training_client(self, base_model, seed=None, user_metadata=None): + """Create a full-parameter (non-LoRA) training client (blocking).""" + return _create_full_training_client_submit( + self, base_model, seed=seed, user_metadata=user_metadata).result() + + +async def _create_full_training_client_async(self, base_model, seed=None, user_metadata=None): + """Async variant of ``create_full_training_client``.""" + return await _create_full_training_client_submit( + self, base_model, seed=seed, user_metadata=user_metadata).result_async() + + def patch_tinker(): """ Apply patches to tinker library. @@ -147,6 +203,7 @@ def patch_tinker(): 2. AsyncTinker.__init__ to bypass 'tml-' prefix validation for api_key 3. ParsedCheckpointTinkerPath.from_tinker_path to support both 'tinker://' and 'twinkle://' prefixes 4. _get_default_headers to inject Twinkle-specific headers + 5. ServiceClient.create_full_training_client to enable full-parameter (non-LoRA) training This patch is idempotent - calling it multiple times has no additional effect. """ @@ -171,6 +228,10 @@ def patch_tinker(): from tinker.lib.public_interfaces.service_client import ServiceClient ServiceClient.__init__ = _make_patched_service_client_init(ServiceClient.__init__) + # Patch 5: add a full-parameter (non-LoRA) training-client entrypoint. + ServiceClient.create_full_training_client = _create_full_training_client + ServiceClient.create_full_training_client_async = _create_full_training_client_async + _patched = True except ImportError: # tinker not installed, skip patching diff --git a/tests/server/integration/test_full_param_e2e.py b/tests/server/integration/test_full_param_e2e.py new file mode 100644 index 000000000..b328b8a5e --- /dev/null +++ b/tests/server/integration/test_full_param_e2e.py @@ -0,0 +1,180 @@ +"""Full-parameter (train_mode: full) e2e: create → SFT → guards → save → sample. + +Drives the tinker-compatible client against a server whose model deployment +runs in exclusive full-parameter mode (no PEFT/MultiLora). Start the server +with a full-parameter variant, e.g.: + + python tests/server/start_e2e_server.py --variant full # transformers DP=2 + python tests/server/start_e2e_server.py --variant full-fsdp2 # transformers FSDP2 + python tests/server/start_e2e_server.py --variant megatron-full # megatron DP=2 PP=2 + python tests/server/start_e2e_server.py --variant megatron-full-tp2pp2 + python tests/server/start_e2e_server.py --variant megatron-full-tp2dp2 + +Phase A — create_full_training_client + SFT training, assert loss decreases +Phase B — create_lora_training_client must be rejected (mode mismatch) +Phase C — second create_full_training_client must be rejected (exclusive busy) +Phase D — save_weights_and_get_sampling_client → sample, non-empty output +Phase E — on-disk checkpoint is a FULL HF checkpoint (no adapter_config.json) + +## How to run + + TWINKLE_TEST_GPU_E2E=1 python -u tests/server/integration/test_full_param_e2e.py + +Expected last line: ``ALL PHASES PASSED``. +""" +from __future__ import annotations + +import dotenv + +dotenv.load_dotenv('.env') + +import glob # noqa: E402 +import os # noqa: E402 +import sys # noqa: E402 +import time # noqa: E402 + +import numpy as np # noqa: E402 +import pytest # noqa: E402 + +pytestmark = pytest.mark.skipif( + os.environ.get('TWINKLE_TEST_GPU_E2E', '0') != '1', + reason='Set TWINKLE_TEST_GPU_E2E=1 to run real GPU E2E tests (requires running server)', +) + +from twinkle import get_logger, init_tinker_client # noqa: E402 +from twinkle.dataloader import DataLoader # noqa: E402 +from twinkle.dataset import Dataset, DatasetMeta # noqa: E402 +from twinkle.preprocessor import SelfCognitionProcessor # noqa: E402 +from twinkle.server.common import input_feature_to_datum # noqa: E402 + +init_tinker_client() + +from tinker import ServiceClient, types # noqa: E402 + +logger = get_logger() + +BASE_MODEL = 'Qwen/Qwen3.5-4B' +BASE_URL = os.environ.get('TWINKLE_SERVER_URL', 'http://localhost:9000') +API_KEY = os.environ.get('TWINKLE_SERVER_TOKEN', 'EMPTY_TOKEN') +OUTPUTS_ROOT = os.path.join(os.path.dirname(__file__), '..', '..', '..', 'outputs') +TRAIN_STEPS = 16 +LEARNING_RATE = 1e-5 +SAMPLE_MAX_TOKENS = 32 + + +def _build_dataloader(batch_size: int = 8): + dataset = Dataset(dataset_meta=DatasetMeta('ms://swift/self-cognition', data_slice=range(500))) + dataset.set_template('Qwen3_5Template', model_id=f'ms://{BASE_MODEL}', max_length=256) + dataset.map(SelfCognitionProcessor('twinkle模型', 'twinkle团队'), load_from_cache_file=False) + dataset.encode(batched=True, load_from_cache_file=False) + return DataLoader(dataset=dataset, batch_size=batch_size) + + +def _loss_per_token(fwdbwd_result, input_datum) -> float: + logprobs = np.concatenate([output['logprobs'].tolist() for output in fwdbwd_result.loss_fn_outputs]) + weights = np.concatenate([example.loss_fn_inputs['weights'].tolist() for example in input_datum]) + return float(-np.dot(logprobs, weights) / weights.sum()) + + +def main() -> int: + t0 = time.time() + service_client = ServiceClient(base_url=BASE_URL, api_key=API_KEY) + + # ── Phase A: full-parameter training ── + logger.info('=' * 60) + logger.info('Phase A: create_full_training_client + SFT (%d steps)', TRAIN_STEPS) + logger.info('=' * 60) + training_client = service_client.create_full_training_client(base_model=BASE_MODEL) + logger.info('Full training client created: model_id=%s', training_client.model_id) + + dataloader = _build_dataloader() + losses = [] + step = 0 + for batch in dataloader: + input_datum = [input_feature_to_datum(f) for f in batch] + fwdbwd_future = training_client.forward_backward(input_datum, 'cross_entropy') + optim_future = training_client.optim_step(types.AdamParams(learning_rate=LEARNING_RATE)) + fwdbwd_result = fwdbwd_future.result() + optim_future.result() + loss = _loss_per_token(fwdbwd_result, input_datum) + losses.append(loss) + step += 1 + logger.info('[A] step=%d loss=%.4f', step, loss) + if step >= TRAIN_STEPS: + break + + first3, last3 = float(np.mean(losses[:3])), float(np.mean(losses[-3:])) + assert last3 < first3, f'Phase A FAIL: loss did not decrease (first3={first3:.4f}, last3={last3:.4f})' + logger.info('Phase A OK: loss %.4f -> %.4f', first3, last3) + + # ── Phase B: LoRA create must be rejected on a full deployment ── + logger.info('Phase B: create_lora_training_client must be rejected') + try: + service_client.create_lora_training_client(base_model=BASE_MODEL, rank=8) + raise AssertionError('Phase B FAIL: LoRA create unexpectedly succeeded on a full deployment') + except AssertionError: + raise + except Exception as e: # noqa: BLE001 + msg = str(e) + assert 'full-parameter' in msg or 'lora_config' in msg, f'Phase B FAIL: unexpected error: {msg[:500]}' + logger.info('Phase B OK: rejected with mode-mismatch error') + + # ── Phase C: second full tenant must be rejected (exclusive) ── + logger.info('Phase C: second create_full_training_client must be rejected (busy)') + second_client = ServiceClient(base_url=BASE_URL, api_key=API_KEY) + try: + second_client.create_full_training_client(base_model=BASE_MODEL) + raise AssertionError('Phase C FAIL: second full tenant unexpectedly succeeded') + except AssertionError: + raise + except Exception as e: # noqa: BLE001 + msg = str(e) + assert 'exclusive' in msg or 'already' in msg, f'Phase C FAIL: unexpected error: {msg[:500]}' + logger.info('Phase C OK: rejected with exclusive-busy error') + + # ── Phase D: save full weights for sampler + sample ── + logger.info('Phase D: save_weights_and_get_sampling_client + sample') + save_start = time.time() + sampling_client = training_client.save_weights_and_get_sampling_client(name='full-e2e') + logger.info('[D] sampler weights saved in %.0fs', time.time() - save_start) + + from twinkle.data_format import Message, Trajectory + from twinkle.template import Template + template = Template(model_id=f'ms://{BASE_MODEL}') + trajectory = Trajectory(messages=[ + Message(role='system', content='You are a helpful assistant'), + Message(role='user', content='你是谁?'), + ]) + input_feature = template.batch_encode([trajectory], add_generation_prompt=True)[0] + prompt = types.ModelInput.from_ints(input_feature['input_ids'].tolist()) + params = types.SamplingParams(max_tokens=SAMPLE_MAX_TOKENS, temperature=0.0) + result = sampling_client.sample(prompt=prompt, sampling_params=params, num_samples=2).result() + assert result.sequences and all(len(seq.tokens) > 0 for seq in result.sequences), \ + 'Phase D FAIL: empty sampling output' + for i, seq in enumerate(result.sequences): + decoded = template.decode(seq.tokens) + logger.info('[D] sample %d: %s', i, decoded[:120].replace('\n', ' ')) + logger.info('Phase D OK: sampled %d sequences on full weights', len(result.sequences)) + + # ── Phase E: on-disk checkpoint format ── + logger.info('Phase E: verify on-disk checkpoint is a full HF checkpoint') + candidates = [] + for pattern in ('model.safetensors', 'model.safetensors.index.json'): + candidates.extend(glob.glob(os.path.join(OUTPUTS_ROOT, '**', pattern), recursive=True)) + fresh = [p for p in candidates if os.path.getmtime(p) > save_start] + assert fresh, f'Phase E FAIL: no fresh full-weight files under {os.path.abspath(OUTPUTS_ROOT)}' + ckpt_dir = os.path.dirname(sorted(fresh, key=os.path.getmtime)[-1]) + assert not os.path.exists(os.path.join(ckpt_dir, 'adapter_config.json')), \ + f'Phase E FAIL: {ckpt_dir} contains adapter_config.json (LoRA format!)' + assert not glob.glob(os.path.join(ckpt_dir, 'adapter_model*')), \ + f'Phase E FAIL: {ckpt_dir} contains adapter_model files (LoRA format!)' + total_gb = sum(os.path.getsize(p) for p in glob.glob(os.path.join(ckpt_dir, '*.safetensors'))) / 1e9 + logger.info('Phase E OK: %s is a full checkpoint (%.1f GB safetensors)', ckpt_dir, total_gb) + + logger.info('Total elapsed: %.0fs', time.time() - t0) + logger.info('ALL PHASES PASSED') + return 0 + + +if __name__ == '__main__': + sys.exit(main()) diff --git a/tests/server/model/test_tinker_handlers.py b/tests/server/model/test_tinker_handlers.py index a5330172f..474ce700f 100644 --- a/tests/server/model/test_tinker_handlers.py +++ b/tests/server/model/test_tinker_handlers.py @@ -52,6 +52,8 @@ async def test_tinker_dpo_forward_backward_requires_per_dp_pairs(): class _SaveWeightsDummyManagement: """Dummy management that actually executes the task to test save_weights_for_sampler logic.""" + is_full_mode = False + def __init__(self): self.model = MagicMock() self.state = MagicMock() @@ -64,6 +66,9 @@ async def _on_request_start(self, request): def get_adapter_name(self, adapter_name=None): return adapter_name + def resolve_model_adapter_name(self, adapter_name): + return adapter_name + def assert_resource_exists(self, adapter_name): pass diff --git a/tests/server/start_e2e_server.py b/tests/server/start_e2e_server.py index dc2a9f694..0e8939ed0 100644 --- a/tests/server/start_e2e_server.py +++ b/tests/server/start_e2e_server.py @@ -1,9 +1,16 @@ """One-click: restart Ray cluster + launch Twinkle server + wait until ready. Usage: - python start_e2e_server.py # default config - python start_e2e_server.py --config my_config.yaml # custom config + python start_e2e_server.py # default config (transformers LoRA) + python start_e2e_server.py --variant megatron-full-tp2pp2 + python start_e2e_server.py --variant full --set device_mesh.fsdp_size=2 --set device_mesh.dp_size=null + python start_e2e_server.py --config my_config.yaml # fully custom config python start_e2e_server.py --kill-only # just kill everything + +Variants are generated from the two base configs by patching the model +application's ``args`` (see VARIANTS); ``--set key.path=value`` adds ad-hoc +overrides on top (value parsed as YAML; ``null`` deletes the key). The +generated file is written to tests/server/config/_generated_e2e.yaml. """ import argparse import os @@ -22,6 +29,29 @@ RAY_TEMP_DIR = "/mnt/nas2/yunlin.myl/ray_logs" SERVER_LOG = os.path.join(WORKDIR, "server_e2e.log") +# ── Config variants ── +# Each variant = (base config, {dotted path in model app args: value}). +# A value of None deletes the key (e.g. to swap device_mesh dimensions). +TRANSFORMERS_BASE = "tests/server/config/server_config_4b_e2e.yaml" +MEGATRON_BASE = "tests/server/config/server_config_4b_e2e_megatron.yaml" +_FULL = { + "train_mode": "full", + # Full-parameter training holds the deployment exclusively; use long + # timeouts so the tenant is not expired mid-test. + "adapter_config.adapter_timeout": 3600, + "adapter_config.adapter_max_lifetime": 36000, +} +VARIANTS = { + "lora": (TRANSFORMERS_BASE, {}), + "full": (TRANSFORMERS_BASE, {**_FULL}), + "full-fsdp2": (TRANSFORMERS_BASE, {**_FULL, "device_mesh.dp_size": None, "device_mesh.fsdp_size": 2}), + "megatron-lora": (MEGATRON_BASE, {}), + "megatron-full": (MEGATRON_BASE, {**_FULL}), + "megatron-full-tp2pp2": (MEGATRON_BASE, {**_FULL, "device_mesh.dp_size": None, "device_mesh.tp_size": 2}), + "megatron-full-tp2dp2": (MEGATRON_BASE, {**_FULL, "device_mesh.pp_size": None, "device_mesh.tp_size": 2}), +} +GENERATED_CONFIG = "tests/server/config/_generated_e2e.yaml" + # ── Server check ── SERVER_URL = "http://localhost:9000/-/routes" READY_KEYWORD = "processor" @@ -29,6 +59,45 @@ POLL_INTERVAL = 5 +def build_config(base: str, overrides: dict, set_args: list[str]) -> str: + """Generate a config from ``base`` by patching the model app's args. + + Returns the path (relative to WORKDIR) of the generated yaml. + """ + import yaml + + with open(os.path.join(WORKDIR, base)) as f: + cfg = yaml.safe_load(f) + + model_apps = [a for a in cfg["applications"] if a.get("import_path") == "model"] + assert len(model_apps) == 1, f"expected exactly one model app in {base}" + args = model_apps[0]["args"] + + merged = dict(overrides) + for item in set_args: + key, sep, raw = item.partition("=") + assert sep, f"--set expects key.path=value, got: {item}" + merged[key.strip()] = yaml.safe_load(raw) + + for dotted, value in merged.items(): + node = args + *parents, leaf = dotted.split(".") + for p in parents: + node = node.setdefault(p, {}) + if value is None: + node.pop(leaf, None) + else: + node[leaf] = value + + out_path = os.path.join(WORKDIR, GENERATED_CONFIG) + with open(out_path, "w") as f: + f.write("# AUTO-GENERATED by start_e2e_server.py — do not edit.\n") + f.write(f"# base: {base}\n# model-args overrides: {merged}\n") + yaml.safe_dump(cfg, f, sort_keys=False) + print(f" Generated config: {GENERATED_CONFIG} (base={base}, overrides={merged})") + return GENERATED_CONFIG + + def run(cmd, env=None, check=True): """Run a shell command, print it, and return CompletedProcess.""" full_env = os.environ.copy() @@ -123,8 +192,13 @@ def wait_ready(): def main(): parser = argparse.ArgumentParser(description="Restart Ray + launch Twinkle server") - parser.add_argument("--config", default=DEFAULT_CONFIG, - help=f"Server config yaml (default: {DEFAULT_CONFIG})") + parser.add_argument("--config", default=None, + help=f"Server config yaml (overrides --variant; default: {DEFAULT_CONFIG})") + parser.add_argument("--variant", default=None, choices=sorted(VARIANTS), + help="Named config variant generated from a base config (see VARIANTS)") + parser.add_argument("--set", dest="set_args", action="append", default=[], + metavar="KEY.PATH=VALUE", + help="Extra model-args override, YAML-parsed (null deletes the key); repeatable") parser.add_argument("--kill-only", action="store_true", help="Only kill server + Ray, don't restart") parser.add_argument("--no-ray", action="store_true", @@ -138,10 +212,21 @@ def main(): print("Done (kill-only)") return 0 + if args.config: + config = args.config + if args.set_args: + config = build_config(config, {}, args.set_args) + else: + base, overrides = VARIANTS[args.variant or "lora"] + if args.variant or args.set_args: + config = build_config(base, overrides, args.set_args) + else: + config = DEFAULT_CONFIG + if not args.no_ray: restart_ray() - launch_server(args.config) + launch_server(config) if wait_ready(): print("\n✓ All set. Run your test:")