diff --git a/cookbook/client/async_rl/server_config.yaml b/cookbook/client/async_rl/server_config.yaml index 103ea523f..6c0eb612d 100644 --- a/cookbook/client/async_rl/server_config.yaml +++ b/cookbook/client/async_rl/server_config.yaml @@ -42,9 +42,6 @@ applications: target_ongoing_requests: 128 ray_actor_options: num_cpus: 0.1 - runtime_env: - env_vars: - TWINKLE_FAIL_FAST: "0" # TransferQueue-backed DataRef service. - name: data-plane @@ -95,7 +92,6 @@ applications: runtime_env: env_vars: TWINKLE_TRUST_REMOTE_CODE: "1" - TWINKLE_FAIL_FAST: "0" # A second GPU hosts vLLM and loads the same local base model. - name: sampler-Qwen3.5-4B @@ -133,7 +129,6 @@ applications: runtime_env: env_vars: TWINKLE_TRUST_REMOTE_CODE: "1" - TWINKLE_FAIL_FAST: "0" - name: processor route_prefix: /api/v1/processor @@ -155,6 +150,3 @@ applications: target_ongoing_requests: 128 ray_actor_options: num_cpus: 0.1 - runtime_env: - env_vars: - TWINKLE_FAIL_FAST: "0" diff --git a/cookbook/client/server/megatron/server_config.yaml b/cookbook/client/server/megatron/server_config.yaml index 24132a07e..81d14e75e 100644 --- a/cookbook/client/server/megatron/server_config.yaml +++ b/cookbook/client/server/megatron/server_config.yaml @@ -54,7 +54,6 @@ applications: env_vars: TWINKLE_TRUST_REMOTE_CODE: "0" TWINKLE_LONG_POLL_TIMEOUT: "120" - TWINKLE_FAIL_FAST: "0" # 3. Sampler Service - Runs inference / sampling using vLLM engine # Used for generating text from the model (e.g., evaluating LoRA results). @@ -98,7 +97,6 @@ applications: env_vars: TWINKLE_TRUST_REMOTE_CODE: "0" TWINKLE_LONG_POLL_TIMEOUT: "120" - TWINKLE_FAIL_FAST: "0" # 2. Model Service - Hosts the base model for training. # Config: PP=2 x DP=2 on 4 GPUs, ~27GB weights/GPU, comfortable for LoRA training @@ -139,4 +137,3 @@ applications: env_vars: TWINKLE_TRUST_REMOTE_CODE: "0" TWINKLE_LONG_POLL_TIMEOUT: "120" - TWINKLE_FAIL_FAST: "0" diff --git a/cookbook/client/server/megatron/server_config_4b.yaml b/cookbook/client/server/megatron/server_config_4b.yaml index 9bdbd5e72..7eed4699d 100644 --- a/cookbook/client/server/megatron/server_config_4b.yaml +++ b/cookbook/client/server/megatron/server_config_4b.yaml @@ -31,9 +31,6 @@ applications: target_ongoing_requests: 128 # Target concurrent requests per replica ray_actor_options: num_cpus: 0.1 # CPU resources allocated to this actor - runtime_env: - env_vars: - TWINKLE_FAIL_FAST: "0" # 2. Model Service (commented out) - Would host the base model for training. # Uncomment and configure if you need a training model worker. @@ -71,7 +68,6 @@ applications: runtime_env: env_vars: TWINKLE_TRUST_REMOTE_CODE: "0" - TWINKLE_FAIL_FAST: "0" # 3. Sampler Service - Runs inference / sampling using vLLM engine # Used for generating text from the model (e.g., evaluating LoRA results). @@ -109,7 +105,6 @@ applications: runtime_env: env_vars: TWINKLE_TRUST_REMOTE_CODE: "0" - TWINKLE_FAIL_FAST: "0" # 4. Processor Service - name: processor @@ -132,6 +127,3 @@ applications: target_ongoing_requests: 128 ray_actor_options: num_cpus: 0.1 - runtime_env: - env_vars: - TWINKLE_FAIL_FAST: "0" diff --git a/cookbook/client/server/transformer/server_config.yaml b/cookbook/client/server/transformer/server_config.yaml index d3ddb2adb..b5d8497fd 100644 --- a/cookbook/client/server/transformer/server_config.yaml +++ b/cookbook/client/server/transformer/server_config.yaml @@ -49,9 +49,6 @@ applications: target_ongoing_requests: 128 # Target concurrent requests per replica ray_actor_options: num_cpus: 0.1 # CPU resources allocated to this actor - runtime_env: - env_vars: - TWINKLE_FAIL_FAST: "0" # 2. Model Service - Hosts the base model for training. - name: models-Qwen3.5-4B @@ -85,7 +82,6 @@ applications: runtime_env: env_vars: TWINKLE_TRUST_REMOTE_CODE: "1" - TWINKLE_FAIL_FAST: "0" # 3. Sampler Service - Runs inference / sampling using vLLM engine # Used for generating text from the model (e.g., evaluating LoRA results). @@ -122,7 +118,6 @@ applications: runtime_env: env_vars: TWINKLE_TRUST_REMOTE_CODE: "1" - TWINKLE_FAIL_FAST: "0" # 4. Processor Service - name: processor @@ -145,6 +140,3 @@ applications: target_ongoing_requests: 128 ray_actor_options: num_cpus: 0.1 - runtime_env: - env_vars: - TWINKLE_FAIL_FAST: "0" diff --git a/docs/source_en/Usage Guide/Server and Client/Server.md b/docs/source_en/Usage Guide/Server and Client/Server.md index f31209b07..5e67e37ca 100644 --- a/docs/source_en/Usage Guide/Server and Client/Server.md +++ b/docs/source_en/Usage Guide/Server and Client/Server.md @@ -450,3 +450,33 @@ twinkle-server check-config -c server_config.yaml | `use_megatron: false` | `backend: transformers` | Additionally, this refactor introduces two new top-level fields — `telemetry` and `persistence` — which did not exist before. Add them as needed. + +## Execution time bounds + +Every backend call has a finite time bound. `T` is the effective task execution +timeout: it equals `execution_timeout`, or `3600s` when that setting is `0`. +`asyncio.wait_for` uses `T`. The Ray wait uses `R`, which is a method's explicit +constant timeout when present and otherwise `T`. The default `T` is `1800s`. + +Two distinct bounds follow, and they must not be collapsed into one number: + +| Bound | Expression | Meaning | +|-------|------------|---------| +| Record-terminal bound | `queue_timeout + T` | After this, a task's future record is guaranteed to be in a terminal state (`completed`/`failed`). Use it for alerting thresholds and client polling total-timeout. | +| Resource-release bound | `Collect_Width × R` from execution start, or `queue_timeout + Collect_Width × R` from submission | After this, the executor thread and the in-flight model-actor call for that task are guaranteed to have finished. Use it for capacity planning. | + +`Collect_Width = len(self._actors) = world_size = tp × pp × dp` — the number of +futures each `remote_function` collection waits on per call. Evidence: +`LazyCollect._get_result` iterates `self._futures`, which come from +`_get_workers(self._actors, execute)` (`infra/__init__.py`), covering every actor — +not just the data-parallel width. On a `tp=8` deployment the execution-start +resource-release bound is therefore `8 × R`, not `R`. + +After the task record becomes terminal, the per-replica Admission_Gate can remain +closed for at most `max(0, Collect_Width × R − T)`: the record is already terminal, +but a leaked executor thread may still hold the gate until its `ray.get` returns or +raises. During that window newly arriving tasks fail fast with a `server`/503 error. + +Each persisted future stores its immutable `absolute_deadline` when it is created. +Cleanup therefore reaches the same decision regardless of which deployment process +holds the cleanup lease. diff --git "a/docs/source_zh/\344\275\277\347\224\250\346\214\207\345\274\225/\346\234\215\345\212\241\347\253\257\345\222\214\345\256\242\346\210\267\347\253\257/\346\234\215\345\212\241\347\253\257.md" "b/docs/source_zh/\344\275\277\347\224\250\346\214\207\345\274\225/\346\234\215\345\212\241\347\253\257\345\222\214\345\256\242\346\210\267\347\253\257/\346\234\215\345\212\241\347\253\257.md" index a4df7a2da..db71e41aa 100644 --- "a/docs/source_zh/\344\275\277\347\224\250\346\214\207\345\274\225/\346\234\215\345\212\241\347\253\257\345\222\214\345\256\242\346\210\267\347\253\257/\346\234\215\345\212\241\347\253\257.md" +++ "b/docs/source_zh/\344\275\277\347\224\250\346\214\207\345\274\225/\346\234\215\345\212\241\347\253\257\345\222\214\345\256\242\346\210\267\347\253\257/\346\234\215\345\212\241\347\253\257.md" @@ -450,3 +450,20 @@ twinkle-server check-config -c server_config.yaml | `use_megatron: false` | `backend: transformers` | 此外本次重构新增了 `telemetry` 和 `persistence` 两个顶层字段(旧版本中不存在),可按需添加。 + +## 执行时间上界 + +每一次 backend 调用都存在有限时间上界。`T` 是任务的有效 execution timeout:等于 task-queue 配置中的 `execution_timeout`;配置为 `0` 时取 `3600` 秒。`asyncio.wait_for` 使用 `T`。Ray 等待使用 `R`:方法显式声明 timeout 时取该常量,否则取 `T`。`T` 的默认值为 `1800` 秒。 + +由此派生出两个**不同**的上界,二者不得合成一个数: + +| 上界 | 表达式 | 含义 | +|------|--------|------| +| 记录终态上界 | `queue_timeout + T` | 超过它后,任务的 future 记录必处于终态(`completed`/`failed`)。用于设置告警阈值与客户端轮询总超时。 | +| 资源释放上界 | 从执行开始为 `Collect_Width × R`;从提交开始为 `queue_timeout + Collect_Width × R` | 超过它后,该任务占用的 executor 线程与 model actor 在飞调用必已结束。用于容量规划。 | + +`Collect_Width = len(self._actors) = world_size = tp × pp × dp`——即每次 `remote_function` 结果收集所等待的 future 个数。证据:`LazyCollect._get_result` 遍历的 `self._futures` 来自 `_get_workers(self._actors, execute)`(`infra/__init__.py`),覆盖全部 actor,而非 data-parallel 宽度。因此在 `tp=8` 的部署上,从执行开始的资源释放上界是 `8 × R` 而非 `R`。 + +任务记录进入终态后,per-replica 准入闸门额外保持关闭的最长时长为 `max(0, Collect_Width × R − T)`:此时记录已是终态,但泄漏的 executor 线程可能仍持有闸门,直到其 `ray.get` 返回或抛出。在该窗口内新到达的任务会以 `server`/503 错误快速失败。 + +每条持久化 future 在创建时写入不可变的 `absolute_deadline`,因此无论哪个 deployment 进程持有 cleanup lease,清理结果都由任务自身契约决定。 diff --git a/src/twinkle/infra/__init__.py b/src/twinkle/infra/__init__.py index 4758b3415..6d6ae8d07 100644 --- a/src/twinkle/infra/__init__.py +++ b/src/twinkle/infra/__init__.py @@ -10,7 +10,7 @@ from twinkle.notifier import Notifier, notify_exception from twinkle.utils import DeviceGroup, DeviceMesh, Platform, check_unsafe, framework_util, get_logger, requires -from .collectors import collect_tensor_dict +from .collectors import collect_tensor_dict as collect_tensor_dict logger = get_logger() @@ -530,7 +530,7 @@ def _run_continous_work(self, func_name: str, execute_method, workers, args, kwa try: ordered: List[Any] = [None] * batch_len for _, indices, ref in submitted: - part = ray.get(ref, timeout=ray_get_timeout) if ray_get_timeout else ray.get(ref) + part = ray.get(ref, timeout=ray_get_timeout) if ray_get_timeout is not None else ray.get(ref) if not isinstance(part, (list, tuple)) or len(part) != len(indices): raise TypeError(f'{func_name}: enable_continous_work needs one result per request, but a worker given ' f'{len(indices)} request(s) returned {type(part).__name__} of length ' @@ -740,7 +740,6 @@ def _get_device_mesh_param(args, kwargs): def _prepare_lazy_collect(args, kwargs): # if a worker received an actor handle, # lazy collect should be false to prevent any outer function receives an object ref - from ._ray import RayHelper if not os.environ.get('WORKER_NAME'): # If this is a driver return args, kwargs @@ -996,7 +995,10 @@ def remote_function(dispatch: Union[Literal['slice', 'all', 'slice_dp', 'last_pp sync: If True, use synchronous execution (execute_all_sync) instead of async. Required for methods with NCCL collective operations (e.g., Megatron forward_backward). lazy_collect: Do lazy collect, this boolean value decides whether this function needs lazy collect. If setting to None, it will follow the global setting. - timeout: Timeout in seconds for ray.get() when collecting results. Instance attribute ``_ray_get_timeout`` overrides this. + timeout: Timeout in seconds for ray.get() when collecting results. The decorator's + explicitly declared value takes priority; the instance attribute ``_ray_get_timeout`` + is the fallback for methods that declare none (``timeout if timeout is not None + else instance``). enable_continous_work: Route each request to the least busy worker instead of slicing the batch over all of them, and return the results in the caller's order. This is what lets a batch smaller than the worker @@ -1044,7 +1046,14 @@ def wrapper(self, *args, **kwargs) -> T1: else: # This is the driver from ._ray import RayHelper - execute_method = RayHelper.execute_all_async if not sync else RayHelper.execute_all_sync + + # Resolve the effective ray.get timeout before choosing execute_method: + # the decorator's explicit value wins, the instance attribute is the + # fallback. ``is not None`` (not ``or``) so that a decorator ``timeout=0`` + # is honored instead of falling back to unbounded waiting. + _rgt = timeout if timeout is not None else getattr(self, '_ray_get_timeout', None) + execute_method = RayHelper.execute_all_async if not sync else functools.partial( + RayHelper.execute_all_sync, timeout=_rgt) # Only classes whose workers run methods side by side need # this; elsewhere Ray already orders calls per actor. _concurrent_actor = bool(getattr(self, '_max_concurrency', None)) @@ -1060,8 +1069,7 @@ def wrapper(self, *args, **kwargs) -> T1: _batch_len = _cw_batch_len(args, kwargs) if _batch_len: return _run_continous_work(self, func.__name__, execute_method, _workers, args, kwargs, - _batch_len, - getattr(self, '_ray_get_timeout', None) or timeout) + _batch_len, _rgt) if RayHelper.has_ref(args, kwargs): # If has any object-ref, dispatch in worker, because we don't know the structure in the ref. # for example, dataloader returns any data list. @@ -1079,7 +1087,6 @@ def wrapper(self, *args, **kwargs) -> T1: # busy. _tracked_refs = _cw_register(self, func.__name__, result) if _concurrent_actor else [] # This is a result future, call it to get the actual result - _rgt = getattr(self, '_ray_get_timeout', None) or timeout result_func = RayHelper.do_get_and_collect_func( _collect_func, collect, result, device_mesh, timeout=_rgt) _local_lazy_collect = _lazy_collect @@ -1090,13 +1097,13 @@ def wrapper(self, *args, **kwargs) -> T1: if func.__name__ == '__len__': # Get the first result and ignore the `lazy_collect` import ray - return ray.get(result[0]) + return ray.get(result[0], timeout=_rgt) if func.__name__ == '__next__': import ray for _res in result: # raise when any worker raises StopIteration - stop = ray.get(_res[1]) + stop = ray.get(_res[1], timeout=_rgt) if stop: raise StopIteration() result = [_res[0] for _res in result] diff --git a/src/twinkle/infra/_ray/ray_helper.py b/src/twinkle/infra/_ray/ray_helper.py index ffd4e1a42..4cc5f6a66 100644 --- a/src/twinkle/infra/_ray/ray_helper.py +++ b/src/twinkle/infra/_ray/ray_helper.py @@ -137,10 +137,17 @@ def is_worker(): return RayHelper.ray_inited() and ray._private.worker.global_worker.mode == ray._private.worker.WORKER_MODE @staticmethod - def execute_all_sync(method_name: str, workers_and_args: List[Tuple[Any, List[Any], Dict[str, Any]]]): - """Execute method and return results.""" + def execute_all_sync(method_name: str, workers_and_args: List[Tuple[Any, List[Any], Dict[str, Any]]], timeout=None): + """Execute method and return results. + + ``timeout`` is passed to ``ray.get(list, timeout=)``, whose semantics are + the **total** wall-clock time to collect the whole list -- different from + ``LazyCollect``'s per-future timing (see ``do_get_and_collect_func``). + The two paths are each bounded on their own; the total-time semantics here + are strictly tighter. + """ import ray - return ray.get(RayHelper.execute_all_async(method_name, workers_and_args)) + return ray.get(RayHelper.execute_all_async(method_name, workers_and_args), timeout=timeout) @staticmethod def execute_all_async(method_name: str, workers_and_args: List[Tuple[Any, List[Any], Dict[str, Any]]]): diff --git a/src/twinkle/loss/grpo.py b/src/twinkle/loss/grpo.py index e146edee6..2bf8d39fc 100644 --- a/src/twinkle/loss/grpo.py +++ b/src/twinkle/loss/grpo.py @@ -202,14 +202,16 @@ def _pad_and_align_to_batch( elif n_sample == n_pos: # Response-only form (e.g. old_logps from vLLM). result[i, pos] = sample - elif n_sample >= seq_len: - # Full-sequence form (e.g. ref_logps right-padded with ignore-value). - result[i, pos] = sample[:seq_len][mask[i]] + elif n_pos == 0 or (n_sample > 0 and pos[-1].item() < n_sample): + # Variable-length full-sequence form. The processor right-pads the + # batch, but per-sample RL fields from Tinker remain unpadded. They + # are valid when every selected mask position exists in this row. + result[i, pos] = sample[pos] else: raise AssertionError(f'data/mask length mismatch at sample {i}: ' f'n_pos={n_pos}, n_sample={n_sample}, seq_len={seq_len} ' - '(expected n_sample == n_pos for response-only form, ' - 'or n_sample >= seq_len for full-sequence form)') + '(expected n_sample == n_pos for response-only form, or all masked positions ' + 'to exist in the per-sample full-sequence form)') return result diff --git a/src/twinkle/model/megatron/megatron.py b/src/twinkle/model/megatron/megatron.py index 851529c6a..4e4e2bbd8 100644 --- a/src/twinkle/model/megatron/megatron.py +++ b/src/twinkle/model/megatron/megatron.py @@ -36,7 +36,6 @@ from twinkle.processor import InputProcessor from twinkle.template import Template from twinkle.utils import construct_class, get_logger, selective_log_softmax -from twinkle.utils.nccl_safe import _is_fail_fast from ._mindspeed_runtime import ensure_mindspeed_adaptor_patched from .strategy import MegatronStrategy @@ -420,62 +419,47 @@ def forward_step_func(data_iterator, model): embeddings = None _loss_instance = loss_instance is_last_pp = mpu.is_pipeline_last_stage(False, unwrapped_model.vp_stage) - try: - if task == 'embedding': - # MegatronEmbeddingPatch already pooled output to [n_seqs, hidden] on last PP stage. - if is_last_pp: - embeddings = output_tensor - elif labels is not None and is_last_pp: - _loss_require_logps = getattr(_loss_instance, 'require_logps', True) - _loss_require_entropy = getattr(_loss_instance, 'require_entropy', False) - _packed = batch.get('packed_seq_params') - cu_seqlens_q = getattr(_packed, 'cu_seqlens_q', None) if _packed is not None else None - if _loss_require_logps: - loss_mask = (labels != -100).bool() - masked_labels = labels.clone() - masked_labels[~loss_mask] = 0 - output_tensor.div_(temperature) - if _loss_require_entropy: - logps, entropies = selective_log_softmax(output_tensor, masked_labels, return_entropy=True) - else: - logps = selective_log_softmax(output_tensor, masked_labels) - # Reconstruct full-length tensors from CP-split shards - logps = processor.postprocess_tensor_cp(logps, cu_seqlens=cu_seqlens_q) - if entropies is not None: - entropies = processor.postprocess_tensor_cp(entropies, cu_seqlens=cu_seqlens_q) - batch['labels'] = processor.postprocess_tensor_cp(labels, cu_seqlens=cu_seqlens_q) - if completion_mask is not None: - # Same index space as labels, so it needs the same CP reassembly. - batch['completion_mask'] = processor.postprocess_tensor_cp( - completion_mask, cu_seqlens=cu_seqlens_q) - if 'position_ids' in batch: - pos = batch['position_ids'] - if pos.dim() == 3: - pos = pos[0] # [2/3, 1, seq] → [1, seq] - batch['position_ids'] = processor.postprocess_tensor_cp(pos, cu_seqlens=cu_seqlens_q) - # Unpack packed sequences into per-sequence batch format - _outputs = {'logps': logps} + if task == 'embedding': + # MegatronEmbeddingPatch already pooled output to [n_seqs, hidden] on last PP stage. + if is_last_pp: + embeddings = output_tensor + elif labels is not None and is_last_pp: + _loss_require_logps = getattr(_loss_instance, 'require_logps', True) + _loss_require_entropy = getattr(_loss_instance, 'require_entropy', False) + _packed = batch.get('packed_seq_params') + cu_seqlens_q = getattr(_packed, 'cu_seqlens_q', None) if _packed is not None else None + if _loss_require_logps: + loss_mask = (labels != -100).bool() + masked_labels = labels.clone() + masked_labels[~loss_mask] = 0 + output_tensor.div_(temperature) + if _loss_require_entropy: + logps, entropies = selective_log_softmax(output_tensor, masked_labels, return_entropy=True) + else: + logps = selective_log_softmax(output_tensor, masked_labels) + # Reconstruct full-length tensors from CP-split shards + logps = processor.postprocess_tensor_cp(logps, cu_seqlens=cu_seqlens_q) if entropies is not None: - _outputs['entropies'] = entropies - if hasattr(_loss_instance, 'require_logits') and _loss_instance.require_logits: - _outputs['logits'] = output_tensor - batch, _outputs = processor.unpack_packed_sequences(batch, _outputs) - logps = _outputs['logps'] - entropies = _outputs.get('entropies', None) - unpacked_logits = _outputs.get('logits', None) - except Exception as e: - # Data processing error (e.g. unpack_packed_sequences dimension mismatch). - # Must catch here inside the scheduler to prevent exception escaping - # and breaking PP P2P communication → NCCL hang. - if _is_fail_fast(): - raise - logger.warning('[nccl_safe] forward_step_func data processing error: ' - '%s: %s', - type(e).__name__, e) - logps = None - unpacked_logits = None - entropies = None - embeddings = None + entropies = processor.postprocess_tensor_cp(entropies, cu_seqlens=cu_seqlens_q) + batch['labels'] = processor.postprocess_tensor_cp(labels, cu_seqlens=cu_seqlens_q) + if completion_mask is not None: + # Same index space as labels, so it needs the same CP reassembly. + batch['completion_mask'] = processor.postprocess_tensor_cp(completion_mask, cu_seqlens=cu_seqlens_q) + if 'position_ids' in batch: + pos = batch['position_ids'] + if pos.dim() == 3: + pos = pos[0] # [2/3, 1, seq] → [1, seq] + batch['position_ids'] = processor.postprocess_tensor_cp(pos, cu_seqlens=cu_seqlens_q) + # Unpack packed sequences into per-sequence batch format + _outputs = {'logps': logps} + if entropies is not None: + _outputs['entropies'] = entropies + if hasattr(_loss_instance, 'require_logits') and _loss_instance.require_logits: + _outputs['logits'] = output_tensor + batch, _outputs = processor.unpack_packed_sequences(batch, _outputs) + logps = _outputs['logps'] + entropies = _outputs.get('entropies', None) + unpacked_logits = _outputs.get('logits', None) return output_tensor, partial( post_loss_function, inputs=batch, @@ -883,7 +867,7 @@ def clip_grad_and_step(self, max_grad_norm: float = 1.0, norm_type=2, **kwargs): self.zero_grad(**kwargs) self.lr_step(**kwargs) - @remote_function(dispatch='all', collect='first', sync=True) + @remote_function(dispatch='all', collect='first', sync=True, timeout=3600) def save(self, name: Optional[str] = None, output_dir: Optional[str] = None, @@ -1486,7 +1470,7 @@ def _patch_adapter(self, adapter_name: str, config_or_dir: Union[PeftConfig, str self._default_tokenizer = self.optimizer_group[adapter_name].template.processor self.active_group = adapter_name - @remote_function(dispatch='all', sync=True) + @remote_function(dispatch='all', sync=True, timeout=3600) def add_adapter_to_model( self, adapter_name: str, diff --git a/src/twinkle/model/megatron/multi_lora_megatron.py b/src/twinkle/model/megatron/multi_lora_megatron.py index ebda91501..8af62df4b 100644 --- a/src/twinkle/model/megatron/multi_lora_megatron.py +++ b/src/twinkle/model/megatron/multi_lora_megatron.py @@ -291,7 +291,7 @@ def _load_multi_lora_optimizer(self, checkpoint_dir: str, adapter_name: str = '' if optimizer_config is not None and 'iteration' in state_dict: optimizer_config.cur_step = state_dict['iteration'] - @remote_function(dispatch='all', collect='first', sync=True) + @remote_function(dispatch='all', collect='first', sync=True, timeout=3600) def save(self, name, output_dir: Optional[str] = None, interval=1, **kwargs): adapter_name = kwargs.pop('adapter_name', None) self._check_adapter_valid(adapter_name) @@ -372,7 +372,7 @@ def load(self, name: str, output_dir: Optional[str] = None, **kwargs): if dist.is_initialized(): dist.barrier() - @remote_function(dispatch='all', collect='first', sync=True) + @remote_function(dispatch='all', collect='first', sync=True, timeout=3600) def resume_from_checkpoint(self, checkpoint_dir, *, resume_only_model=False, **kwargs): adapter_name = kwargs.pop('adapter_name', None) self._check_adapter_valid(adapter_name) @@ -403,7 +403,7 @@ def get_state_dict(self, **kwargs): self._check_adapter_valid(kwargs.get('adapter_name')) return self.multi_adapter.get_state_dict(**kwargs) - @remote_function(dispatch='all', sync=True) + @remote_function(dispatch='all', sync=True, timeout=3600) def add_adapter_to_model( self, adapter_name: str, diff --git a/src/twinkle/model/multi_lora.py b/src/twinkle/model/multi_lora.py index ff7766102..6eb80ed23 100644 --- a/src/twinkle/model/multi_lora.py +++ b/src/twinkle/model/multi_lora.py @@ -193,12 +193,14 @@ def _after(_module): _before(_module) else: _before(self.module) - yield adapter_name - if isinstance(self.module, list): - for _module in self.module: - _after(_module) - else: - _after(self.module) + try: + yield adapter_name + finally: + if isinstance(self.module, list): + for _module in self.module: + _after(_module) + else: + _after(self.module) # self.deactivate_adapter() def check_length( diff --git a/src/twinkle/model/optimizer_group.py b/src/twinkle/model/optimizer_group.py index f5177d672..150e694bc 100644 --- a/src/twinkle/model/optimizer_group.py +++ b/src/twinkle/model/optimizer_group.py @@ -48,12 +48,6 @@ class BaseOptimizerGroup: _device_mesh: DeviceMesh = None _last_grad_norm: float = 0.0 - def __setattr__(self, name, value): - if name == 'loss_instance' and value is not None: - from twinkle.utils.nccl_safe import safe_loss - value = safe_loss(value) - super().__setattr__(name, value) - def do_grad_sync(self, gradient_accumulation_steps: Optional[int] = None) -> bool: if gradient_accumulation_steps is None: gradient_accumulation_steps = self.gradient_accumulation_steps diff --git a/src/twinkle/model/transformers/transformers.py b/src/twinkle/model/transformers/transformers.py index 0087cc7f4..fcf856f86 100644 --- a/src/twinkle/model/transformers/transformers.py +++ b/src/twinkle/model/transformers/transformers.py @@ -1557,7 +1557,7 @@ def _restore_training_state(self, checkpoint_dir, *, adapter_name=''): return trainer_state - @remote_function(dispatch='all', collect='first', sync=True) + @remote_function(dispatch='all', collect='first', sync=True, timeout=3600) def resume_from_checkpoint(self, checkpoint_dir, *, resume_only_model=False, **kwargs): adapter_name = kwargs.get('adapter_name', '') diff --git a/src/twinkle/sampler/vllm_sampler/vllm_sampler.py b/src/twinkle/sampler/vllm_sampler/vllm_sampler.py index 7877ddb55..6438fa660 100644 --- a/src/twinkle/sampler/vllm_sampler/vllm_sampler.py +++ b/src/twinkle/sampler/vllm_sampler/vllm_sampler.py @@ -251,6 +251,14 @@ async def _sample_single( else: feat['input_ids'] = response.prompt_token_ids feat['labels'] = [-100] * len(response.prompt_token_ids) + # A sampling prompt (e.g. a tinker ModelInput) carries input_ids but no labels; + # concat_input_feature would then derive a zero-length prefix completion_mask + # and raise. The prompt is pure context, so materialise aligned all-context + # labels when they are missing or length-mismatched. Present, aligned labels are + # left untouched (preserving provenance), and the logprobs-only path -- which + # never concatenates -- is not touched. + if not logprobs_only and 'input_ids' in feat and len(feat.get('labels') or []) != len(feat['input_ids']): + feat['labels'] = [-100] * len(feat['input_ids']) sequences = [] for seq in response.sequences: if logprobs_only: @@ -495,7 +503,7 @@ def unload_adapter_paths(self, adapter_paths: list[str]) -> None: """Unload policy snapshots from vLLM and clear cached requests.""" self._run_in_loop(self.engine.unload_lora_paths(adapter_paths)) - @remote_function(dispatch='all', collect='first', lazy_collect=False) + @remote_function(dispatch='all', collect='first', lazy_collect=False, timeout=3600) def load_full_weights_from_path(self, path: Optional[str] = None) -> int: """Load a full (non-LoRA) HF checkpoint into the engine's base model. diff --git a/src/twinkle/server/gateway/tinker_handlers.py b/src/twinkle/server/gateway/tinker_handlers.py index 9ef98ce15..af2e8adb0 100644 --- a/src/twinkle/server/gateway/tinker_handlers.py +++ b/src/twinkle/server/gateway/tinker_handlers.py @@ -19,6 +19,7 @@ from twinkle.hub import HubOperation from twinkle.server.checkpoint import create_checkpoint_manager, create_training_run_manager +from twinkle.server.utils.task_errors import error_payload_from_stored from twinkle.server.utils.task_queue import QueueState from twinkle.server.utils.validation import get_token_from_request from twinkle.utils.logger import get_logger @@ -119,8 +120,8 @@ async def retrieve_future(request: Request, } if status == 'failed': - result = record.get('result', {}) - return {'error': result.get('error', 'Unknown error'), 'category': result.get('category', 'Server')} + payload = error_payload_from_stored(record.get('result'), request_id=request_id) + return payload.model_dump(mode='json', exclude_none=True) result = record.get('result') if result is None: diff --git a/src/twinkle/server/launcher/env_propagation.py b/src/twinkle/server/launcher/env_propagation.py index 3b0c9b956..4d4a7a474 100644 --- a/src/twinkle/server/launcher/env_propagation.py +++ b/src/twinkle/server/launcher/env_propagation.py @@ -21,10 +21,6 @@ 'TWINKLE_MODEL_ID_ALIASES', ) -# NCCL-safe env var keys: controls fault tolerance behavior in distributed -# training (safe_loss / @nccl_safe). Must reach model worker actors. -NCCL_SAFE_ENV_KEYS: tuple[str, ...] = ('TWINKLE_FAIL_FAST', ) - def build_telemetry_env_vars() -> dict[str, str]: """Collect telemetry env vars from ``os.environ`` for worker propagation.""" @@ -37,15 +33,9 @@ def build_persistence_env_vars() -> dict[str, str]: return {k: os.environ[k] for k in PERSISTENCE_ENV_KEYS if k in os.environ} -def build_nccl_safe_env_vars() -> dict[str, str]: - """Collect NCCL-safe env vars from ``os.environ`` for worker propagation.""" - return {k: os.environ[k] for k in NCCL_SAFE_ENV_KEYS if k in os.environ} - - def build_propagated_env_vars() -> dict[str, str]: """Aggregate all env vars that must reach Ray worker processes.""" merged: dict[str, str] = {} merged.update(build_telemetry_env_vars()) merged.update(build_persistence_env_vars()) - merged.update(build_nccl_safe_env_vars()) return merged diff --git a/src/twinkle/server/model/app.py b/src/twinkle/server/model/app.py index b3589ffee..0b73cbc60 100644 --- a/src/twinkle/server/model/app.py +++ b/src/twinkle/server/model/app.py @@ -7,7 +7,8 @@ """ from __future__ import annotations -from fastapi import FastAPI, Request +import asyncio +from fastapi import FastAPI, HTTPException, Request from ray import serve from ray.serve.config import RequestRouterConfig from typing import Any @@ -131,9 +132,18 @@ async def __init__(self, from twinkle.server.data_plane import DataPlaneProxy self.data_plane = DataPlaneProxy(data_plane_url) self._replica_registered = False - - # Initialize mixins - self._init_task_queue(queue_config, deployment_name='Model') + self._model_unhealthy = False + self._health_probe_task = None + + actors = getattr(self.model, '_actors', None) + self._init_task_queue( + queue_config, + deployment_name='Model', + enable_admission_gate=True, + on_backend_timeout=self._probe_after_timeout, + collect_width=len(actors) if actors else 1, + ) + self.model._ray_get_timeout = self._task_queue_config.effective_execution_timeout self._init_adapter_manager(**(adapter_config or {})) await self._register_replica_on_startup() # Note: countdown task is started lazily in _ensure_sticky() @@ -152,6 +162,7 @@ async def _register_replica_on_startup(self) -> None: """Register this replica's capacity before Ray Serve marks it ready.""" if not self._replica_registered: await self.state.register_replica(self.replica_id, self.max_loras) + await self.state.touch_replica_last_seen(self.replica_id) self._replica_registered = True @serve.multiplexed(max_num_models_per_replica=5) @@ -166,9 +177,11 @@ async def _ensure_sticky(self): async def _on_request_start(self, request: Request) -> str: await self._ensure_sticky() + await self.state.touch_replica_last_seen(self.replica_id) await self._ensure_state_cleanup_started() - token = get_token_from_request(request) - return token + if self._model_unhealthy: + raise HTTPException(status_code=503, detail='Model actors are unavailable') + return get_token_from_request(request) async def shutdown(self) -> None: """Explicit async cleanup — called via FastAPI shutdown event.""" @@ -176,32 +189,49 @@ async def shutdown(self) -> None: await self.state.unregister_replica(self.replica_id) except Exception: pass + await self.shutdown_task_queue() await self.data_plane.close() - def check_model_health(self) -> dict: - """Probe model actors liveness via a lightweight ping. - - Returns a dict with 'healthy' (bool) and 'detail' (str). - If the model actors are dead (e.g. OOM/SIGSEGV), the ping call - will raise RayActorError, signalling the watchdog to restart. - """ + async def _run_model_health_probe(self) -> dict: try: - result = self.model.ping() + result = await self.call_backend(self.model.ping, admit=False) if result is True: + self._model_unhealthy = False return {'healthy': True, 'detail': 'model actors alive'} + self._model_unhealthy = True return {'healthy': False, 'detail': f'unexpected ping result: {result}'} except Exception as e: + self._model_unhealthy = True return {'healthy': False, 'detail': f'model actor unreachable: {e}'} + async def check_model_health(self) -> dict: + """Run one coalesced actor probe outside the event loop.""" + current = getattr(self, '_health_probe_task', None) + if current is None or current.done(): + self._health_probe_task = asyncio.create_task(self._run_model_health_probe()) + return await asyncio.shield(self._health_probe_task) + + def mark_unhealthy(self) -> None: + """Flag the deployment unhealthy; /healthz returns 503 until a probe recovers it.""" + self._model_unhealthy = True + + async def _probe_after_timeout(self) -> None: + """Fired by ComputeWorker on a backend timeout: probe and log liveness (R3#2).""" + result = await self.check_model_health() + logger.warning('[Model] post-timeout liveness probe: %s', result) + async def _cleanup_adapter(self, adapter_name: str) -> None: if self.get_resource_info(adapter_name): - self.clear_resource_state(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() + # Takes the Admission_Gate: this path is driven by the background + # countdown and never enters Task_Queue, so the gate is what keeps + # it from colliding with an in-flight training call. + await self.call_backend(self.model.reload_initial_weights) else: - self.model.remove_adapter(adapter_name) + await self.call_backend(self.model.remove_adapter, adapter_name) + self.clear_resource_state(adapter_name) self.unregister_resource(adapter_name) await self.state.unload_model(adapter_name) diff --git a/src/twinkle/server/model/backends/megatron_model.py b/src/twinkle/server/model/backends/megatron_model.py index 56f9acc98..15dd3470c 100644 --- a/src/twinkle/server/model/backends/megatron_model.py +++ b/src/twinkle/server/model/backends/megatron_model.py @@ -11,7 +11,6 @@ """ import torch from tinker import types -from typing import TYPE_CHECKING, Any, List, Optional, Tuple, Union from twinkle import remote_class, remote_function from twinkle.data_format import InputFeature, Trajectory @@ -33,7 +32,7 @@ class in the MRO. For full-parameter training the ``adapter_name`` is the """ @remote_function(dispatch='slice_dp', collect=collect_forward_backward_results, sync=True) - @nccl_safe_megatron(tinker=True) + @nccl_safe_megatron def tinker_forward_backward(self, *, inputs: list[types.Datum], adapter_name: str, loss_fn: str, **kwargs): """Combined forward and backward pass.""" self._tinker_setup_loss(loss_fn, inputs, adapter_name, kwargs) @@ -54,7 +53,7 @@ def tinker_forward_backward(self, *, inputs: list[types.Datum], adapter_name: st return [results, loss] @remote_function(dispatch='slice_dp', collect=collect_forward_backward_results) - @nccl_safe_megatron(tinker=True) + @nccl_safe_megatron def tinker_forward_only(self, *, inputs: list[types.Datum], adapter_name: str = None, **kwargs): """Forward pass without gradient computation.""" template = self.get_template(adapter_name) @@ -102,7 +101,7 @@ def tinker_calculate_metric(self, is_training, **kwargs): metric = super().calculate_metric(is_training, **kwargs) return clean_metrics(metric) - @remote_function(dispatch='all', sync=True) + @remote_function(dispatch='all', sync=True, timeout=3600) def tinker_load(self, checkpoint_dir: str, **kwargs): """Load checkpoint with token-based isolation support.""" token = kwargs.pop('token', None) @@ -121,7 +120,7 @@ def tinker_load(self, checkpoint_dir: str, **kwargs): # ------------------------------------------------------------------ @remote_function(dispatch='slice_dp', collect=collect_tensor_dict) - @nccl_safe_megatron(forward_only=True) + @nccl_safe_megatron def forward_only(self, *, inputs: InputFeature | list[InputFeature] | Trajectory | list[Trajectory], **kwargs): """Forward-only for twinkle-native clients (InputFeature/Trajectory I/O).""" output = super().forward_only(inputs=inputs, **kwargs) @@ -135,7 +134,7 @@ def forward_backward(self, *, inputs: InputFeature | list[InputFeature] | Trajec output = super().forward_backward(inputs=inputs, **kwargs) return to_cpu_safe_output(output) - @remote_function(collect='first', lazy_collect=False) + @remote_function(collect='first', lazy_collect=False, sync=True, timeout=4) def ping(self) -> bool: """Lightweight liveness probe for watchdog health checks.""" return True diff --git a/src/twinkle/server/model/backends/mock_model.py b/src/twinkle/server/model/backends/mock_model.py index 7bf9866cc..c9bc12ff2 100644 --- a/src/twinkle/server/model/backends/mock_model.py +++ b/src/twinkle/server/model/backends/mock_model.py @@ -240,7 +240,7 @@ def remove_adapter(self, adapter_name: str) -> None: def has_adapter(self, adapter_name: str) -> bool: return adapter_name in self._adapters - @remote_function(collect='first', lazy_collect=False) + @remote_function(collect='first', lazy_collect=False, sync=True, timeout=4) def ping(self) -> bool: """Lightweight liveness probe for watchdog health checks.""" return True diff --git a/src/twinkle/server/model/backends/transformers_model.py b/src/twinkle/server/model/backends/transformers_model.py index 8dc503bb0..ff709a692 100644 --- a/src/twinkle/server/model/backends/transformers_model.py +++ b/src/twinkle/server/model/backends/transformers_model.py @@ -13,7 +13,6 @@ (InputFeature/Trajectory-based I/O) via /twinkle/* endpoints. """ from tinker import types -from typing import List, Union from twinkle import remote_class, remote_function from twinkle.data_format import InputFeature, Trajectory @@ -23,7 +22,6 @@ 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 class _TransformersTinkerCompatMixin(TwinkleCompatModelBase): @@ -48,7 +46,6 @@ def tinker_forward_only(self, *, inputs: list[types.Datum], adapter_name: str = return [results, 0.0] @remote_function(dispatch='slice_dp', collect=collect_forward_backward_results) - @nccl_safe(tinker=True) def tinker_forward_backward(self, *, inputs: list[types.Datum], adapter_name: str, loss_fn: str, **kwargs): self._tinker_setup_loss(loss_fn, inputs, adapter_name, kwargs) template = self.get_template(adapter_name) @@ -107,14 +104,13 @@ def forward_only(self, *, inputs: InputFeature | list[InputFeature] | Trajectory return to_cpu_safe_output(output) @remote_function(dispatch='slice_dp', collect=collect_tensor_dict) - @nccl_safe def forward_backward(self, *, inputs: InputFeature | list[InputFeature] | Trajectory | list[Trajectory], **kwargs): """Forward+backward for twinkle-native clients (InputFeature/Trajectory I/O).""" self._normalize_ref_outputs(kwargs) output = super().forward_backward(inputs=inputs, **kwargs) return to_cpu_safe_output(output) - @remote_function(collect='first', lazy_collect=False) + @remote_function(collect='first', lazy_collect=False, sync=True, timeout=4) def ping(self) -> bool: """Lightweight liveness probe for watchdog health checks.""" return True diff --git a/src/twinkle/server/model/tinker_handlers.py b/src/twinkle/server/model/tinker_handlers.py index c3df70efc..d868a0d20 100644 --- a/src/twinkle/server/model/tinker_handlers.py +++ b/src/twinkle/server/model/tinker_handlers.py @@ -10,9 +10,8 @@ import traceback from collections.abc import Callable from fastapi import Depends, FastAPI, Request -from peft import LoraConfig from tinker import types -from typing import TYPE_CHECKING, Any +from typing import TYPE_CHECKING if TYPE_CHECKING: from .app import ModelManagement @@ -20,6 +19,7 @@ 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.server.utils.task_queue.types import UserTaskError from twinkle.utils.logger import get_logger logger = get_logger() @@ -45,17 +45,11 @@ async def _create_adapter(): 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, - ) + raise UserTaskError('This deployment runs in full-parameter (exclusive) mode; do not pass ' + 'lora_config. Use create_full_training_client (or omit lora_config).') 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, - ) + raise UserTaskError('This deployment runs in LoRA mode; lora_config is required. ' + 'Use create_lora_training_client.') # Exclusive full-parameter training: reject early (before touching # state) if another tenant already holds the deployment. if self.is_full_mode: @@ -69,34 +63,33 @@ async def _create_adapter(): 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) + await self.call_backend( + self.model.set_template, template, adapter_name=model_adapter, model_id=self.base_model) + await self.call_backend(self.model.set_processor, 'InputProcessor', adapter_name=model_adapter) + await self.call_backend(self.model.set_optimizer, 'Adam', adapter_name=model_adapter) self.set_resource_state(adapter_name, 'grad_ready', False) else: - # TODO: Make LoraConfig more flexible + from peft import LoraConfig lora_cfg = LoraConfig(r=body.lora_config.rank, target_modules='all-linear') 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) - 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) + await self.call_backend( + self.model.add_adapter_to_model, adapter_name=adapter_name, config_or_dir=lora_cfg) + await self.call_backend( + self.model.set_template, template, adapter_name=adapter_name, model_id=self.base_model) + await self.call_backend(self.model.set_processor, 'InputProcessor', adapter_name=adapter_name) + await self.call_backend(self.model.set_optimizer, 'Adam', adapter_name=adapter_name) self.set_resource_state(adapter_name, 'grad_ready', False) 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) + raise UserTaskError(str(e)) from e except Exception: if _model_id: adapter_name = self.get_adapter_name(adapter_name=_model_id) await self._cleanup_adapter(adapter_name) logger.error(traceback.format_exc()) - return types.RequestFailedResponse( - error=traceback.format_exc(), - category=types.RequestErrorCategory.Server, - ) + raise return await self.schedule_task(_create_adapter, token=token, task_type='create_model') @@ -147,8 +140,8 @@ async def _do_forward(): 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=model_adapter, **loss_fn_config) + output, loss = await self.call_backend( + self.model.tinker_forward_only, inputs=datum_list, adapter_name=model_adapter, **loss_fn_config) return types.ForwardBackwardOutput( loss_fn_output_type='CrossEntropyLossReturn', loss_fn_outputs=output, @@ -156,10 +149,7 @@ async def _do_forward(): ) except Exception: logger.error(traceback.format_exc()) - return types.RequestFailedResponse( - error=traceback.format_exc(), - category=types.RequestErrorCategory.Server, - ) + raise datum_list = body.forward_input.data input_tokens = sum(len(d.model_input.to_ints()) for d in datum_list) @@ -190,8 +180,12 @@ async def _do_forward_backward(): 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=model_adapter, loss_fn=loss_fn, **loss_fn_config) + output, loss = await self.call_backend( + self.model.tinker_forward_backward, + 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) @@ -202,10 +196,7 @@ async def _do_forward_backward(): ) except Exception: logger.error(traceback.format_exc()) - return types.RequestFailedResponse( - error=traceback.format_exc(), - category=types.RequestErrorCategory.Server, - ) + raise datum_list = body.forward_backward_input.data input_tokens = sum(len(d.model_input.to_ints()) for d in datum_list) @@ -238,16 +229,15 @@ async def _do_optim(): 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=model_adapter) + await self.call_backend( + 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=model_adapter) + metrics = await self.call_backend( + self.model.tinker_calculate_metric, is_training=True, adapter_name=model_adapter) return types.OptimStepResponse(metrics=metrics) except Exception: logger.error(traceback.format_exc()) - return types.RequestFailedResponse( - error=traceback.format_exc(), - category=types.RequestErrorCategory.Server, - ) + raise return await self.schedule_task(_do_optim, model_id=body.model_id, token=token, task_type='optim_step') @@ -267,16 +257,17 @@ async def _do_save(): 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=model_adapter, save_optimizer=True) + await self.call_backend( + self.model.save, + 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: logger.error(traceback.format_exc()) - return types.RequestFailedResponse( - error=traceback.format_exc(), - category=types.RequestErrorCategory.Server, - ) + raise return await self.schedule_task(_do_save, model_id=body.model_id, token=token, task_type='save_weights') @@ -298,7 +289,8 @@ 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( + await self.call_backend( + self.model.save, name='latest', output_dir=save_dir, adapter_name=self.resolve_model_adapter_name(adapter_name), @@ -320,10 +312,7 @@ async def _do_save_for_sampler(): path=tinker_path, sampling_session_id=sampling_session_id) except Exception: logger.error(traceback.format_exc()) - return types.RequestFailedResponse( - error=traceback.format_exc(), - category=types.RequestErrorCategory.Server, - ) + raise return await self.schedule_task( _do_save_for_sampler, model_id=body.model_id, token=token, task_type='save_weights_for_sampler') @@ -341,7 +330,8 @@ async def _do_load(): assert self.model is not None, 'Model not loaded, please load model first' adapter_name = self.get_adapter_name(adapter_name=body.model_id) self.assert_resource_exists(adapter_name) - self.model.tinker_load( + await self.call_backend( + self.model.tinker_load, checkpoint_dir=body.path, load_optimizer=body.optimizer, adapter_name=self.resolve_model_adapter_name(adapter_name), @@ -350,9 +340,6 @@ async def _do_load(): return types.LoadWeightsResponse(path=body.path, type='load_weights') except Exception: logger.error(traceback.format_exc()) - return types.RequestFailedResponse( - error=traceback.format_exc(), - category=types.RequestErrorCategory.Server, - ) + raise return await self.schedule_task(_do_load, model_id=body.model_id, token=token, task_type='load_weights') diff --git a/src/twinkle/server/model/twinkle_handlers.py b/src/twinkle/server/model/twinkle_handlers.py index 582aa5c00..c743c5d3e 100644 --- a/src/twinkle/server/model/twinkle_handlers.py +++ b/src/twinkle/server/model/twinkle_handlers.py @@ -8,13 +8,11 @@ """ from __future__ import annotations -import asyncio import torch import traceback from collections.abc import Callable from fastapi import Depends, FastAPI, HTTPException, Request from pathlib import Path -from peft import LoraConfig from typing import TYPE_CHECKING, Any if TYPE_CHECKING: @@ -29,7 +27,6 @@ select_output_rows) 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 logger = get_logger() @@ -71,8 +68,8 @@ async def model_healthz( self: ModelManagement = Depends(self_fn), ) -> dict: """Deep health probe: pings underlying model actors to verify liveness.""" - result = self.check_model_health() - if not result['healthy']: + result = await self.check_model_health() + if self._model_unhealthy or not result['healthy']: from fastapi.responses import JSONResponse return JSONResponse(status_code=503, content=result) return result @@ -108,8 +105,11 @@ 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=self.resolve_model_adapter_name(adapter_name), **extra_kwargs) + ret = await self.call_backend( + 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] @@ -139,7 +139,8 @@ async def _task(): self.assert_resource_exists(adapter_name) raw_inputs, field_kwargs = await resolve_data_plane_model_inputs(body, self.data_plane) kwargs = merge_forward_kwargs(body.model_extra or {}, field_kwargs) - ret = self.model.forward( + ret = await self.call_backend( + self.model.forward, inputs=_parse_inputs(raw_inputs), adapter_name=adapter_name, **kwargs, @@ -193,8 +194,11 @@ 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=self.resolve_model_adapter_name(adapter_name), **extra_kwargs) + ret = await self.call_backend( + 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] @@ -222,7 +226,7 @@ async def _task(): raw_inputs, field_kwargs = await resolve_data_plane_model_inputs(body, self.data_plane) inputs = _parse_inputs(raw_inputs) kwargs = merge_forward_kwargs(body.model_extra or {}, field_kwargs) - ret = self.model.forward_only(inputs=inputs, adapter_name=adapter_name, **kwargs) + ret = await self.call_backend(self.model.forward_only, inputs=inputs, adapter_name=adapter_name, **kwargs) if body.output_ref is not None: rows = select_output_rows( ret, @@ -257,7 +261,8 @@ 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=self.resolve_model_adapter_name(adapter_name), **extra_kwargs) + ret = await self.call_backend( + self.model.calculate_loss, adapter_name=self.resolve_model_adapter_name(adapter_name), **extra_kwargs) return {'result': ret} return await run_task( @@ -271,7 +276,8 @@ 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=self.resolve_model_adapter_name(adapter_name), **extra_kwargs) + await self.call_backend( + 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')) @@ -299,8 +305,11 @@ 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=self.resolve_model_adapter_name(adapter_name), **extra_kwargs) + ret = await self.call_backend( + 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] @@ -330,7 +339,8 @@ async def _task(): self.assert_resource_exists(adapter_name) raw_inputs, field_kwargs = await resolve_data_plane_model_inputs(body, self.data_plane) kwargs = merge_forward_kwargs(body.model_extra or {}, field_kwargs) - ret = self.model.forward_backward( + ret = await self.call_backend( + self.model.forward_backward, inputs=_parse_inputs(raw_inputs), adapter_name=adapter_name, **kwargs, @@ -361,7 +371,8 @@ 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=self.resolve_model_adapter_name(adapter_name), **extra_kwargs) + ret = await self.call_backend( + self.model.clip_grad_norm, adapter_name=self.resolve_model_adapter_name(adapter_name), **extra_kwargs) return {'result': str(ret)} return await run_task( @@ -375,7 +386,8 @@ 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=self.resolve_model_adapter_name(adapter_name), **extra_kwargs) + await self.call_backend( + 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')) @@ -387,7 +399,8 @@ 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=self.resolve_model_adapter_name(adapter_name), **extra_kwargs) + await self.call_backend( + 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')) @@ -399,7 +412,8 @@ 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=self.resolve_model_adapter_name(adapter_name), **extra_kwargs) + await self.call_backend( + 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')) @@ -415,7 +429,8 @@ async def clip_grad_and_step( async def _task(): self.assert_resource_exists(adapter_name) extra_kwargs = body.model_extra or {} - self.model.clip_grad_and_step( + await self.call_backend( + self.model.clip_grad_and_step, max_grad_norm=body.max_grad_norm, norm_type=body.norm_type, adapter_name=self.resolve_model_adapter_name(adapter_name), @@ -437,8 +452,10 @@ 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=self.resolve_model_adapter_name(adapter_name), **extra_kwargs) + ret = await self.call_backend( + self.model.get_train_configs, + adapter_name=self.resolve_model_adapter_name(adapter_name), + **extra_kwargs) return {'result': ret} return await run_task( @@ -452,8 +469,11 @@ 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=self.resolve_model_adapter_name(adapter_name), **extra_kwargs) + await self.call_backend( + 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')) @@ -469,8 +489,11 @@ 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=self.resolve_model_adapter_name(adapter_name), **extra_kwargs) + await self.call_backend( + 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')) @@ -487,8 +510,11 @@ 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=self.resolve_model_adapter_name(adapter_name), **extra_kwargs) + await self.call_backend( + 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')) @@ -510,7 +536,8 @@ async def _task(): model_id=adapter_name, name=checkpoint_name, is_sampler=body.is_sampler) # For sampler weights the actual data is always written to 'latest/'. model_save_name = 'latest' if body.is_sampler else checkpoint_name - checkpoint_dir = self.model.save( + checkpoint_dir = await self.call_backend( + self.model.save, name=model_save_name, output_dir=save_dir, adapter_name=self.resolve_model_adapter_name(adapter_name), @@ -530,7 +557,8 @@ async def _task(): extra_kwargs = body.model_extra or {} checkpoint_manager = create_checkpoint_manager(token, client_type='twinkle') resolved = checkpoint_manager.resolve_load_path(body.name) - self.model.load( + await self.call_backend( + self.model.load, name=resolved.checkpoint_name, output_dir=resolved.checkpoint_dir, adapter_name=self.resolve_model_adapter_name(adapter_name), @@ -556,7 +584,8 @@ async def _task(): checkpoint_dir = ( Path(resolved.checkpoint_dir, resolved.checkpoint_name).as_posix() if resolved.checkpoint_dir else body.name) - ret = self.model.resume_from_checkpoint( + ret = await self.call_backend( + self.model.resume_from_checkpoint, checkpoint_dir, resume_only_model=body.resume_only_model, adapter_name=self.resolve_model_adapter_name(adapter_name), @@ -588,11 +617,7 @@ async def _task(): checkpoint_manager.get_ckpt_dir(model_id=model_id_to_load, checkpoint_id=checkpoint_id)) else: checkpoint_dir = body.checkpoint_dir - # Run blocking upload in thread pool so the event loop is not blocked. - # async_upload is intentionally ignored here: the task queue + client polling - # already provide the fire-and-forget / wait semantics without holding the - # HTTP connection open for the full duration of the upload. - await asyncio.to_thread( + await self.call_backend( self.model.upload_to_hub, checkpoint_dir=checkpoint_dir, hub_model_id=body.hub_model_id, @@ -640,6 +665,9 @@ async def add_adapter_to_model( raise HTTPException(status_code=400, detail=str(exc)) async def _task(): + from peft import LoraConfig + + from twinkle_client.common.serialize import deserialize_object config = deserialize_object(body.config) extra_kwargs = body.model_extra or {} training_run_manager = create_training_run_manager(token, client_type='twinkle') @@ -682,7 +710,7 @@ async def _task(): # 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) + await self.call_backend(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) @@ -703,11 +731,15 @@ async def apply_patch( adapter_name = _get_twinkle_adapter_name(request, body.adapter_name) async def _task(): + from twinkle_client.common.serialize import deserialize_object 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=self.resolve_model_adapter_name(adapter_name), **extra_kwargs) + await self.call_backend( + 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')) @@ -721,10 +753,12 @@ async def add_metric( adapter_name = _get_twinkle_adapter_name(request, body.adapter_name) async def _task(): + from twinkle_client.common.serialize import deserialize_object self.assert_resource_exists(adapter_name) extra_kwargs = body.model_extra or {} metric_cls = deserialize_object(body.metric_cls) - self.model.add_metric( + await self.call_backend( + self.model.add_metric, metric_cls, is_training=body.is_training, adapter_name=self.resolve_model_adapter_name(adapter_name), @@ -744,8 +778,11 @@ 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=self.resolve_model_adapter_name(adapter_name), **extra_kwargs) + await self.call_backend( + 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')) @@ -761,8 +798,11 @@ 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=self.resolve_model_adapter_name(adapter_name), **extra_kwargs) + await self.call_backend( + 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')) @@ -779,7 +819,8 @@ async def calculate_metric( async def _task(): self.assert_resource_exists(adapter_name) extra_kwargs = body.model_extra or {} - ret = self.model.calculate_metric( + ret = await self.call_backend( + self.model.calculate_metric, is_training=body.is_training, adapter_name=self.resolve_model_adapter_name(adapter_name), **extra_kwargs) @@ -800,7 +841,8 @@ 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=self.resolve_model_adapter_name(adapter_name), **extra_kwargs) + ret = await self.call_backend( + 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/processor/twinkle_handlers.py b/src/twinkle/server/processor/twinkle_handlers.py index d9da305ba..e8f8792c4 100644 --- a/src/twinkle/server/processor/twinkle_handlers.py +++ b/src/twinkle/server/processor/twinkle_handlers.py @@ -23,7 +23,6 @@ from twinkle.server.telemetry.tracing import traced_operation from twinkle.server.utils.validation import get_session_id_from_request, get_token_from_request from twinkle.utils.logger import get_logger -from twinkle_client.common.serialize import deserialize_object logger = get_logger() @@ -61,6 +60,7 @@ async def create( _kwargs.pop('remote_group', None) _kwargs.pop('device_mesh', None) + from twinkle_client.common.serialize import deserialize_object resolved_kwargs = {} for key, value in _kwargs.items(): if isinstance(value, str) and value.startswith('pid:'): @@ -107,6 +107,7 @@ async def call( assert function is not None, f'`{function_name}` not found in {processor.__class__}' assert hasattr(function, '_execute'), f'Cannot call inner method of {processor.__class__}' + from twinkle_client.common.serialize import deserialize_object resolved_kwargs = {} for key, value in _kwargs.items(): if isinstance(value, str) and value.startswith('pid:'): diff --git a/src/twinkle/server/sampler/app.py b/src/twinkle/server/sampler/app.py index 8941a40bf..3e7d01454 100644 --- a/src/twinkle/server/sampler/app.py +++ b/src/twinkle/server/sampler/app.py @@ -7,7 +7,6 @@ """ from __future__ import annotations -import asyncio from fastapi import FastAPI, Request from ray import serve from typing import Any @@ -103,12 +102,12 @@ def __init__(self, self.sampler_type = sampler_type self.model_id = model_id replica_context = serve.get_replica_context() - replica_id = replica_context.replica_id.unique_id + self.replica_id = replica_context.replica_id.unique_id sampler_kwargs: dict[str, Any] = { 'model_id': model_id, 'remote_group': self.device_group.name, - 'instance_id': replica_id, + 'instance_id': self.replica_id, } if sampler_type != 'mock': sampler_kwargs.update( @@ -127,14 +126,25 @@ def __init__(self, from twinkle.server.data_plane import DataPlaneProxy self.data_plane = DataPlaneProxy(data_plane_url) - # Initialize task queue mixin - self._init_task_queue(queue_config, deployment_name='Sampler') + actors = getattr(self.sampler, '_actors', None) + self._init_task_queue( + queue_config, + deployment_name='Sampler', + collect_width=len(actors) if actors else 1, + ) + self.sampler._ray_get_timeout = self._task_queue_config.effective_execution_timeout async def shutdown(self) -> None: - cancel_all = getattr(self.sampler, 'cancel_all_generations', None) - if callable(cancel_all): - await asyncio.to_thread(cancel_all) - await self.data_plane.close() + try: + cancel_all = getattr(self.sampler, 'cancel_all_generations', None) + if callable(cancel_all): + await self.call_backend(cancel_all) + finally: + try: + await self.state.unregister_replica(self.replica_id) + finally: + await self.shutdown_task_queue() + await self.data_plane.close() @serve.multiplexed(max_num_models_per_replica=5) async def _sticky_entry(self, sticky_key: str): @@ -146,9 +156,9 @@ async def _ensure_sticky(self): async def _on_request_start(self, request: Request) -> str: await self._ensure_sticky() + await self.state.touch_replica_last_seen(self.replica_id) await self._ensure_state_cleanup_started() - token = get_token_from_request(request) - return token + return get_token_from_request(request) def build_sampler_app(model_id: str, diff --git a/src/twinkle/server/sampler/tinker_handlers.py b/src/twinkle/server/sampler/tinker_handlers.py index de43b4b93..3f8790010 100644 --- a/src/twinkle/server/sampler/tinker_handlers.py +++ b/src/twinkle/server/sampler/tinker_handlers.py @@ -19,11 +19,28 @@ from twinkle.data_format import SamplingParams from twinkle.server.checkpoint import create_checkpoint_manager from twinkle.server.utils import get_template_for_model +from twinkle.server.utils.task_queue.types import UserTaskError from twinkle.utils.logger import get_logger logger = get_logger() +def _sampled_sequence(*, stop_reason, tokens, logprobs): + return types.SampledSequence( + stop_reason=stop_reason, + tokens=tokens, + logprobs=logprobs, + ) + + +def _sample_response(*, sequences, prompt_logprobs, topk_prompt_logprobs): + return types.SampleResponse( + sequences=sequences, + prompt_logprobs=prompt_logprobs, + topk_prompt_logprobs=topk_prompt_logprobs, + ) + + def _register_tinker_sampler_routes(app: FastAPI, self_fn: Callable[[], SamplerManagement]) -> None: """Register the tinker sampler route on the given FastAPI app. @@ -52,9 +69,9 @@ async def _do_sample(): # Set template for sampler based on model type template = get_template_for_model(self.model_id) - self.sampler.set_template(template, model_id=self.model_id) + await self.call_backend(self.sampler.set_template, template, model_id=self.model_id) # Reset prefix cache for new weights - self.sampler.reset_prefix_cache() + await self.call_backend(self.sampler.reset_prefix_cache) # Get model_path from body or sampling session model_path = body.model_path @@ -71,10 +88,7 @@ async def _do_sample(): # Base-model sampling is valid when no model_path was provided. if adapter_uri and not os.path.exists(adapter_uri): - return types.RequestFailedResponse( - error=f'Adapter URI {model_path} does not exist. Please check the model_path.', - category=types.RequestErrorCategory.User, - ) + raise UserTaskError(f'Adapter URI {model_path} does not exist. Please check the model_path.') # Convert tinker SamplingParams to twinkle SamplingParams if needed sampling_params = None @@ -85,6 +99,10 @@ async def _do_sample(): top_p=body.sampling_params.top_p, top_k=body.sampling_params.top_k, stop=body.sampling_params.stop, + # tinker 0.16.1 has no SamplingParams.logprobs field, but its + # SampledSequence contract and GRPO training require one + # chosen-token logprob per generated token. + logprobs=1, ) # A resolved checkpoint is either a LoRA adapter dir (has @@ -96,9 +114,10 @@ async def _do_sample(): 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) + await self.call_backend(self.sampler.load_full_weights_from_path, adapter_uri) - responses = self.sampler.sample( + responses = await self.call_backend( + self.sampler.sample, inputs=[prompt_inputs] * body.num_samples, sampling_params=sampling_params, adapter_path=lora_path, @@ -116,25 +135,26 @@ async def _do_sample(): flattened = [float(lp_list[0][1]) for lp_list in seq.logprobs if lp_list] except (IndexError, TypeError): flattened = [] - if flattened and len(flattened) == len(seq.logprobs): + if len(flattened) == len(seq.tokens): logprobs = flattened + else: + raise RuntimeError( + f'Sampler returned {len(flattened)} logprobs for {len(seq.tokens)} generated ' + 'tokens; refusing to return a misaligned Tinker SampledSequence.') tinker_sequences.append( - types.SampledSequence( + _sampled_sequence( stop_reason=seq.stop_reason, tokens=list(seq.tokens), logprobs=logprobs, )) - return types.SampleResponse( + return _sample_response( sequences=tinker_sequences, prompt_logprobs=responses[0].prompt_logprobs, topk_prompt_logprobs=responses[0].topk_prompt_logprobs, ) except Exception: logger.error(traceback.format_exc()) - return types.RequestFailedResponse( - error=traceback.format_exc(), - category=types.RequestErrorCategory.Server, - ) + raise input_tokens = len(body.prompt.to_ints()) return await self.schedule_task( diff --git a/src/twinkle/server/sampler/twinkle_handlers.py b/src/twinkle/server/sampler/twinkle_handlers.py index 609c25f21..c5428b055 100644 --- a/src/twinkle/server/sampler/twinkle_handlers.py +++ b/src/twinkle/server/sampler/twinkle_handlers.py @@ -15,8 +15,6 @@ from fastapi.responses import StreamingResponse from typing import TYPE_CHECKING -from twinkle_client.common.serialize import deserialize_object - if TYPE_CHECKING: from .app import SamplerManagement @@ -24,8 +22,9 @@ import twinkle_client.types as types from twinkle.data_format import InputFeature, SamplingParams, Trajectory -from twinkle.server.telemetry.correlation import MODEL_ID, TOKEN_ID +from twinkle.server.telemetry.correlation import MODEL_ID from twinkle.server.telemetry.tracing import traced_operation +from twinkle.server.utils.task_errors import task_error_payload from twinkle.server.utils.validation import get_session_id_from_request from twinkle.utils.logger import get_logger from twinkle_client.common.json_utils import json_safe @@ -132,31 +131,55 @@ def _submission_states(value) -> list[dict]: return value if isinstance(value, list) else [value] -async def _await_generation( - sampler, - submission_id: str, -): - """Poll an admitted generation without occupying the sampler admission queue.""" - collected = False +async def _stream_queue(q, sentinel, request_id: str, total_timeout: float, single_get_timeout: float = 60.0): + loop = asyncio.get_running_loop() + start = loop.time() try: + while True: + remaining = total_timeout - (loop.time() - start) + if remaining <= 0: + payload = task_error_payload( + 'sample_stream exceeded the execution time bound', request_id=request_id, error_code=504) + yield json.dumps(payload) + '\n' + break + try: + item = await asyncio.wait_for( + loop.run_in_executor(None, q.get), timeout=min(single_get_timeout, remaining)) + except asyncio.TimeoutError: + payload = task_error_payload( + 'sample_stream timed out waiting for the next token', request_id=request_id, error_code=504) + yield json.dumps(payload) + '\n' + break + if item == sentinel: + break + if isinstance(item, Exception): + payload = task_error_payload(f'{type(item).__name__}: {item}', request_id=request_id, error_code=500) + yield json.dumps(payload) + '\n' + break + delta, reason = item + yield json.dumps({'delta': delta, 'finish_reason': reason}) + '\n' + finally: + try: + q.shutdown(force=True) + except Exception: + pass + + +async def _await_generation(service: SamplerManagement, submission_id: str, timeout: float): + """Poll one admitted generation through the backend boundary.""" + collected = False + + async def poll(): + nonlocal collected poll_interval = 0.01 while True: try: - states = _submission_states(await asyncio.to_thread(sampler.get_generation_status, submission_id)) + states = _submission_states(await service.call_backend(service.sampler.get_generation_status, + submission_id)) except Exception as error: - # A pending read-only actor call can be cancelled by Ray while - # the generation submitted just above remains alive. Treating - # that as a generation failure makes the finally block discard - # otherwise valid rollout work. Retry only Ray's explicit task - # cancellation; actor death and application errors must still - # propagate immediately. from ray.exceptions import TaskCancelledError if not isinstance(error, TaskCancelledError): raise - logger.warning( - 'Generation status poll was cancelled; retrying submission %s', - submission_id, - ) await asyncio.sleep(poll_interval) poll_interval = min(poll_interval * 1.5, 0.25) continue @@ -168,15 +191,19 @@ async def _await_generation( error = failed.get('error') or failed.get('status', 'unknown failure') raise RuntimeError(f'generation {submission_id} failed: {error}') if states and all(state.get('status') == 'completed' for state in states): - responses = await asyncio.to_thread(sampler.collect_generation, submission_id) + responses = await service.call_backend(service.sampler.collect_generation, submission_id) collected = True return responses await asyncio.sleep(poll_interval) poll_interval = min(poll_interval * 1.5, 0.25) + + try: + return await asyncio.wait_for(poll(), timeout=timeout) finally: if not collected: try: - await asyncio.to_thread(sampler.cancel_generation, submission_id) + await asyncio.wait_for( + service.call_backend(service.sampler.cancel_generation, submission_id), timeout=4.0) except Exception: logger.warning('Failed to cancel generation %s', submission_id, exc_info=True) @@ -231,13 +258,13 @@ async def _task(): checkpoint_manager = create_checkpoint_manager(token, client_type='twinkle') _, resolved_uri = checkpoint_manager.parse_adapter_uri(body.adapter_uri) # Reset prefix cache only when new weights are loaded - self.sampler.reset_prefix_cache() + await self.call_backend(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) + await self.call_backend(self.sampler.load_full_weights_from_path, resolved_uri) # Parse inputs inputs = body.inputs @@ -259,7 +286,8 @@ async def _task(): params = SamplingParams.from_dict(body.sampling_params) # Sample - responses = self.sampler.sample( + responses = await self.call_backend( + self.sampler.sample, inputs, params, adapter_name=full_adapter_name, @@ -314,7 +342,7 @@ async def sample_to_data_plane( submission_id = uuid.uuid4().hex async def _admit(): - await asyncio.to_thread( + await self.call_backend( self.sampler.submit_generation, submission_id, inputs, @@ -337,7 +365,7 @@ async def _admit(): task_type='sample_admission', )) - responses = await _await_generation(self.sampler, submission_id) + responses = await _await_generation(self, submission_id, self._task_queue_config.effective_execution_timeout) rows, tags = _build_rollout_rows_and_tags( _to_sample_response_models(responses), group_ids=body.group_ids, @@ -367,7 +395,7 @@ async def unload_adapter_paths( resolved_paths.append(adapter_path) unload = getattr(self.sampler, 'unload_adapter_paths', None) if unload is not None: - unload(resolved_paths) + await self.call_backend(unload, resolved_paths) return {'status': 'ok'} @app.post('/twinkle/set_template', response_model=types.SetTemplateResponse) @@ -379,7 +407,7 @@ async def set_template( """Set the chat template for encoding Trajectory inputs.""" extra_kwargs = body.model_extra or {} with traced_operation('sampler.set_template'): - self.sampler.set_template(body.template_cls, **extra_kwargs) + await self.call_backend(self.sampler.set_template, body.template_cls, **extra_kwargs) return types.SetTemplateResponse() @app.post('/twinkle/add_adapter_to_sampler', response_model=types.AddAdapterResponse) @@ -396,7 +424,7 @@ async def add_adapter_to_sampler( config = LoraConfig(**body.config) if isinstance(body.config, dict) else body.config with traced_operation('sampler.add_adapter_to_sampler', attrs={MODEL_ID: self.model_id}): - self.sampler.add_adapter_to_sampler(full_adapter_name, config) + await self.call_backend(self.sampler.add_adapter_to_sampler, full_adapter_name, config) return types.AddAdapterResponse(adapter_name=full_adapter_name) @@ -406,10 +434,11 @@ async def apply_patch( body: types.ApplyPatchRequest, self: SamplerManagement = Depends(self_fn), ) -> None: + from twinkle_client.common.serialize import deserialize_object extra_kwargs = body.model_extra or {} patch_cls = deserialize_object(body.patch_cls) with traced_operation('sampler.apply_patch'): - self.sampler.apply_patch(patch_cls, **extra_kwargs) + await self.call_backend(self.sampler.apply_patch, patch_cls, **extra_kwargs) @app.post('/twinkle/sample_stream') async def sample_stream( @@ -437,11 +466,11 @@ async def sample_stream( from twinkle.server.checkpoint import create_checkpoint_manager checkpoint_manager = create_checkpoint_manager(token, client_type='twinkle') _, resolved_uri = checkpoint_manager.parse_adapter_uri(body.adapter_uri) - self.sampler.reset_prefix_cache() + await self.call_backend(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) + await self.call_backend(self.sampler.load_full_weights_from_path, resolved_uri) inputs = body.inputs if isinstance(inputs, list): @@ -464,8 +493,17 @@ async def sample_stream( from .backends import STREAM_SENTINEL + request_id = f'req_{uuid.uuid4().hex}' + actors = self.sampler._actors + if not actors: + + async def _no_actor_generator(): + payload = task_error_payload('No available sampler actor', request_id=request_id, error_code=503) + yield json.dumps(payload) + '\n' + + return StreamingResponse(_no_actor_generator(), media_type='application/x-ndjson') q = Queue(maxsize=128) - actor = self.sampler._actors[0] + actor = actors[0] actor.sample_stream_to_queue.remote( q, inputs_parsed, @@ -474,16 +512,12 @@ async def sample_stream( adapter_path=adapter_path, ) - async def _stream_generator(): - loop = asyncio.get_event_loop() - while True: - item = await loop.run_in_executor(None, q.get) - if item == STREAM_SENTINEL: - break - if isinstance(item, Exception): - yield json.dumps({'error': str(item)}) + '\n' - break - delta, reason = item - yield json.dumps({'delta': delta, 'finish_reason': reason}) + '\n' - - return StreamingResponse(_stream_generator(), media_type='application/x-ndjson') + return StreamingResponse( + _stream_queue( + q, + STREAM_SENTINEL, + request_id, + self._task_queue_config.effective_execution_timeout, + ), + media_type='application/x-ndjson', + ) diff --git a/src/twinkle/server/state/future_manager.py b/src/twinkle/server/state/future_manager.py index 5adf89ccb..3a9bd917d 100644 --- a/src/twinkle/server/state/future_manager.py +++ b/src/twinkle/server/state/future_manager.py @@ -2,12 +2,16 @@ from __future__ import annotations import functools -from datetime import datetime +import time from typing import Any +from twinkle.server.utils.task_errors import task_error_payload +from twinkle.utils.logger import get_logger from .backend.base import StateBackend from .base import BaseManager -from .models import FutureRecord +from .models import FutureRecord, _now_iso + +logger = get_logger() # Status sets used by the do-not-regress guard inside the atomic transform. _TERMINAL_STATUSES = frozenset({'completed', 'failed'}) @@ -17,12 +21,15 @@ def _future_record_transform( existing: dict | None, *, + request_id: str, new_status: str, model_id: str | None, reason: str | None, result: Any, queue_state: str | None, queue_state_reason: str | None, + replica_id: str | None, + absolute_deadline: float | None, now: str, ) -> dict | None: """Atomic transform body for :meth:`FutureManager.store_status`. @@ -30,12 +37,16 @@ def _future_record_transform( Module-level so it remains picklable when forwarded across the Ray actor boundary (closures and lambdas cannot be). - Drops the write entirely (returns ``None``) when ``new_status`` would - regress a terminal status — the StateBackend.update_atomic contract treats - a ``None`` return as "keep the current value", which is what stops stale - retries from clobbering a freshly committed terminal state. + A record already in a terminal state is never overwritten (returns ``None``, + which ``update_atomic`` treats as "keep current value"). A write of a + *different* terminal state is logged; a write of the *same* terminal state is + dropped silently (State_Backend idempotent retries produce these and they + indicate no defect). """ - if (existing is not None and existing.get('status') in _TERMINAL_STATUSES and new_status in _NON_TERMINAL_STATUSES): + existing_status = existing.get('status') if existing is not None else None + if existing_status in _TERMINAL_STATUSES: + if new_status != existing_status: + logger.warning('future %s already terminal as %r; refusing %r', request_id, existing_status, new_status) return None if existing is None: @@ -46,6 +57,8 @@ def _future_record_transform( result=result, queue_state=queue_state, queue_state_reason=queue_state_reason, + replica_id=replica_id, + absolute_deadline=absolute_deadline, created_at=now, updated_at=now, ) @@ -55,6 +68,7 @@ def _future_record_transform( updated['status'] = new_status updated['model_id'] = model_id updated['updated_at'] = now + # replica_id is set at creation and is deliberately NOT overwritten here. if reason is not None: updated['reason'] = reason if result is not None: @@ -67,10 +81,7 @@ def _future_record_transform( class FutureManager(BaseManager[FutureRecord]): - """Manages async task futures / request statuses. - - Expiry is based on `updated_at` (falls back to `created_at`). - """ + """Manage future state, terminal retention, and immutable task deadlines.""" def __init__(self, backend: StateBackend, expiration_timeout: float) -> None: super().__init__(backend, 'future::', FutureRecord, expiration_timeout) @@ -86,6 +97,8 @@ async def store_status( result: Any = None, queue_state: str | None = None, queue_state_reason: str | None = None, + replica_id: str | None = None, + absolute_deadline: float | None = None, ) -> None: """Create or update a future record with the latest status. @@ -99,39 +112,94 @@ async def store_status( if result is not None and hasattr(result, 'model_dump'): result = result.model_dump() - now = datetime.now().isoformat() + now = _now_iso() await self._backend.update_atomic( self._make_key(request_id), functools.partial( _future_record_transform, + request_id=request_id, new_status=status, model_id=model_id, reason=reason, result=result, queue_state=queue_state, queue_state_reason=queue_state_reason, + replica_id=replica_id, + absolute_deadline=absolute_deadline, now=now, ), ) # ----- Cleanup ----- - async def cleanup_expired(self, cutoff_time: float, **kwargs) -> int: - """Remove futures whose last update is older than cutoff_time. + async def cleanup_expired( + self, + cutoff_time: float, + *, + alive_replica_ids: set[str] | None = None, + ) -> int: + """Expire future records without ever deleting a non-terminal one. + + Processing matrix (design §5.2): + + | status | replica alive | past deadline | action | + |--------------|---------------|---------------|-------------------| + | Terminal | — | ts < cutoff | delete | + | non-Terminal | yes | no | keep (untouched) | + | non-Terminal | yes | yes | write ``failed`` | + | non-Terminal | no | — | write ``failed`` | Args: - cutoff_time: Unix timestamp threshold. + cutoff_time: Unix timestamp; terminal records older than it are deleted. + alive_replica_ids: replicas currently considered alive. ``None`` disables + the orphan check (every non-terminal record is treated as owned). Returns: - Number of futures removed. + Number of terminal records removed (records written ``failed`` are not + counted here; they are removed on a later pass once terminal). """ all_records = await self.get_all() - expired_ids = [] + now = time.time() + expired_ids: list[str] = [] for request_id, record in all_records.items(): - timestamp_str = record.updated_at or record.created_at - timestamp = self._parse_timestamp(timestamp_str) - if timestamp < cutoff_time: - expired_ids.append(request_id) + if record.status in _TERMINAL_STATUSES: + timestamp = self._parse_timestamp(record.updated_at or record.created_at) + if timestamp < cutoff_time: + expired_ids.append(request_id) + continue + + # Non-terminal records are never deleted -- only ever written ``failed``. + # replica_id None (pre-upgrade) => ownership unknown => treated as alive. + replica_id = record.replica_id + replica_alive = (replica_id is None or alive_replica_ids is None or replica_id in alive_replica_ids) + if not replica_alive: + await self.store_status( + request_id, + 'failed', + record.model_id, + result=task_error_payload( + 'The replica that owned this task is no longer available.', + request_id=request_id, + error_code=503, + ), + replica_id=replica_id, + ) + continue + deadline = record.absolute_deadline + if deadline is None: + deadline = self._parse_timestamp(record.created_at) + self.expiration_timeout + if now > deadline: + await self.store_status( + request_id, + 'failed', + record.model_id, + result=task_error_payload( + 'Task exceeded the absolute survival bound without reaching a terminal state.', + request_id=request_id, + error_code=500, + ), + replica_id=replica_id, + ) for request_id in expired_ids: await self.remove(request_id) diff --git a/src/twinkle/server/state/model_manager.py b/src/twinkle/server/state/model_manager.py index ae75621c4..58edf59b2 100644 --- a/src/twinkle/server/state/model_manager.py +++ b/src/twinkle/server/state/model_manager.py @@ -11,6 +11,7 @@ from __future__ import annotations import functools +import time from .backend.base import StateBackend from .base import BaseManager @@ -101,6 +102,28 @@ async def unregister_replica(self, replica_id: str) -> None: await self.remove(model_id) await self._replicas.unregister(replica_id) + async def touch_replica_last_seen(self, replica_id: str) -> None: + """Refresh a replica's liveness timestamp (R4#6).""" + await self._replicas.touch_last_seen(replica_id) + + async def get_alive_replica_ids(self, liveness_threshold: float) -> set[str]: + """Return replicas considered alive (R4#7, R4#8). + + A replica is alive when it has a ``last_seen`` within ``liveness_threshold``, + OR when it has a ``max_loras`` entry but no ``last_seen`` yet (registered + before this spec / before its first request -- treated as alive so an + upgrade does not orphan in-flight tasks). + """ + registered = await self._replicas.get_all() + last_seen = await self._replicas.get_all_last_seen() + now = time.time() + alive: set[str] = set() + for rid in set(registered) | set(last_seen): + ls = last_seen.get(rid) + if (ls is None and rid in registered) or (ls is not None and (now - ls) <= liveness_threshold): + alive.add(rid) + return alive + async def get_available_replica_ids(self, candidate_ids: list[str]) -> list[str]: """Return the subset of ``candidate_ids`` that still have capacity. diff --git a/src/twinkle/server/state/models.py b/src/twinkle/server/state/models.py index 71279b894..7b11813ae 100644 --- a/src/twinkle/server/state/models.py +++ b/src/twinkle/server/state/models.py @@ -2,13 +2,16 @@ from __future__ import annotations import time -from datetime import datetime +from datetime import datetime, timezone from pydantic import BaseModel, Field from typing import Any def _now_iso() -> str: - return datetime.now().isoformat() + # UTC-aware so _parse_timestamp (which reads timestamps back as UTC) agrees with + # it and with time.time(); a naive local string would be misread as UTC and skew + # every expiry comparison by the host's UTC offset. + return datetime.now(timezone.utc).isoformat() class SessionRecord(BaseModel): @@ -53,5 +56,8 @@ class FutureRecord(BaseModel): result: Any = None queue_state: str | None = None queue_state_reason: str | None = None + # Replica ownership and deadline are fixed when the record is created. + replica_id: str | None = None + absolute_deadline: float | None = None created_at: str = Field(default_factory=_now_iso) updated_at: str = Field(default_factory=_now_iso) diff --git a/src/twinkle/server/state/replica_registry.py b/src/twinkle/server/state/replica_registry.py index b2e11d13b..d3eda8d85 100644 --- a/src/twinkle/server/state/replica_registry.py +++ b/src/twinkle/server/state/replica_registry.py @@ -1,28 +1,28 @@ # Copyright (c) ModelScope Contributors. All rights reserved. -"""Backend-backed registry of replica capacity. +"""Backend-backed registry of replica capacity and liveness. -Each entry persists to ``replica::::max_loras`` in the configured -:class:`StateBackend` (Redis or the actor-wrapped RayActorBackend), so every -Ray Serve worker sees one consistent view of the cluster's capacity even -though each worker holds its own ``ServerState`` instance. - -The registry knows *only* about declared capacity. The current loaded-model -count is derived by querying the persisted ``model::*`` records directly — -nothing here caches that count, so concurrent writes from different workers -cannot drift into an inconsistent local index. +Capacity and ``last_seen`` use separate keys so sampler liveness does not alter +the model-capacity data shape. """ from __future__ import annotations +import time + from .backend.base import StateBackend REPLICA_PREFIX = 'replica::' _MAX_LORAS_SUFFIX = '::max_loras' +_LAST_SEEN_SUFFIX = '::last_seen' def _make_key(replica_id: str) -> str: return f'{REPLICA_PREFIX}{replica_id}{_MAX_LORAS_SUFFIX}' +def _last_seen_key(replica_id: str) -> str: + return f'{REPLICA_PREFIX}{replica_id}{_LAST_SEEN_SUFFIX}' + + def _replica_id_from_key(key: str) -> str | None: if not key.startswith(REPLICA_PREFIX) or not key.endswith(_MAX_LORAS_SUFFIX): return None @@ -30,7 +30,7 @@ def _replica_id_from_key(key: str) -> str | None: class ReplicaRegistry: - """Read/write replica capacity through the shared :class:`StateBackend`.""" + """Read/write replica capacity and liveness through the shared backend.""" def __init__(self, backend: StateBackend) -> None: self._backend = backend @@ -42,6 +42,36 @@ async def register(self, replica_id: str, max_loras: int) -> None: async def unregister(self, replica_id: str) -> None: """Remove the capacity entry for ``replica_id`` (idempotent).""" await self._backend.delete(_make_key(replica_id)) + await self._backend.delete(_last_seen_key(replica_id)) + + async def touch_last_seen(self, replica_id: str) -> None: + """Refresh the replica's liveness timestamp (separate key from max_loras).""" + await self._backend.set(_last_seen_key(replica_id), time.time()) + + async def get_last_seen(self, replica_id: str) -> float | None: + """Return the replica's last-seen unix time, or ``None`` if never refreshed.""" + value = await self._backend.get(_last_seen_key(replica_id)) + if value is None: + return None + try: + return float(value) + except (TypeError, ValueError): + return None + + async def get_all_last_seen(self) -> dict[str, float]: + """Return every replica's last-seen timestamp.""" + keys = await self._backend.keys(f'{REPLICA_PREFIX}*{_LAST_SEEN_SUFFIX}') + out: dict[str, float] = {} + for key in keys: + if not key.startswith(REPLICA_PREFIX) or not key.endswith(_LAST_SEEN_SUFFIX): + continue + rid = key[len(REPLICA_PREFIX):-len(_LAST_SEEN_SUFFIX)] + value = await self._backend.get(key) + try: + out[rid] = float(value) + except (TypeError, ValueError): + continue + return out async def get_max_loras(self, replica_id: str) -> int | None: """Return the declared capacity, or ``None`` if the replica is unknown.""" diff --git a/src/twinkle/server/state/server_state.py b/src/twinkle/server/state/server_state.py index 1fb548f54..fc8014a03 100644 --- a/src/twinkle/server/state/server_state.py +++ b/src/twinkle/server/state/server_state.py @@ -289,6 +289,8 @@ async def store_future_status( result: Any = None, queue_state: str | None = None, queue_state_reason: str | None = None, + replica_id: str | None = None, + absolute_deadline: float | None = None, ) -> None: """Store task status with optional result. @@ -317,6 +319,8 @@ async def store_future_status( result=result, queue_state=queue_state, queue_state_reason=queue_state_reason, + replica_id=replica_id, + absolute_deadline=absolute_deadline, ) # ----- Configuration Management ----- @@ -370,7 +374,9 @@ async def cleanup_expired_resources(self) -> dict[str, int]: models_removed = await self._model_mgr.cleanup_expired(cutoff_time, expired_session_ids=expired_session_ids) samplings_removed = await self._sampling_mgr.cleanup_expired( cutoff_time, expired_session_ids=expired_session_ids) - futures_removed = await self._future_mgr.cleanup_expired(cutoff_time) + + alive_replica_ids = await self._model_mgr.get_alive_replica_ids(self.expiration_timeout) + futures_removed = await self._future_mgr.cleanup_expired(cutoff_time, alive_replica_ids=alive_replica_ids) return { 'sessions': sessions_removed, @@ -379,6 +385,10 @@ async def cleanup_expired_resources(self) -> dict[str, int]: 'futures': futures_removed, } + async def touch_replica_last_seen(self, replica_id: str) -> None: + """Refresh a replica's liveness timestamp in the shared registry (R4#6).""" + await self._model_mgr.touch_replica_last_seen(replica_id) + async def _cleanup_loop(self) -> None: """Background task that periodically cleans up expired resources. diff --git a/src/twinkle/server/utils/task_errors.py b/src/twinkle/server/utils/task_errors.py index 478745c03..031f71cf0 100644 --- a/src/twinkle/server/utils/task_errors.py +++ b/src/twinkle/server/utils/task_errors.py @@ -1,5 +1,73 @@ # Copyright (c) ModelScope Contributors. All rights reserved. +"""Construction and backward-compatible reading of failure payloads. +``ErrorPayload`` is the single representation of a failure both on the wire and in +state (R5). This module owns the two entry points that produce/repair it. +""" +from __future__ import annotations -def task_error_payload(error: str) -> dict[str, str]: - return {'error': error, 'category': 'Server'} +from collections.abc import Mapping +from typing import Any + +from twinkle_client.types.errors import ErrorCategory, ErrorPayload + +_ERROR_MAX = 1024 +_TRACEBACK_MAX = 65536 +_TRUNCATION_MARKER = '...[traceback truncated, tail kept]...\n' + + +def _trim_traceback(text: str) -> str: + """Keep the tail of an over-long traceback (innermost frames are densest).""" + if len(text) <= _TRACEBACK_MAX: + return text + keep = _TRACEBACK_MAX - len(_TRUNCATION_MARKER) + return _TRUNCATION_MARKER + text[-keep:] + + +def task_error_payload( + error: str, + *, + request_id: str, + error_code: int = 500, + category: ErrorCategory | str = ErrorCategory.Server, + traceback_text: str | None = None, +) -> dict[str, Any]: + """Build an ``ErrorPayload`` and return it as a JSON-safe dict for storage. + + Traceback splitting and length trimming happen here so over-long text is never + written to State_Backend. A ``user`` category carries no traceback. + """ + if isinstance(category, str): + category = ErrorCategory(category.lower()) + tb = _trim_traceback(traceback_text) if category is ErrorCategory.Server and traceback_text else None + lines = str(error).splitlines() + summary = (lines[0] if lines else '')[:_ERROR_MAX] + payload = ErrorPayload( + error=summary, + category=category, + error_code=error_code, + request_id=request_id, + traceback=tb, + ) + return payload.model_dump(mode='json', exclude_none=True) + + +def error_payload_from_stored(stored: Any, *, request_id: str) -> ErrorPayload: + """Build an ``ErrorPayload`` from whatever is sitting in ``FutureRecord.result``. + + Records written before this spec have only ``{error, category}``. Missing + ``error_code`` / ``request_id`` / ``category`` are backfilled with ``500`` / + the caller-supplied value / ``Unknown`` so a rolling upgrade never raises + ``pydantic.ValidationError``. + """ + if isinstance(stored, Mapping): + data = dict(stored) + else: + data = {'error': 'Unknown error' if stored is None else str(stored)} + data.setdefault('category', ErrorCategory.Unknown) + data.setdefault('error_code', 500) + data.setdefault('request_id', request_id) + category = str(data['category']).lower() + if category != ErrorCategory.Server.value: + data.pop('traceback', None) + return ErrorPayload.model_validate(data) diff --git a/src/twinkle/server/utils/task_queue/__init__.py b/src/twinkle/server/utils/task_queue/__init__.py index 5c90d3181..d77dfe182 100644 --- a/src/twinkle/server/utils/task_queue/__init__.py +++ b/src/twinkle/server/utils/task_queue/__init__.py @@ -12,12 +12,13 @@ from .config import TaskQueueConfig from .mixin import TaskQueueMixin from .rate_limiter import RateLimiter -from .types import QueuedTask, QueueState, TaskStatus +from .types import QueuedTask, QueueState, TaskStatus, UserTaskError from .worker import ComputeWorker __all__ = [ 'TaskStatus', 'QueueState', + 'UserTaskError', 'QueuedTask', 'TaskQueueConfig', 'TaskQueueMixin', diff --git a/src/twinkle/server/utils/task_queue/config.py b/src/twinkle/server/utils/task_queue/config.py index 79d62095a..4de98e0da 100644 --- a/src/twinkle/server/utils/task_queue/config.py +++ b/src/twinkle/server/utils/task_queue/config.py @@ -11,6 +11,11 @@ from pydantic import BaseModel, ConfigDict, Field +# Finite bounds used when configuration omits a limit and by long-running methods. +_ZERO_EXECUTION_TIMEOUT_FALLBACK: float = 3600.0 +_MAX_DECLARED_BACKEND_TIMEOUT: float = 3600.0 +_ABSOLUTE_TTL_MULTIPLIER: int = 2 + class TaskQueueConfig(BaseModel): """Configuration for task queue and rate limiting. @@ -20,7 +25,9 @@ class TaskQueueConfig(BaseModel): tps_limit: Maximum input tokens per second per user token. ``0`` disables. window_seconds: Sliding window for rate-limit calculations. Must be > 0. queue_timeout: Maximum time a task can wait in queue (seconds). - execution_timeout: Maximum time a task can execute (seconds). 0 means no limit. + execution_timeout: Maximum time a task can execute (seconds). ``0`` means "no + configured limit"; a finite bound of 3600s is substituted instead of + unbounded waiting (see ``effective_execution_timeout``). enabled: Whether rate limiting is enabled. token_cleanup_multiplier: Multiplier for token cleanup threshold. token_cleanup_interval: How often to run cleanup task (seconds). @@ -33,8 +40,28 @@ class TaskQueueConfig(BaseModel): tps_limit: float = Field(default=16000.0, ge=0) window_seconds: float = Field(default=1.0, gt=0) queue_timeout: float = Field(default=300.0, ge=0) - execution_timeout: float = Field(default=120.0, ge=0) + execution_timeout: float = Field(default=1800.0, ge=0) enabled: bool = True token_cleanup_multiplier: float = Field(default=10.0, ge=0) token_cleanup_interval: float = Field(default=60.0, ge=0) max_input_tokens: int = Field(default=16000, ge=1) + + @property + def effective_execution_timeout(self) -> float: + """The single source of the execution time bound. + + ``0`` is not rejected (that would fail existing deployments); it is read + as "no configured limit" and replaced by a finite fallback so the bound + is always positive. This value feeds both ``_ray_get_timeout`` and the + ComputeWorker's ``asyncio.wait_for`` -- there is no second, independently + configurable timeout. + """ + if self.execution_timeout > 0: + return self.execution_timeout + return _ZERO_EXECUTION_TIMEOUT_FALLBACK + + def absolute_future_ttl(self, collect_width: int) -> float: + """Conservative lifetime for a non-terminal future record.""" + ray_timeout = max(self.effective_execution_timeout, _MAX_DECLARED_BACKEND_TIMEOUT) + resource_bound = max(1, collect_width) * ray_timeout + return _ABSOLUTE_TTL_MULTIPLIER * (self.queue_timeout + resource_bound) diff --git a/src/twinkle/server/utils/task_queue/mixin.py b/src/twinkle/server/utils/task_queue/mixin.py index dcdf805ce..83526c399 100644 --- a/src/twinkle/server/utils/task_queue/mixin.py +++ b/src/twinkle/server/utils/task_queue/mixin.py @@ -8,18 +8,22 @@ from __future__ import annotations import asyncio +import contextlib +import functools import time import traceback import uuid from collections.abc import Callable, Coroutine +from concurrent.futures import ThreadPoolExecutor from typing import TYPE_CHECKING, Any from twinkle.server.telemetry.middleware import get_task_metrics from twinkle.server.utils.task_errors import task_error_payload from twinkle.utils.logger import get_logger +from twinkle_client.types.errors import ErrorCategory from .config import TaskQueueConfig from .rate_limiter import RateLimiter -from .types import QueuedTask, QueueState, TaskStatus +from .types import BackendBusyError, QueuedTask, QueueState, TaskStatus from .worker import ComputeWorker if TYPE_CHECKING: @@ -50,16 +54,38 @@ class TaskQueueMixin: state: ServerState - def _init_task_queue(self, config: TaskQueueConfig | None = None, deployment_name: str = '') -> None: + def _init_task_queue( + self, + config: TaskQueueConfig | None = None, + deployment_name: str = '', + *, + enable_admission_gate: bool = False, + on_backend_timeout: Callable[[], Coroutine[Any, Any, None]] | None = None, + collect_width: int = 1, + ) -> None: """Initialise the task queue, rate limiter, and compute worker. ``config`` must be a typed :class:`TaskQueueConfig` (the launcher passes the instance straight through). ``None`` constructs a default config. + + ``enable_admission_gate`` turns on the per-replica Admission_Gate + (:meth:`call_backend`). ``ModelManagement`` enables it; ``SamplerManagement`` + does not (vllm sampler owns its own concurrency and the weight-update / + generation mutual exclusion is covered by infra ``_cw_barrier``). + + ``on_backend_timeout`` runs after a backend timeout. ``collect_width`` is the + number of actor results a backend call may collect and determines the persisted + future deadline. """ self._task_queue_config = config if config is not None else TaskQueueConfig() + if self._task_queue_config.execution_timeout == 0: + logger.warning( + '[TaskQueue] execution_timeout=0: a finite %.0fs bound has replaced unbounded waiting ' + '(deployment=%s).', self._task_queue_config.effective_execution_timeout, deployment_name or 'unknown') self._deployment_name = deployment_name self._task_metrics = get_task_metrics(deployment_name) if deployment_name else None + self._future_absolute_ttl = self._task_queue_config.absolute_future_ttl(collect_width) self._rate_limiter = RateLimiter( rps_limit=self._task_queue_config.rps_limit, @@ -77,10 +103,95 @@ def _init_task_queue(self, config: TaskQueueConfig | None = None, deployment_nam config=self._task_queue_config, task_metrics=self._task_metrics, deployment_name=deployment_name, + on_backend_timeout=on_backend_timeout, ) + self._backend_executor = ThreadPoolExecutor(thread_name_prefix='twinkle-backend') + self._backend_probe_executor = ThreadPoolExecutor(max_workers=1, thread_name_prefix='twinkle-backend-probe') + self._backend_admission: asyncio.Lock | None = asyncio.Lock() if enable_admission_gate else None + self._backend_poisoned = asyncio.Event() self._event_loop: asyncio.AbstractEventLoop | None = None + async def _acquire_backend_gate(self, gate: asyncio.Lock) -> None: + if self._backend_poisoned.is_set(): + raise BackendBusyError('This replica is waiting for a timed-out backend call to exit.') + if not gate.locked(): + await gate.acquire() + else: + acquire_task = asyncio.create_task(gate.acquire()) + poison_task = asyncio.create_task(self._backend_poisoned.wait()) + try: + done, _ = await asyncio.wait((acquire_task, poison_task), return_when=asyncio.FIRST_COMPLETED) + except asyncio.CancelledError: + acquire_task.cancel() + poison_task.cancel() + await asyncio.gather(acquire_task, poison_task, return_exceptions=True) + if acquire_task.done() and not acquire_task.cancelled() and acquire_task.result(): + gate.release() + raise + if poison_task in done and self._backend_poisoned.is_set(): + if acquire_task.done() and not acquire_task.cancelled() and acquire_task.result(): + gate.release() + else: + acquire_task.cancel() + with contextlib.suppress(asyncio.CancelledError): + await acquire_task + raise BackendBusyError('This replica is waiting for a timed-out backend call to exit.') + poison_task.cancel() + with contextlib.suppress(asyncio.CancelledError): + await poison_task + await acquire_task + if self._backend_poisoned.is_set(): + gate.release() + raise BackendBusyError('This replica is waiting for a timed-out backend call to exit.') + + async def call_backend(self, fn: Callable[..., Any], /, *args: Any, admit: bool = True, **kwargs: Any) -> Any: + """Run one backend call outside the event loop. + + Normal model calls serialize through the admission gate. If the awaiting + task times out while its thread is still running, the gate is poisoned: + waiters fail immediately until that thread exits. Health probes bypass the + gate and use a reserved executor thread. Sampler deployments disable the + gate because their backend owns request concurrency. + """ + loop = asyncio.get_running_loop() + gate = self._backend_admission if admit else None + if gate is not None: + await self._acquire_backend_gate(gate) + + executor = self._backend_executor if admit else self._backend_probe_executor + try: + concurrent_future = executor.submit(functools.partial(fn, *args, **kwargs)) + except Exception: + if gate is not None and gate.locked(): + gate.release() + raise + + if gate is not None: + + def release_gate(_future) -> None: + + def release() -> None: + self._backend_poisoned.clear() + if gate.locked(): + gate.release() + + with contextlib.suppress(RuntimeError): + loop.call_soon_threadsafe(release) + + concurrent_future.add_done_callback(release_gate) + + try: + return await asyncio.wrap_future(concurrent_future, loop=loop) + except asyncio.CancelledError: + if gate is not None and concurrent_future.running(): + self._backend_poisoned.set() + raise + + def _future_deadline(self) -> float: + ttl = getattr(self, '_future_absolute_ttl', self._task_queue_config.absolute_future_ttl(1)) + return time.time() + ttl + @staticmethod def _queue_key(model_id: str | None, token: str | None) -> str: if model_id: @@ -108,7 +219,13 @@ async def _perform_preflight_checks( return None async def reject(error_msg: str, queue_state: str) -> dict[str, Any]: - error_payload = {'error': error_msg, 'category': 'User'} + error_code = 429 if queue_state == QueueState.PAUSED_RATE_LIMIT.value else 400 + error_payload = task_error_payload( + error_msg, + request_id=request_id, + error_code=error_code, + category=ErrorCategory.User, + ) if persist_failure: await self.state.store_future_status( request_id, @@ -117,6 +234,7 @@ async def reject(error_msg: str, queue_state: str) -> dict[str, Any]: result=error_payload, queue_state=queue_state, queue_state_reason=error_msg, + replica_id=getattr(self, 'replica_id', None), ) return {'request_id': request_id, 'model_id': model_id} # Private marker consumed by schedule_task_and_wait(). It is not @@ -192,6 +310,8 @@ async def _schedule_task( TaskStatus.PENDING.value, model_id, queue_state=QueueState.ACTIVE.value, + replica_id=getattr(self, 'replica_id', None), + absolute_deadline=self._future_deadline(), ) queue_key = self._queue_key(model_id=model_id, token=token) @@ -342,6 +462,8 @@ async def schedule_background_task( TaskStatus.RUNNING.value, model_id, queue_state=QueueState.ACTIVE.value, + replica_id=getattr(self, 'replica_id', None), + absolute_deadline=self._future_deadline(), ) async def _run() -> None: @@ -355,8 +477,13 @@ async def _run() -> None: queue_state=QueueState.ACTIVE.value, ) logger.info(f'[TaskQueue] Background task {request_id} completed, type={task_type or "unknown"}') - except Exception: - error_payload = task_error_payload(traceback.format_exc()) + except Exception as exc: + error_payload = task_error_payload( + f'{type(exc).__name__}: {exc}', + request_id=request_id, + error_code=500, + traceback_text=traceback.format_exc(), + ) await self.state.store_future_status( request_id, TaskStatus.FAILED.value, @@ -413,4 +540,9 @@ async def shutdown_task_queue(self) -> None: """Gracefully shut down the compute queue and release resources.""" await self._rate_limiter.stop_cleanup_task() await self._compute_worker.stop() + # Do not wait on threads that may be leaked on a timed-out backend call. + if getattr(self, '_backend_executor', None) is not None: + self._backend_executor.shutdown(wait=False, cancel_futures=True) + if getattr(self, '_backend_probe_executor', None) is not None: + self._backend_probe_executor.shutdown(wait=False, cancel_futures=True) logger.debug('[TaskQueue] Task queue shutdown complete') diff --git a/src/twinkle/server/utils/task_queue/types.py b/src/twinkle/server/utils/task_queue/types.py index daf8d2bb2..c7f51462b 100644 --- a/src/twinkle/server/utils/task_queue/types.py +++ b/src/twinkle/server/utils/task_queue/types.py @@ -26,6 +26,20 @@ class TaskStatus(Enum): RATE_LIMITED = 'rate_limited' # Task rejected due to rate limiting +class UserTaskError(ValueError): + """A queued operation rejected because of caller input or usage.""" + + +class BackendBusyError(RuntimeError): + """Raised when the per-replica Admission_Gate is held by a leaked backend call. + + A new backend call arriving while the gate is closed (its holder is a call that + already exceeded ``asyncio.wait_for`` but whose executor thread has not yet + returned) fails fast with this error instead of queueing behind it. The worker + maps it to ``ErrorPayload(category='server', error_code=503)``. + """ + + class QueueState(Enum): """Queue state for tinker client compatibility. diff --git a/src/twinkle/server/utils/task_queue/worker.py b/src/twinkle/server/utils/task_queue/worker.py index fdbb36d16..b1e276cf6 100644 --- a/src/twinkle/server/utils/task_queue/worker.py +++ b/src/twinkle/server/utils/task_queue/worker.py @@ -12,14 +12,15 @@ import time import traceback from collections import deque -from typing import TYPE_CHECKING, Any, Deque +from typing import TYPE_CHECKING, Any, Callable, Deque from twinkle.server.telemetry.correlation import MODEL_ID, TOKEN_ID from twinkle.server.telemetry.tracing import traced_operation from twinkle.server.utils.task_errors import task_error_payload from twinkle.utils.logger import get_logger +from twinkle_client.types.errors import ErrorCategory from .config import TaskQueueConfig -from .types import QueuedTask, QueueState, TaskStatus +from .types import BackendBusyError, QueuedTask, QueueState, TaskStatus, UserTaskError if TYPE_CHECKING: from twinkle.server.state import ServerState @@ -27,6 +28,13 @@ logger = get_logger() +# Ray_Get_Timeout is classified the same as asyncio.TimeoutError: 504/Server (R5#8). +try: + from ray.exceptions import GetTimeoutError as _RayGetTimeout + _TIMEOUT_EXCEPTIONS: tuple[type[BaseException], ...] = (asyncio.TimeoutError, _RayGetTimeout) +except Exception: # pragma: no cover - ray always present in server runtime + _TIMEOUT_EXCEPTIONS = (asyncio.TimeoutError, ) + class ComputeWorker: """Serial background worker that processes GPU compute tasks. @@ -46,11 +54,14 @@ def __init__( config: TaskQueueConfig, task_metrics: TaskMetrics | None, deployment_name: str, + on_backend_timeout: Callable[[], Any] | None = None, ) -> None: self._state = state self._config = config self._task_metrics = task_metrics self._deployment_name = deployment_name + # Optional coroutine-returning callback fired on a backend timeout (R3#2). + self._on_backend_timeout = on_backend_timeout self.task_queues: dict[str, asyncio.Queue] = {} self.queue_order: Deque[str] = deque() @@ -141,14 +152,24 @@ async def _store_task_failed( error: str, queue_state: str, queue_state_reason: str | None = None, + *, + error_code: int = 500, + category: ErrorCategory = ErrorCategory.Server, + traceback_text: str | None = None, ) -> None: - """Store FAILED status with a standardised error payload.""" + """Store FAILED status with a standardised ``ErrorPayload``.""" if task.persist_status: await self._state.store_future_status( task.request_id, TaskStatus.FAILED.value, task.model_id, - result=task_error_payload(error), + result=task_error_payload( + error, + request_id=task.request_id, + error_code=error_code, + category=category, + traceback_text=traceback_text, + ), queue_state=queue_state, queue_state_reason=queue_state_reason, ) @@ -232,10 +253,9 @@ async def _execute_task(self, task: QueuedTask, queue_key: str, q: asyncio.Queue f'type={task_type}, queue_key={queue_key}') with traced_operation(handler_span_name, attrs=handler_attrs): coro = task.coro_factory() - if self._config.execution_timeout > 0: - result = await asyncio.wait_for(coro, timeout=self._config.execution_timeout) - else: - result = await coro + # effective_execution_timeout is always positive (0 -> finite fallback), + # so wait_for is always in effect. + result = await asyncio.wait_for(coro, timeout=self._config.effective_execution_timeout) exec_time = time.monotonic() - exec_start logger.info(f'[ComputeWorker] Task {task.request_id} completed in {exec_time:.2f}s, type={task_type}') if task.persist_status: @@ -247,21 +267,50 @@ async def _execute_task(self, task: QueuedTask, queue_key: str, q: asyncio.Queue queue_state=QueueState.ACTIVE.value, ) self._complete_result(task, result) - except asyncio.TimeoutError: + except _TIMEOUT_EXCEPTIONS: task_status = 'timeout' exec_time = time.monotonic() - exec_start - error = (f'Execution timeout exceeded: {self._config.execution_timeout}s, ' + error = (f'Backend call timed out (bound {self._config.effective_execution_timeout}s), ' f'actual execution time: {exec_time:.2f}s') logger.error(f'[ComputeWorker] Task {task.request_id} TIMEOUT after {exec_time:.2f}s, ' f'type={task_type}, queue_key={queue_key}') - await self._store_task_failed(task, error, QueueState.ACTIVE.value) - except Exception: + # asyncio.TimeoutError and Ray_Get_Timeout are 504/Server (R5#8). + await self._store_task_failed(task, error, QueueState.ACTIVE.value, error_code=504) + # Probe actor liveness after a timeout so an operator learns the replica's + # state without waiting for a second request to also time out (R3#2). + if self._on_backend_timeout is not None: + try: + await self._on_backend_timeout() + except Exception: + logger.error(f'[ComputeWorker] backend-timeout probe failed:\n{traceback.format_exc(limit=3)}') + except UserTaskError as exc: + task_status = 'failed' + exec_time = time.monotonic() - exec_start + await self._store_task_failed( + task, + f'{type(exc).__name__}: {exc}', + QueueState.UNKNOWN.value, + error_code=400, + category=ErrorCategory.User, + ) + except BackendBusyError as exc: + task_status = 'failed' + exec_time = time.monotonic() - exec_start + error = str(exc) + logger.error(f'[ComputeWorker] Task {task.request_id} REFUSED (admission gate held) after ' + f'{exec_time:.2f}s, type={task_type}, queue_key={queue_key}') + # Gate held by a leaked timed-out call -> 503/Server (R2#4). + await self._store_task_failed(task, error, QueueState.ACTIVE.value, error_code=503) + except Exception as exc: task_status = 'failed' exec_time = time.monotonic() - exec_start - error = traceback.format_exc() + # error is a single-line summary; the full traceback goes only to the + # traceback field, never into `error` (R5#7). + error = f'{type(exc).__name__}: {exc}' logger.error(f'[ComputeWorker] Task {task.request_id} FAILED after {exec_time:.2f}s, ' f'type={task_type}:\n{traceback.format_exc(limit=3)}') - await self._store_task_failed(task, error, QueueState.ACTIVE.value) + await self._store_task_failed( + task, error, QueueState.ACTIVE.value, error_code=500, traceback_text=traceback.format_exc()) finally: q.task_done() self._record_execution_time(task_type, exec_time) @@ -296,12 +345,29 @@ async def _try_run_one(self) -> bool: await self._fail_timed_out_task(task, queue_wait, q) continue # try the next queue + # A record already in a Terminal_State (e.g. written 'failed' by the + # state-hygiene orphan handling) must not be executed again (R3#8). + if task.persist_status and await self._is_record_terminal(task.request_id): + logger.info(f'[ComputeWorker] Task {task.request_id} already terminal on dequeue; skipping.') + q.task_done() + continue + # Execute the task (serial: stops after the first execution) await self._execute_task(task, queue_key, q) return True return False + async def _is_record_terminal(self, request_id: str) -> bool: + """True if the future record already holds a Terminal_State.""" + try: + record = await self._state.get_future(request_id) + except Exception: + return False + if not record: + return False + return record.get('status') in (TaskStatus.COMPLETED.value, TaskStatus.FAILED.value) + # ------------------------------------------------------------------ # Main worker loop # ------------------------------------------------------------------ diff --git a/src/twinkle/utils/nccl_safe.py b/src/twinkle/utils/nccl_safe.py index 8bafeb1f2..a7c7eef0c 100644 --- a/src/twinkle/utils/nccl_safe.py +++ b/src/twinkle/utils/nccl_safe.py @@ -1,339 +1,73 @@ # Copyright (c) ModelScope Contributors. All rights reserved. -"""NCCL-safe utilities for production distributed training. - -Provides three layers of protection to prevent NCCL hangs: - -Layer 1 - safe_loss(): - Wraps loss instances to catch computation errors and return - graph-connected zero loss (ensures FSDP ReduceScatter can proceed). - -Layer 2 - @nccl_safe decorator: - Wraps forward_backward methods to ensure backward() always executes - after forward() has started, even if intermediate code raises. - -Layer 3 - @nccl_safe_megatron decorator: - Wraps Megatron backend methods (forward_only, forward_backward) where - the entire function body involves NCCL communication (sync=True). - Catches pre-communication errors (e.g. data preprocessing failures) - that would otherwise leave other DP ranks waiting at a collective. - -Controlled by environment variable: - TWINKLE_FAIL_FAST=1 (default, development): all protection is transparent, - exceptions propagate normally. - TWINKLE_FAIL_FAST=0 (production): protection activated, exceptions in - NCCL-critical sections are caught and handled gracefully. +"""NCCL critical-section failure logging. + +Single responsibility: inside a Megatron NCCL-critical method, log a +rank-attributed failure and then re-raise it unchanged. + +This module does NOT prevent asymmetric-failure blocking -- nothing at this layer +can. A rank that swallows its exception still does not enter the collective, so the +other ranks stay blocked regardless. The time bound for an asymmetric failure comes +from Ray_Get_Timeout (the effective execution timeout applied per future), not from +this decorator. Diagnosability is the only reason this wrapper exists. + +Coverage removed together with the former Layer 1 (the loss-instance wrapper) and +Layer 2 (the forward/backward decorator) silent degradation: under FSDP, the window +between ``calculate_loss``'s loss call and its surrounding bookkeeping (metric +accumulation, ``status.num_tokens``), between the three calls inside a +``forward_backward`` body, and numerical problems inside a loss (NaN, shape mismatch) +may each constitute a "forward ran, backward did not" asymmetric-failure window. That +window is no longer covered by any silent degradation; its time bound is the two +bounds documented for the task queue (record-terminal = ``queue_timeout + T``; +resource-release = ``Collect_Width * T``). """ import functools -import os -from twinkle.data_format import LossOutput -from twinkle.loss import Loss from twinkle.utils.logger import get_logger logger = get_logger() +# Errors are logged with at most this many trailing characters of traceback. +_TRACEBACK_LIMIT = 8192 -def _is_fail_fast() -> bool: - """Check if fail-fast mode is enabled (default: enabled). - - Returns True (fail-fast/development mode) unless TWINKLE_FAIL_FAST - is explicitly set to a falsy value. - """ - val = os.getenv('TWINKLE_FAIL_FAST', '1').upper() - return val not in ('0', 'NO', 'FALSE', 'OFF') - - -# ─── Layer 1: safe_loss ──────────────────────────────────────────────────── +def _global_rank() -> int: + """Best-effort global rank for failure attribution; -1 if unavailable.""" + try: + from twinkle.utils import Platform + return Platform.get_rank() + except Exception: + return -1 -def safe_loss(loss_instance): - """Wrap loss instance for production graceful degradation. - Always wraps the loss instance (idempotent). The fail-fast check is deferred - to call time so that TWINKLE_FAIL_FAST can be set after wrapping (e.g. in - Ray actor processes where env vars may not be inherited from the launcher). +def nccl_safe_megatron(func): + """Log a rank-attributed failure inside the NCCL critical section, then re-raise. - When TWINKLE_FAIL_FAST=1 (default, development): wrapper is transparent, - exceptions propagate normally. - When TWINKLE_FAIL_FAST=0 (production): wrapper catches exceptions and - returns a graph-connected zero loss (ensures FSDP ReduceScatter proceeds). - - Idempotent: already-wrapped instances are returned as-is. - """ - if getattr(loss_instance, '_nccl_safe_wrapped', False): - return loss_instance - return SafeLossWrapper(loss_instance) - - -class SafeLossWrapper(Loss): - """Loss subclass that catches computation errors and returns graph-connected zero loss. - - Inherits from :class:`twinkle.loss.Loss` so ``isinstance(wrapper, Loss)`` - assertions in the training pipeline continue to pass. + This decorator does *not* prevent asymmetric-failure blocking -- nothing at this + layer can. A rank that swallows its exception still does not enter the collective. + The time bound for that case comes from Ray_Get_Timeout. Diagnosability is the only + reason this wrapper still exists. Its behavior is unconditional: no environment + variable or config switch affects it, and it returns no degraded value. """ - def __init__(self, loss_instance): - super().__init__() - self._loss_instance = loss_instance - self.require_logps = getattr(loss_instance, 'require_logps', True) - self.require_entropy = getattr(loss_instance, 'require_entropy', False) - self.require_logits = getattr(loss_instance, 'require_logits', False) - self.enable_sampling_replay = getattr(loss_instance, 'enable_sampling_replay', False) - self.require_values = getattr(loss_instance, 'require_values', False) - self.reduction = getattr(loss_instance, 'reduction', 'mean') - self._nccl_safe_wrapped = True - - def __call__(self, inputs, outputs, **kwargs): - if _is_fail_fast(): - return self._loss_instance(inputs, outputs, **kwargs) + @functools.wraps(func) + def wrapper(self, *args, **kwargs): try: - return self._loss_instance(inputs, outputs, **kwargs) - except Exception as e: + return func(self, *args, **kwargs) + except Exception as exc: import traceback - logger.warning('[nccl_safe] Loss computation skipped due to error: ' - '%s: %s\n%s', - type(e).__name__, e, traceback.format_exc()) - return _zero_loss(outputs) - - def micro_batch_scale(self, inputs, indices): - """Preserve the wrapped loss's micro-batch reduction semantics.""" - return self._loss_instance.micro_batch_scale(inputs, indices) - - -def _zero_loss(outputs) -> 'LossOutput': - """Create a graph-connected zero loss for FSDP compatibility. - - Finds a gradient-bearing tensor from outputs to maintain graph connectivity, - ensuring backward hooks (ReduceScatter) fire. - """ - import torch - if isinstance(outputs, dict): - for key in ('logps', 'values', 'logits', 'loss'): - t = outputs.get(key) - if t is not None and isinstance(t, torch.Tensor) and t.requires_grad: - return LossOutput(loss=(t.flatten()[:1] * 0).sum(), num_tokens=0) - # Fallback: standalone zero tensor (may not trigger FSDP hooks) - device = 'cpu' - if isinstance(outputs, dict): - for v in outputs.values(): - if hasattr(v, 'device'): - device = v.device - break - return LossOutput(loss=torch.zeros((), device=device, requires_grad=True), num_tokens=0) - - -# ─── Layer 2: @nccl_safe decorator ────────────────────────────────────────── - - -def nccl_safe(func=None, *, tinker=False): - """Decorator ensuring backward() executes if forward() has already run. - - Detects forward completion by comparing train_status.outputs before/after - the wrapped function call. If an exception occurs after forward has run - but before backward completes, forces a zero-gradient backward pass to - prevent NCCL hang (other ranks waiting for ReduceScatter). - - Args: - func: The function to decorate (when used without arguments). - tinker: If True, fallback returns ``[[], 0.0]`` (tinker format). - If False, fallback returns outputs dict with ``loss=0.0``. - - Usage:: - - @remote_function(dispatch='slice_dp', collect=...) - @nccl_safe(tinker=True) - def tinker_forward_backward(self, *, inputs, adapter_name, ...): - # method body completely unchanged - ... - - @remote_function(dispatch='slice_dp', collect=...) - @nccl_safe - def forward_backward(self, *, inputs, **kwargs): - # method body completely unchanged - ... - """ - - def decorator(fn): - - @functools.wraps(fn) - def wrapper(self, *args, **kwargs): - if _is_fail_fast(): - return fn(self, *args, **kwargs) - - # Extract adapter_name for state tracking - adapter_name = kwargs.get('adapter_name') - if adapter_name is None and hasattr(self, '_get_default_group'): - adapter_name = self._get_default_group() - - og = self.optimizer_group.get(adapter_name) if adapter_name else None - if og is None: - # Cannot track state without optimizer group, passthrough - return fn(self, *args, **kwargs) - - # Snapshot state before call to detect forward completion - outputs_before = og.train_status.outputs - - try: - return fn(self, *args, **kwargs) - except Exception as e: - outputs_after = og.train_status.outputs - forward_ran = (outputs_after is not None and outputs_after is not outputs_before) - - if not forward_ran: - # Pre-forward failure: no NCCL ops started, safe to propagate - raise - - # Forward completed. Check if backward already ran. - # TransformersModel.backward() clears loss_value to None. - backward_done = (og.train_status.loss_value is None) - - if backward_done: - # Post-backward failure (e.g. output formatting) - # No NCCL hang risk, just return gracefully - logger.warning(f'[nccl_safe] Post-backward error (no NCCL risk): ' - f'{type(e).__name__}: {e}') - else: - # CRITICAL: forward ran but backward didn't → NCCL hang risk! - logger.warning(f'[nccl_safe] Forcing zero backward to prevent NCCL hang: ' - f'{type(e).__name__}: {e}') - _force_zero_backward(self, og, adapter_name, kwargs) - - # Return fallback result - if tinker: - return [[], 0.0] - outputs_after['loss'] = 0.0 - return outputs_after - - return wrapper - - if func is not None: - # @nccl_safe without arguments - return decorator(func) - # @nccl_safe(tinker=True) with arguments - return decorator - - -def _iter_model_params(model): - """Iterate parameters from ``model.model``, supporting single model or list of models.""" - raw_model = getattr(model, 'model', None) - if raw_model is None: - return iter([]) - if isinstance(raw_model, (list, tuple)): - for m in raw_model: - yield from m.parameters() - else: - yield from raw_model.parameters() - - -def _force_zero_backward(model, og, adapter_name, kwargs): - """Force a zero-gradient backward pass to prevent NCCL hang. - - Creates a graph-connected zero loss tensor and calls backward(), - ensuring FSDP ReduceScatter hooks fire on all ranks. - """ - import torch - - outputs = og.train_status.outputs - - # Find a graph-connected tensor for zero loss - zero_loss = None - if outputs is not None and isinstance(outputs, dict): - for key in ('logps', 'values', 'logits', 'loss'): - t = outputs.get(key) - if t is not None and isinstance(t, torch.Tensor) and t.requires_grad: - zero_loss = (t.flatten()[:1] * 0).sum() - break - - if zero_loss is None: - # Fallback: use first model parameter to maintain graph connectivity. - # Do NOT detach() the parameter -- the zero loss must remain connected - # to the model's autograd graph so FSDP ReduceScatter hooks fire. - # Use lazy iteration to avoid materializing the full parameter list. - try: - param = next((p for p in _iter_model_params(model) if p.requires_grad), None) - if param is not None: - zero_loss = (param.flatten()[0] * 0).sum() + rank = _global_rank() + context = f'twinkle backend method={func.__name__}, global_rank={rank}' + if hasattr(exc, 'add_note'): + exc.add_note(context) + elif exc.args: + exc.args = (f'{exc.args[0]} [{context}]', *exc.args[1:]) else: - zero_loss = torch.zeros((), device='cuda', requires_grad=True) - except Exception: - zero_loss = torch.zeros((), device='cuda', requires_grad=True) - - og.train_status.loss_value = zero_loss - - # Call backward with minimal kwargs - bwd_kwargs = {'adapter_name': adapter_name} - gas = kwargs.get('gradient_accumulation_steps') - if gas is not None: - bwd_kwargs['gradient_accumulation_steps'] = gas - model.backward(**bwd_kwargs) - - -# ─── Layer 3: @nccl_safe_megatron decorator ────────────────────────────────── - - -def nccl_safe_megatron(func=None, *, tinker=False, forward_only=False): - """Decorator for Megatron backend methods where the entire body is NCCL-critical. - - Unlike @nccl_safe (which detects forward/backward boundaries), this decorator - treats the **entire function** as a NCCL-critical section. In Megatron, - forward_only and forward_backward both call get_forward_backward_func() which - requires all DP ranks to enter synchronously. If one rank fails during data - preprocessing (before entering Megatron's scheduler), other ranks will hang - waiting for the collective. - - This decorator catches ALL exceptions (when TWINKLE_FAIL_FAST=0) and returns - a safe fallback value, preventing NCCL hang from asymmetric failures. - - Args: - func: The function to decorate (when used without arguments). - tinker: If True, fallback returns ``[[], 0.0]`` (tinker format). - forward_only: If True, fallback returns empty dict ``{}`` (forward_only format). - - Usage:: - - @remote_function(dispatch='slice_dp', collect=..., sync=True) - @nccl_safe_megatron - def forward_backward(self, *, inputs, **kwargs): - ... - - @remote_function(dispatch='slice_dp', collect=...) - @nccl_safe_megatron(forward_only=True) - def forward_only(self, *, inputs, **kwargs): - ... - - @remote_function(dispatch='slice_dp', collect=..., sync=True) - @nccl_safe_megatron(tinker=True) - def tinker_forward_backward(self, *, inputs, **kwargs): - ... - """ - - def decorator(fn): - - @functools.wraps(fn) - def wrapper(self, *args, **kwargs): - if _is_fail_fast(): - return fn(self, *args, **kwargs) - - try: - return fn(self, *args, **kwargs) - except Exception as e: - import traceback - logger.warning(f'[nccl_safe_megatron] Exception in Megatron method ' - f'{fn.__name__}: {type(e).__name__}: {e}\n' - f'{traceback.format_exc()}') - - # Return safe fallback to prevent NCCL hang on other ranks - if tinker: - return [[], 0.0] - if forward_only: - return {} - # forward_backward fallback: return dict with loss=0.0 - return {'loss': 0.0} - - return wrapper - - if func is not None: - # @nccl_safe_megatron without arguments - return decorator(func) - # @nccl_safe_megatron(tinker=True) with arguments - return decorator + exc.args = (context, ) + tb = traceback.format_exc() + if len(tb) > _TRACEBACK_LIMIT: + tb = tb[-_TRACEBACK_LIMIT:] + logger.error('[nccl_safe_megatron] %s in %s on global rank %s:\n%s', + type(exc).__name__, func.__name__, rank, tb) + raise + + return wrapper diff --git a/src/twinkle_client/types/base.py b/src/twinkle_client/types/base.py new file mode 100644 index 000000000..ff06b8f15 --- /dev/null +++ b/src/twinkle_client/types/base.py @@ -0,0 +1,90 @@ +# Copyright (c) ModelScope Contributors. All rights reserved. +"""Shared pydantic base classes and the naming rulings for the wire contract. + +This module is a public contract carrier imported across packages (Twinkle_Server +reverse-imports ``twinkle_client.types``); it therefore intentionally carries **no** +underscore prefix. + +Naming rulings (authoritative for all three split specs; kept in code, not only in +the spec, so a later reader cannot merge these away): + +1. Schema_Module modules imported across packages do NOT use an underscore prefix. + ``base.py`` / ``errors.py`` / ``lifecycle.py`` / ``data.py`` are public-contract + carriers; an underscore means "package-private", and a cross-package import of a + private module is a violation. Modules used only inside Twinkle_Client (never + imported by Twinkle_Server) are exempt. +2. New twinkle-native request models do NOT reuse a class name already present in + ``tinker.types``. Known collision to avoid: ``ForwardBackwardRequest``. Two + handlers import ``types`` from twinkle_client and from tinker respectively; a + same-named model is distinguished only by the import alias and is easy to + misread in a review diff. +3. The field expressing a failure-semantic category is named ``error_code``, NOT + ``status_code`` -- an execution-time failure is delivered with HTTP 200, so the + value is systematically unequal to the response status code. +4. A closed value set on a wire field is declared as ``Literal`` / enum, never a + bare ``str`` (see ``QueueStateLiteral`` in ``errors.py``). + +These three base classes are DEFINED here but NOT applied to any existing model by +this spec: applying ``extra='forbid'`` would immediately reject an old client's +request, which would break the zero-wire-change guarantee. +""" +from __future__ import annotations + +from pydantic import BaseModel, ConfigDict, Field +from pydantic.fields import FieldInfo +from typing import Any, Optional + + +class StrictRequest(BaseModel): + """Request bodies. Typos fail loudly.""" + + model_config = ConfigDict(frozen=True, extra='forbid') + + +class ResponseModel(BaseModel): + """Response bodies. An old client tolerates new server fields.""" + + model_config = ConfigDict(frozen=True, extra='ignore') + + +class DataModel(BaseModel): + """Data-plane models (InputFeature / Trajectory on the wire). + + Same ConfigDict as ResponseModel, different reason -- which is why this is a + separate class and not an alias. ResponseModel's ``ignore`` exists so an old + client tolerates new response fields. DataModel's ``ignore`` exists so a + user's Preprocessor / Template may leave harmless extra keys (the original + columns left by ``dataset.map``, say) without the request being rejected. + + Do NOT "fix" this to inherit StrictRequest. Doing so rejects those extra keys + and breaks a large number of existing datasets. + """ + + model_config = ConfigDict(frozen=True, extra='ignore') + + +# Key under which backend-applicability metadata is stored in a field's +# ``json_schema_extra``. A single constant, helper and reader -- kept here with the +# base classes rather than in a module of their own (no isolation benefit). +BACKEND_ONLY_KEY = 'twinkle_backend_only' + + +def backend_only(*backends: str, **field_kwargs: Any) -> FieldInfo: + """Mark a model field as applicable only to the given backend(s). + + Attaches the backend tuple to the field's ``json_schema_extra`` under + ``BACKEND_ONLY_KEY``; read it back with :func:`read_backend_only`. + """ + extra = dict(field_kwargs.pop('json_schema_extra', None) or {}) + extra[BACKEND_ONLY_KEY] = tuple(backends) + return Field(json_schema_extra=extra, **field_kwargs) + + +def read_backend_only(field_info: FieldInfo) -> Optional[tuple[str, ...]]: + """Return the backend tuple a field was tagged with, or ``None`` if untagged.""" + extra = getattr(field_info, 'json_schema_extra', None) + if isinstance(extra, dict): + value = extra.get(BACKEND_ONLY_KEY) + if value is not None: + return tuple(value) + return None diff --git a/src/twinkle_client/types/errors.py b/src/twinkle_client/types/errors.py new file mode 100644 index 000000000..f1f4e3e17 --- /dev/null +++ b/src/twinkle_client/types/errors.py @@ -0,0 +1,68 @@ +# Copyright (c) ModelScope Contributors. All rights reserved. +"""Structured failure payload -- the single representation of a failure. + +Twinkle <-> tinker exception mapping (verified, kept here so a future new exception +class can be lined up against its tinker counterpart): + +- Tinker 0.29.0 ``RequestFailedError`` (``tinker/_exceptions.py``; carries + ``message`` / ``request_id`` / ``category``) is the "the task completed in a failed + terminal state" exception. Its wire values are ``unknown`` / ``server`` / + ``user``, matching :class:`ErrorCategory`; legacy TitleCase values are normalized. + +``twinkle_client/utils/patch_tinker.py`` shows two SDKs can coexist in one process, +so a semantically-equal but differently-named exception must be lookup-able. +""" +from __future__ import annotations + +from enum import StrEnum +from pydantic import Field, field_validator, model_validator +from typing import Any, Literal, Optional + +from .base import ResponseModel + +# Closed value set, kept in sync with the server-side ``QueueState`` enum values +# (a consistency test asserts the two sets are equal). Wire fields carrying a queue +# state declare this alias, never a bare ``str`` (naming ruling 4). +QueueStateLiteral = Literal['active', 'paused_rate_limit', 'paused_capacity', 'unknown'] + + +class ErrorCategory(StrEnum): + """Error attribution. Matches tinker's ``RequestErrorCategory``.""" + + Unknown = 'unknown' + Server = 'server' + User = 'user' + + +class ErrorPayload(ResponseModel): + """The single representation of a failure, on the wire and in state. + + ``error_code``, not ``status_code``: once server-request-lifecycle lands, an + execution-time failure is delivered with HTTP 200, so this value is + *systematically* unequal to the response status code. Keeping the name + ``status_code`` would make every reader misparse it once. The 400-599 range is + kept to reuse HTTP's semantic space, not to align with response codes. + + Inherits ``ResponseModel`` (``extra='ignore'``), so a future added field does + not make an old client fail to parse it. + """ + + error: str = Field(max_length=1024) + category: ErrorCategory + error_code: int = Field(ge=400, le=599) + request_id: str + traceback: Optional[str] = Field(default=None, max_length=65536) + details: Optional[list[dict[str, Any]]] = None + + @field_validator('category', mode='before') + @classmethod + def normalize_legacy_category(cls, value: Any) -> Any: + if isinstance(value, str): + return value.lower() + return value + + @model_validator(mode='after') + def traceback_is_server_only(self) -> 'ErrorPayload': + if self.traceback is not None and self.category is not ErrorCategory.Server: + raise ValueError('traceback is only valid for server errors') + return self diff --git a/tests/infra/test_ray_get_timeout.py b/tests/infra/test_ray_get_timeout.py new file mode 100644 index 000000000..00a3b79c8 --- /dev/null +++ b/tests/infra/test_ray_get_timeout.py @@ -0,0 +1,137 @@ +# Copyright (c) ModelScope Contributors. All rights reserved. +"""Unit tests for the Sync_Dispatch_Path time bound (spec T1.6 / R9#1). + +These exercise only ``twinkle.infra`` against a plain sleeping Ray actor. They +depend on neither GPU, Megatron, nor any ``src/twinkle/server/**`` component. +""" +from __future__ import annotations + +import pytest + +ray = pytest.importorskip('ray') + +import twinkle.infra as infra # noqa: E402 +from twinkle.infra import remote_function # noqa: E402 +from twinkle.infra._ray.ray_helper import RayHelper # noqa: E402 + + +@ray.remote +class _Sleeper: + """A plain Ray actor whose only method sleeps for a caller-supplied time.""" + + def slow(self, seconds: float): + import time + time.sleep(seconds) + return seconds + + def slow_batch(self, seconds: list[float]): + import time + time.sleep(seconds[0]) + return seconds + + def _twinkle_async_slow_batch(self, seconds: list[float]): + return self.slow_batch(seconds) + + +@pytest.fixture(scope='module', autouse=True) +def _ray_and_ray_mode(): + """Bring up Ray and put infra into 'ray' mode for the driver-side path.""" + ray.init(ignore_reinit_error=True, num_cpus=2, logging_level='ERROR') + prev_mode = infra._mode + infra._mode = 'ray' + try: + yield + finally: + infra._mode = prev_mode + + +def _make_driver(): + """A minimal stand-in for a remote_class handle: one actor, no concurrency.""" + driver = type('Driver', (), {})() + driver._actors = [_Sleeper.remote()] + driver._max_concurrency = None + return driver + + +def test_execute_all_sync_times_out(_ray_and_ray_mode): + """R9#1: execute_all_sync(timeout=) raises when the remote does not return in time.""" + actor = _Sleeper.remote() + workers_and_args = [(actor, [3.0], {})] + with pytest.raises(ray.exceptions.GetTimeoutError): + RayHelper.execute_all_sync('slow', workers_and_args, timeout=0.5) + + +def test_execute_all_sync_returns_within_timeout(_ray_and_ray_mode): + actor = _Sleeper.remote() + workers_and_args = [(actor, [0.1], {})] + assert RayHelper.execute_all_sync('slow', workers_and_args, timeout=10.0) == [0.1] + + +def test_decorator_timeout_takes_priority_over_instance(): + """A small decorator timeout wins over a large instance ``_ray_get_timeout``.""" + + def slow(self, seconds): # body runs in the worker, not here + return seconds + + wrapped = remote_function(dispatch='all', collect='first', sync=True, timeout=0.5)(slow) + driver = _make_driver() + driver._ray_get_timeout = 100.0 # would allow the call if it were consulted + with pytest.raises(ray.exceptions.GetTimeoutError): + wrapped(driver, 3.0) + + +def test_decorator_timeout_wins_when_larger_than_instance(): + """The decorator value wins even when it is the *larger* of the two. + + A large decorator timeout with a tiny instance value must NOT time out -- + proving the instance value is ignored when the decorator declares one. + """ + + def slow(self, seconds): + return seconds + + wrapped = remote_function(dispatch='all', collect='first', sync=True, timeout=100.0)(slow) + driver = _make_driver() + driver._ray_get_timeout = 0.3 # would time out if it were consulted + result = wrapped(driver, 1.0) + # sync collect may hand back a lazy-collect callable; resolving it must not time out. + assert (result() if callable(result) else result) == 1.0 + + +def test_instance_timeout_is_fallback_when_decorator_absent(): + """With no decorator timeout, the instance ``_ray_get_timeout`` applies.""" + + def slow(self, seconds): + return seconds + + wrapped = remote_function(dispatch='all', collect='first', sync=True)(slow) + driver = _make_driver() + driver._ray_get_timeout = 0.3 + with pytest.raises(ray.exceptions.GetTimeoutError): + wrapped(driver, 2.0) + + +def test_continuous_work_timeout_zero_is_not_treated_as_falsy(): + + def slow_batch(self, seconds): + return seconds + + wrapped = remote_function( + dispatch='all', collect='first', timeout=0, enable_continous_work=True)(slow_batch) + driver = _make_driver() + driver._ray_get_timeout = 100.0 + with pytest.raises(ray.exceptions.GetTimeoutError): + wrapped(driver, [1.0]) + + +def test_decorator_timeout_zero_is_not_treated_as_falsy(): + """timeout=0 means 'time out immediately', not 'fall back to unbounded'.""" + + def slow(self, seconds): + return seconds + + wrapped = remote_function(dispatch='all', collect='first', sync=True, timeout=0)(slow) + driver = _make_driver() + driver._ray_get_timeout = 100.0 # the old ``or`` bug would fall back here + with pytest.raises(ray.exceptions.GetTimeoutError): + wrapped(driver, 1.0) diff --git a/tests/loss/test_grpo_gkd.py b/tests/loss/test_grpo_gkd.py index 3d0f55120..1f9daa707 100644 --- a/tests/loss/test_grpo_gkd.py +++ b/tests/loss/test_grpo_gkd.py @@ -82,6 +82,27 @@ def test_grpo_list_advantages(self): result = loss_fn(inputs, outputs, old_logps=old_logps, advantages=adv_list) assert torch.isfinite(result['loss']) + def test_pad_variable_length_full_sequence_rows(self): + """Unpadded per-sample rows align against a right-padded batch mask.""" + mask = torch.tensor([ + [False, True, True, False, False], + [False, False, True, True, True], + ]) + rows = [ + [-9.0, 0.1, 0.2], + [-9.0, -9.0, 0.3, 0.4, 0.5], + ] + + got = GRPOLoss()._pad_and_align_to_batch(rows, mask, mask.device, torch.float32) + + assert torch.equal(got[0], torch.tensor([0.0, 0.1, 0.2, 0.0, 0.0])) + assert torch.equal(got[1], torch.tensor([0.0, 0.0, 0.3, 0.4, 0.5])) + + def test_pad_rejects_full_sequence_missing_a_masked_position(self): + mask = torch.tensor([[False, False, True, True, True]]) + with pytest.raises(AssertionError, match='all masked positions'): + GRPOLoss()._pad_and_align_to_batch([[0.1, 0.2, 0.3, 0.4]], mask, mask.device, torch.float32) + def test_grpo_weights_sequences_equally(self): labels = torch.tensor([ [1, -100, -100], diff --git a/tests/model/test_micro_batch.py b/tests/model/test_micro_batch.py index 8d21c1dbb..ec4083e96 100644 --- a/tests/model/test_micro_batch.py +++ b/tests/model/test_micro_batch.py @@ -8,7 +8,6 @@ from twinkle.model.micro_batch import MicroBatchConfig, plan_micro_batches from twinkle.model.transformers.transformers import TransformersModel from twinkle.processor import InputProcessor -from twinkle.utils.nccl_safe import safe_loss @pytest.mark.parametrize('packing_algorithm', ['ffd', 'kk']) @@ -72,17 +71,6 @@ def test_sample_mean_and_token_sum_micro_batch_scales(): assert CrossEntropyLoss(reduction='sum').micro_batch_scale(inputs, [0]) == 1.0 -def test_safe_loss_preserves_wrapped_micro_batch_scale(): - inputs = [ - {'labels': [1, -100]}, - {'labels': [2, 3]}, - {'labels': [4, -100]}, - {'labels': [5, 6]}, - ] - - assert safe_loss(GRPOLoss()).micro_batch_scale(inputs, [0, 2]) == .5 - - def test_loss_without_micro_batch_semantics_fails_when_split(): with pytest.raises(NotImplementedError, match='does not support micro-batching'): Loss().micro_batch_scale([{}, {}], [0]) diff --git a/tests/model/test_multi_lora_target_parameters.py b/tests/model/test_multi_lora_target_parameters.py index b28ef6b36..c76854cea 100644 --- a/tests/model/test_multi_lora_target_parameters.py +++ b/tests/model/test_multi_lora_target_parameters.py @@ -2,6 +2,7 @@ import sys import types +import pytest import torch from peft import LoraConfig, get_peft_model from peft.utils import set_peft_model_state_dict @@ -45,6 +46,42 @@ def forward(self, x, expert_idx=0): return self.mlp.experts(x, expert_idx=expert_idx) +class FakePeftModule: + + def __init__(self, peft_config, active_adapter): + self.peft_config = peft_config + self.active_adapter = active_adapter + + +def test_save_context_restores_peft_state_when_load_fails(): + from twinkle.model.multi_lora import LoraTenant, MultiLora + + original_config = {'lora_0': object(), 'lora_1': object()} + modules = [ + FakePeftModule(original_config, 'lora_1'), + FakePeftModule(original_config, 'lora_0'), + ] + multi_lora = MultiLora(max_loras=2, max_r=4) + multi_lora.module = modules + multi_lora.loras = [ + LoraTenant( + index=0, + adapter_name='lora_0', + config=_make_target_cfg(), + tenant_adapter_name='tenant', + tenant_config=_make_target_cfg(), + ) + ] + + with pytest.raises(RuntimeError, match='load failed'): + with multi_lora.save_context('tenant') as adapter_name: + assert adapter_name == 'lora_0' + raise RuntimeError('load failed') + + assert [module.peft_config for module in modules] == [original_config, original_config] + assert [module.active_adapter for module in modules] == ['lora_1', 'lora_0'] + + def test_peft_target_parameter_key_shapes_for_3d_experts(): model = FakeModel() cfg = LoraConfig( @@ -266,4 +303,4 @@ def test_multilora_transformers_installs_target_parameters_once(): assert test_target_parameter_multi_lora_updates_only_active_adapter() == True assert test_multilora_releases_target_parameter_slot_to_initial_weights() == True assert test_multilora_state_dict_round_trips_target_parameters() == True - assert test_multilora_transformers_installs_target_parameters_once() == True \ No newline at end of file + assert test_multilora_transformers_installs_target_parameters_once() == True diff --git a/tests/server/conftest.py b/tests/server/conftest.py index 4ac4d24fd..fea2e62ca 100644 --- a/tests/server/conftest.py +++ b/tests/server/conftest.py @@ -1,5 +1,5 @@ # Copyright (c) ModelScope Contributors. All rights reserved. -"""Shared Ray runtime + per-test isolation for ``tests/server`` (state, cli, ...). +"""Shared Ray runtime, per-test isolation, and evidence boundaries. ``RayActorBackend`` is a forwarding wrapper around a detached Ray actor; instantiating one without an initialized Ray runtime raises @@ -12,6 +12,15 @@ of the actor wrapper. To keep tests independent we clear that actor's store before each test function. Tests that pin a non-default ``key_prefix`` get their own actor; this fixture intentionally leaves those alone. + +Evidence boundary (spec T8.4 / R9#10): every mock-model backend method accepts +``**kwargs`` without argument validation, and the mock enters no real collective. +A mock-backed test therefore proves neither request/argument validation nor NCCL +behavior (asymmetric failure, collective mis-pairing, ReduceScatter, etc.). It may +prove only backend dispatch, task-queue behavior, timeout/admission mechanisms, +and event-loop responsiveness. Validation and NCCL claims require the GPU-gated +``test_nccl_safe_*_e2e.py`` tests against a real server. The contract suite covers +all five apps; Tinker compatibility follows the pinned 0.16.1 SDK wire values. """ from __future__ import annotations diff --git a/tests/server/contract/client_api_baseline.json b/tests/server/contract/client_api_baseline.json index 65f9db147..44672f779 100644 --- a/tests/server/contract/client_api_baseline.json +++ b/tests/server/contract/client_api_baseline.json @@ -3,42 +3,406 @@ "paths": { "/twinkle/append": { "POST": { - "operationId": "append_twinkle_append_post", - "parameters": [], - "responses": [ - "200", - "422" - ] + "body": [ + { + "name": "body", + "required": true, + "schema": { + "$defs": { + "DataRef": { + "description": "Opaque reference to rows stored in the server-side TransferQueue.", + "properties": { + "fields": { + "items": { + "type": "string" + }, + "title": "Fields", + "type": "array" + }, + "kind": { + "default": "data", + "title": "Kind", + "type": "string" + }, + "num_tokens": { + "default": 0, + "title": "Num Tokens", + "type": "integer" + }, + "ref_id": { + "title": "Ref Id", + "type": "string" + }, + "size": { + "title": "Size", + "type": "integer" + } + }, + "required": [ + "ref_id", + "size" + ], + "title": "DataRef", + "type": "object" + } + }, + "properties": { + "ref": { + "$ref": "#/$defs/DataRef" + }, + "rows": { + "items": { + "additionalProperties": true, + "type": "object" + }, + "title": "Rows", + "type": "array" + }, + "tags": { + "anyOf": [ + { + "items": { + "additionalProperties": true, + "type": "object" + }, + "type": "array" + }, + { + "type": "null" + } + ], + "default": null, + "title": "Tags" + } + }, + "required": [ + "ref", + "rows" + ], + "title": "DataAppendRequest", + "type": "object" + } + } + ], + "cookies": [], + "headers": [], + "operationId": "append", + "path": [], + "query": [], + "response": { + "description": "Opaque reference to rows stored in the server-side TransferQueue.", + "properties": { + "fields": { + "items": { + "type": "string" + }, + "title": "Fields", + "type": "array" + }, + "kind": { + "default": "data", + "title": "Kind", + "type": "string" + }, + "num_tokens": { + "default": 0, + "title": "Num Tokens", + "type": "integer" + }, + "ref_id": { + "title": "Ref Id", + "type": "string" + }, + "size": { + "title": "Size", + "type": "integer" + } + }, + "required": [ + "ref_id", + "size" + ], + "title": "DataRef", + "type": "object" + }, + "responses": {}, + "statusCode": 200 } }, "/twinkle/get": { "POST": { - "operationId": "get_twinkle_get_post", - "parameters": [], - "responses": [ - "200", - "422" - ] + "body": [ + { + "name": "body", + "required": true, + "schema": { + "$defs": { + "DataRef": { + "description": "Opaque reference to rows stored in the server-side TransferQueue.", + "properties": { + "fields": { + "items": { + "type": "string" + }, + "title": "Fields", + "type": "array" + }, + "kind": { + "default": "data", + "title": "Kind", + "type": "string" + }, + "num_tokens": { + "default": 0, + "title": "Num Tokens", + "type": "integer" + }, + "ref_id": { + "title": "Ref Id", + "type": "string" + }, + "size": { + "title": "Size", + "type": "integer" + } + }, + "required": [ + "ref_id", + "size" + ], + "title": "DataRef", + "type": "object" + } + }, + "properties": { + "fields": { + "anyOf": [ + { + "items": { + "type": "string" + }, + "type": "array" + }, + { + "type": "null" + } + ], + "default": null, + "title": "Fields" + }, + "include_tags": { + "default": false, + "title": "Include Tags", + "type": "boolean" + }, + "ref": { + "$ref": "#/$defs/DataRef" + } + }, + "required": [ + "ref" + ], + "title": "DataGetRequest", + "type": "object" + } + } + ], + "cookies": [], + "headers": [], + "operationId": "get", + "path": [], + "query": [], + "response": { + "properties": { + "rows": { + "items": { + "additionalProperties": true, + "type": "object" + }, + "title": "Rows", + "type": "array" + }, + "tags": { + "items": { + "additionalProperties": true, + "type": "object" + }, + "title": "Tags", + "type": "array" + } + }, + "required": [ + "rows" + ], + "title": "DataRowsResponse", + "type": "object" + }, + "responses": {}, + "statusCode": 200 } }, "/twinkle/put": { "POST": { - "operationId": "put_twinkle_put_post", - "parameters": [], - "responses": [ - "200", - "422" - ] + "body": [ + { + "name": "body", + "required": true, + "schema": { + "properties": { + "kind": { + "default": "data", + "title": "Kind", + "type": "string" + }, + "rows": { + "items": { + "additionalProperties": true, + "type": "object" + }, + "title": "Rows", + "type": "array" + }, + "tags": { + "anyOf": [ + { + "items": { + "additionalProperties": true, + "type": "object" + }, + "type": "array" + }, + { + "type": "null" + } + ], + "default": null, + "title": "Tags" + } + }, + "required": [ + "rows" + ], + "title": "DataPutRequest", + "type": "object" + } + } + ], + "cookies": [], + "headers": [], + "operationId": "put", + "path": [], + "query": [], + "response": { + "description": "Opaque reference to rows stored in the server-side TransferQueue.", + "properties": { + "fields": { + "items": { + "type": "string" + }, + "title": "Fields", + "type": "array" + }, + "kind": { + "default": "data", + "title": "Kind", + "type": "string" + }, + "num_tokens": { + "default": 0, + "title": "Num Tokens", + "type": "integer" + }, + "ref_id": { + "title": "Ref Id", + "type": "string" + }, + "size": { + "title": "Size", + "type": "integer" + } + }, + "required": [ + "ref_id", + "size" + ], + "title": "DataRef", + "type": "object" + }, + "responses": {}, + "statusCode": 200 } }, "/twinkle/release": { "POST": { - "operationId": "release_twinkle_release_post", - "parameters": [], - "responses": [ - "200", - "422" - ] + "body": [ + { + "name": "body", + "required": true, + "schema": { + "$defs": { + "DataRef": { + "description": "Opaque reference to rows stored in the server-side TransferQueue.", + "properties": { + "fields": { + "items": { + "type": "string" + }, + "title": "Fields", + "type": "array" + }, + "kind": { + "default": "data", + "title": "Kind", + "type": "string" + }, + "num_tokens": { + "default": 0, + "title": "Num Tokens", + "type": "integer" + }, + "ref_id": { + "title": "Ref Id", + "type": "string" + }, + "size": { + "title": "Size", + "type": "integer" + } + }, + "required": [ + "ref_id", + "size" + ], + "title": "DataRef", + "type": "object" + } + }, + "properties": { + "ref": { + "$ref": "#/$defs/DataRef" + } + }, + "required": [ + "ref" + ], + "title": "DataReleaseRequest", + "type": "object" + } + } + ], + "cookies": [], + "headers": [], + "operationId": "release", + "path": [], + "query": [], + "response": { + "additionalProperties": { + "type": "string" + }, + "type": "object" + }, + "responses": {}, + "statusCode": 200 } } } @@ -47,531 +411,3690 @@ "paths": { "/asample": { "POST": { - "operationId": "asample_asample_post", - "parameters": [], - "responses": [ - "200", - "422" - ] + "body": [ + { + "name": "body", + "required": true, + "schema": { + "$defs": { + "EncodedTextChunk": { + "additionalProperties": false, + "properties": { + "tokens": { + "items": { + "type": "integer" + }, + "title": "Tokens", + "type": "array" + }, + "type": { + "const": "encoded_text", + "default": "encoded_text", + "title": "Type", + "type": "string" + } + }, + "required": [ + "tokens" + ], + "title": "EncodedTextChunk", + "type": "object" + }, + "ImageAssetPointerChunk": { + "additionalProperties": false, + "properties": { + "expected_tokens": { + "anyOf": [ + { + "type": "integer" + }, + { + "type": "null" + } + ], + "default": null, + "title": "Expected Tokens" + }, + "format": { + "enum": [ + "png", + "jpeg" + ], + "title": "Format", + "type": "string" + }, + "location": { + "title": "Location", + "type": "string" + }, + "type": { + "const": "image_asset_pointer", + "default": "image_asset_pointer", + "title": "Type", + "type": "string" + } + }, + "required": [ + "format", + "location" + ], + "title": "ImageAssetPointerChunk", + "type": "object" + }, + "ImageChunk": { + "additionalProperties": false, + "properties": { + "data": { + "format": "binary", + "title": "Data", + "type": "string" + }, + "expected_tokens": { + "anyOf": [ + { + "type": "integer" + }, + { + "type": "null" + } + ], + "default": null, + "title": "Expected Tokens" + }, + "format": { + "enum": [ + "png", + "jpeg" + ], + "title": "Format", + "type": "string" + }, + "type": { + "const": "image", + "default": "image", + "title": "Type", + "type": "string" + } + }, + "required": [ + "data", + "format" + ], + "title": "ImageChunk", + "type": "object" + }, + "ModelInput": { + "additionalProperties": false, + "properties": { + "chunks": { + "items": { + "anyOf": [ + { + "$ref": "#/$defs/EncodedTextChunk" + }, + { + "$ref": "#/$defs/ImageAssetPointerChunk" + }, + { + "$ref": "#/$defs/ImageChunk" + } + ] + }, + "title": "Chunks", + "type": "array" + } + }, + "required": [ + "chunks" + ], + "title": "ModelInput", + "type": "object" + }, + "SamplingParams": { + "properties": { + "max_tokens": { + "anyOf": [ + { + "type": "integer" + }, + { + "type": "null" + } + ], + "default": null, + "title": "Max Tokens" + }, + "seed": { + "anyOf": [ + { + "type": "integer" + }, + { + "type": "null" + } + ], + "default": null, + "title": "Seed" + }, + "stop": { + "anyOf": [ + { + "type": "string" + }, + { + "items": { + "type": "string" + }, + "type": "array" + }, + { + "items": { + "type": "integer" + }, + "type": "array" + }, + { + "type": "null" + } + ], + "default": null, + "title": "Stop" + }, + "temperature": { + "default": 1, + "title": "Temperature", + "type": "number" + }, + "top_k": { + "default": -1, + "title": "Top K", + "type": "integer" + }, + "top_p": { + "default": 1, + "title": "Top P", + "type": "number" + } + }, + "title": "SamplingParams", + "type": "object" + } + }, + "additionalProperties": false, + "properties": { + "base_model": { + "anyOf": [ + { + "type": "string" + }, + { + "type": "null" + } + ], + "default": null, + "title": "Base Model" + }, + "model_path": { + "anyOf": [ + { + "type": "string" + }, + { + "type": "null" + } + ], + "default": null, + "title": "Model Path" + }, + "num_samples": { + "default": 1, + "title": "Num Samples", + "type": "integer" + }, + "prompt": { + "$ref": "#/$defs/ModelInput" + }, + "prompt_logprobs": { + "anyOf": [ + { + "type": "boolean" + }, + { + "type": "null" + } + ], + "default": null, + "title": "Prompt Logprobs" + }, + "sampling_params": { + "$ref": "#/$defs/SamplingParams" + }, + "sampling_session_id": { + "anyOf": [ + { + "type": "string" + }, + { + "type": "null" + } + ], + "default": null, + "title": "Sampling Session Id" + }, + "seq_id": { + "anyOf": [ + { + "type": "integer" + }, + { + "type": "null" + } + ], + "default": null, + "title": "Seq Id" + }, + "topk_prompt_logprobs": { + "default": 0, + "title": "Topk Prompt Logprobs", + "type": "integer" + }, + "type": { + "const": "sample", + "default": "sample", + "title": "Type", + "type": "string" + } + }, + "required": [ + "prompt", + "sampling_params" + ], + "title": "SampleRequest", + "type": "object" + } + } + ], + "cookies": [], + "headers": [], + "operationId": "asample", + "path": [], + "query": [], + "response": {}, + "responses": {}, + "statusCode": 200 } }, "/chat/completions": { "POST": { - "operationId": "chat_completions_chat_completions_post", - "parameters": [], - "responses": [ - "200" - ] + "body": [], + "cookies": [], + "headers": [], + "operationId": "chat_completions", + "path": [], + "query": [], + "response": { + "type": "null" + }, + "responses": {}, + "statusCode": 200 } }, "/create_model": { "POST": { - "operationId": "create_model_create_model_post", - "parameters": [], - "responses": [ - "200", - "422" - ] + "body": [ + { + "name": "body", + "required": true, + "schema": { + "$defs": { + "LoraConfig": { + "additionalProperties": false, + "properties": { + "rank": { + "title": "Rank", + "type": "integer" + }, + "seed": { + "anyOf": [ + { + "type": "integer" + }, + { + "type": "null" + } + ], + "default": null, + "title": "Seed" + }, + "train_attn": { + "default": true, + "title": "Train Attn", + "type": "boolean" + }, + "train_mlp": { + "default": true, + "title": "Train Mlp", + "type": "boolean" + }, + "train_unembed": { + "default": true, + "title": "Train Unembed", + "type": "boolean" + } + }, + "required": [ + "rank" + ], + "title": "LoraConfig", + "type": "object" + } + }, + "additionalProperties": false, + "properties": { + "base_model": { + "title": "Base Model", + "type": "string" + }, + "lora_config": { + "anyOf": [ + { + "$ref": "#/$defs/LoraConfig" + }, + { + "type": "null" + } + ], + "default": null + }, + "model_seq_id": { + "title": "Model Seq Id", + "type": "integer" + }, + "session_id": { + "title": "Session Id", + "type": "string" + }, + "type": { + "const": "create_model", + "default": "create_model", + "title": "Type", + "type": "string" + }, + "user_metadata": { + "anyOf": [ + { + "additionalProperties": true, + "type": "object" + }, + { + "type": "null" + } + ], + "default": null, + "title": "User Metadata" + } + }, + "required": [ + "session_id", + "model_seq_id", + "base_model" + ], + "title": "CreateModelRequest", + "type": "object" + } + } + ], + "cookies": [], + "headers": [], + "operationId": "create_model", + "path": [], + "query": [], + "response": {}, + "responses": {}, + "statusCode": 200 } }, "/create_sampling_session": { "POST": { - "operationId": "create_sampling_session_create_sampling_session_post", - "parameters": [], - "responses": [ - "200", - "422" - ] + "body": [ + { + "name": "body", + "required": true, + "schema": { + "additionalProperties": false, + "properties": { + "base_model": { + "anyOf": [ + { + "type": "string" + }, + { + "type": "null" + } + ], + "default": null, + "title": "Base Model" + }, + "model_path": { + "anyOf": [ + { + "type": "string" + }, + { + "type": "null" + } + ], + "default": null, + "title": "Model Path" + }, + "sampling_session_seq_id": { + "title": "Sampling Session Seq Id", + "type": "integer" + }, + "session_id": { + "title": "Session Id", + "type": "string" + }, + "type": { + "const": "create_sampling_session", + "default": "create_sampling_session", + "title": "Type", + "type": "string" + } + }, + "required": [ + "session_id", + "sampling_session_seq_id" + ], + "title": "CreateSamplingSessionRequest", + "type": "object" + } + } + ], + "cookies": [], + "headers": [], + "operationId": "create_sampling_session", + "path": [], + "query": [], + "response": { + "properties": { + "sampling_session_id": { + "title": "Sampling Session Id", + "type": "string" + }, + "type": { + "const": "create_sampling_session", + "default": "create_sampling_session", + "title": "Type", + "type": "string" + } + }, + "required": [ + "sampling_session_id" + ], + "title": "CreateSamplingSessionResponse", + "type": "object" + }, + "responses": {}, + "statusCode": 200 } }, "/create_session": { "POST": { - "operationId": "create_session_create_session_post", - "parameters": [], - "responses": [ - "200", - "422" - ] + "body": [ + { + "name": "body", + "required": true, + "schema": { + "additionalProperties": false, + "properties": { + "project_id": { + "anyOf": [ + { + "type": "string" + }, + { + "type": "null" + } + ], + "default": null, + "title": "Project Id" + }, + "sdk_version": { + "title": "Sdk Version", + "type": "string" + }, + "tags": { + "items": { + "type": "string" + }, + "title": "Tags", + "type": "array" + }, + "type": { + "const": "create_session", + "default": "create_session", + "title": "Type", + "type": "string" + }, + "user_metadata": { + "anyOf": [ + { + "additionalProperties": true, + "type": "object" + }, + { + "type": "null" + } + ], + "title": "User Metadata" + } + }, + "required": [ + "tags", + "user_metadata", + "sdk_version" + ], + "title": "CreateSessionRequest", + "type": "object" + } + } + ], + "cookies": [], + "headers": [], + "operationId": "create_session", + "path": [], + "query": [], + "response": { + "properties": { + "error_message": { + "anyOf": [ + { + "type": "string" + }, + { + "type": "null" + } + ], + "default": null, + "title": "Error Message" + }, + "info_message": { + "anyOf": [ + { + "type": "string" + }, + { + "type": "null" + } + ], + "default": null, + "title": "Info Message" + }, + "session_id": { + "title": "Session Id", + "type": "string" + }, + "type": { + "const": "create_session", + "default": "create_session", + "title": "Type", + "type": "string" + }, + "warning_message": { + "anyOf": [ + { + "type": "string" + }, + { + "type": "null" + } + ], + "default": null, + "title": "Warning Message" + } + }, + "required": [ + "session_id" + ], + "title": "CreateSessionResponse", + "type": "object" + }, + "responses": {}, + "statusCode": 200 } }, "/forward": { "POST": { - "operationId": "forward_forward_post", - "parameters": [], - "responses": [ - "200", - "422" - ] + "body": [ + { + "name": "body", + "required": true, + "schema": { + "$defs": { + "Datum": { + "additionalProperties": false, + "properties": { + "loss_fn_inputs": { + "additionalProperties": { + "$ref": "#/$defs/TensorData" + }, + "title": "Loss Fn Inputs", + "type": "object" + }, + "model_input": { + "$ref": "#/$defs/ModelInput" + } + }, + "required": [ + "loss_fn_inputs", + "model_input" + ], + "title": "Datum", + "type": "object" + }, + "EncodedTextChunk": { + "additionalProperties": false, + "properties": { + "tokens": { + "items": { + "type": "integer" + }, + "title": "Tokens", + "type": "array" + }, + "type": { + "const": "encoded_text", + "default": "encoded_text", + "title": "Type", + "type": "string" + } + }, + "required": [ + "tokens" + ], + "title": "EncodedTextChunk", + "type": "object" + }, + "ForwardBackwardInput": { + "additionalProperties": false, + "properties": { + "data": { + "items": { + "$ref": "#/$defs/Datum" + }, + "title": "Data", + "type": "array" + }, + "loss_fn": { + "enum": [ + "cross_entropy", + "importance_sampling", + "ppo", + "cispo", + "dro" + ], + "title": "Loss Fn", + "type": "string" + }, + "loss_fn_config": { + "anyOf": [ + { + "additionalProperties": { + "type": "number" + }, + "type": "object" + }, + { + "type": "null" + } + ], + "default": null, + "title": "Loss Fn Config" + } + }, + "required": [ + "data", + "loss_fn" + ], + "title": "ForwardBackwardInput", + "type": "object" + }, + "ImageAssetPointerChunk": { + "additionalProperties": false, + "properties": { + "expected_tokens": { + "anyOf": [ + { + "type": "integer" + }, + { + "type": "null" + } + ], + "default": null, + "title": "Expected Tokens" + }, + "format": { + "enum": [ + "png", + "jpeg" + ], + "title": "Format", + "type": "string" + }, + "location": { + "title": "Location", + "type": "string" + }, + "type": { + "const": "image_asset_pointer", + "default": "image_asset_pointer", + "title": "Type", + "type": "string" + } + }, + "required": [ + "format", + "location" + ], + "title": "ImageAssetPointerChunk", + "type": "object" + }, + "ImageChunk": { + "additionalProperties": false, + "properties": { + "data": { + "format": "binary", + "title": "Data", + "type": "string" + }, + "expected_tokens": { + "anyOf": [ + { + "type": "integer" + }, + { + "type": "null" + } + ], + "default": null, + "title": "Expected Tokens" + }, + "format": { + "enum": [ + "png", + "jpeg" + ], + "title": "Format", + "type": "string" + }, + "type": { + "const": "image", + "default": "image", + "title": "Type", + "type": "string" + } + }, + "required": [ + "data", + "format" + ], + "title": "ImageChunk", + "type": "object" + }, + "ModelInput": { + "additionalProperties": false, + "properties": { + "chunks": { + "items": { + "anyOf": [ + { + "$ref": "#/$defs/EncodedTextChunk" + }, + { + "$ref": "#/$defs/ImageAssetPointerChunk" + }, + { + "$ref": "#/$defs/ImageChunk" + } + ] + }, + "title": "Chunks", + "type": "array" + } + }, + "required": [ + "chunks" + ], + "title": "ModelInput", + "type": "object" + }, + "TensorData": { + "additionalProperties": false, + "properties": { + "data": { + "anyOf": [ + { + "items": { + "type": "integer" + }, + "type": "array" + }, + { + "items": { + "type": "number" + }, + "type": "array" + } + ], + "title": "Data" + }, + "dtype": { + "enum": [ + "int64", + "float32" + ], + "title": "Dtype", + "type": "string" + }, + "shape": { + "anyOf": [ + { + "items": { + "type": "integer" + }, + "type": "array" + }, + { + "type": "null" + } + ], + "default": null, + "title": "Shape" + } + }, + "required": [ + "data", + "dtype" + ], + "title": "TensorData", + "type": "object" + } + }, + "additionalProperties": false, + "properties": { + "forward_input": { + "$ref": "#/$defs/ForwardBackwardInput" + }, + "model_id": { + "title": "Model Id", + "type": "string" + }, + "seq_id": { + "anyOf": [ + { + "type": "integer" + }, + { + "type": "null" + } + ], + "default": null, + "title": "Seq Id" + } + }, + "required": [ + "forward_input", + "model_id" + ], + "title": "ForwardRequest", + "type": "object" + } + } + ], + "cookies": [], + "headers": [], + "operationId": "forward", + "path": [], + "query": [], + "response": {}, + "responses": {}, + "statusCode": 200 } }, "/forward_backward": { "POST": { - "operationId": "forward_backward_forward_backward_post", - "parameters": [], - "responses": [ - "200", - "422" - ] + "body": [ + { + "name": "body", + "required": true, + "schema": { + "$defs": { + "Datum": { + "additionalProperties": false, + "properties": { + "loss_fn_inputs": { + "additionalProperties": { + "$ref": "#/$defs/TensorData" + }, + "title": "Loss Fn Inputs", + "type": "object" + }, + "model_input": { + "$ref": "#/$defs/ModelInput" + } + }, + "required": [ + "loss_fn_inputs", + "model_input" + ], + "title": "Datum", + "type": "object" + }, + "EncodedTextChunk": { + "additionalProperties": false, + "properties": { + "tokens": { + "items": { + "type": "integer" + }, + "title": "Tokens", + "type": "array" + }, + "type": { + "const": "encoded_text", + "default": "encoded_text", + "title": "Type", + "type": "string" + } + }, + "required": [ + "tokens" + ], + "title": "EncodedTextChunk", + "type": "object" + }, + "ForwardBackwardInput": { + "additionalProperties": false, + "properties": { + "data": { + "items": { + "$ref": "#/$defs/Datum" + }, + "title": "Data", + "type": "array" + }, + "loss_fn": { + "enum": [ + "cross_entropy", + "importance_sampling", + "ppo", + "cispo", + "dro" + ], + "title": "Loss Fn", + "type": "string" + }, + "loss_fn_config": { + "anyOf": [ + { + "additionalProperties": { + "type": "number" + }, + "type": "object" + }, + { + "type": "null" + } + ], + "default": null, + "title": "Loss Fn Config" + } + }, + "required": [ + "data", + "loss_fn" + ], + "title": "ForwardBackwardInput", + "type": "object" + }, + "ImageAssetPointerChunk": { + "additionalProperties": false, + "properties": { + "expected_tokens": { + "anyOf": [ + { + "type": "integer" + }, + { + "type": "null" + } + ], + "default": null, + "title": "Expected Tokens" + }, + "format": { + "enum": [ + "png", + "jpeg" + ], + "title": "Format", + "type": "string" + }, + "location": { + "title": "Location", + "type": "string" + }, + "type": { + "const": "image_asset_pointer", + "default": "image_asset_pointer", + "title": "Type", + "type": "string" + } + }, + "required": [ + "format", + "location" + ], + "title": "ImageAssetPointerChunk", + "type": "object" + }, + "ImageChunk": { + "additionalProperties": false, + "properties": { + "data": { + "format": "binary", + "title": "Data", + "type": "string" + }, + "expected_tokens": { + "anyOf": [ + { + "type": "integer" + }, + { + "type": "null" + } + ], + "default": null, + "title": "Expected Tokens" + }, + "format": { + "enum": [ + "png", + "jpeg" + ], + "title": "Format", + "type": "string" + }, + "type": { + "const": "image", + "default": "image", + "title": "Type", + "type": "string" + } + }, + "required": [ + "data", + "format" + ], + "title": "ImageChunk", + "type": "object" + }, + "ModelInput": { + "additionalProperties": false, + "properties": { + "chunks": { + "items": { + "anyOf": [ + { + "$ref": "#/$defs/EncodedTextChunk" + }, + { + "$ref": "#/$defs/ImageAssetPointerChunk" + }, + { + "$ref": "#/$defs/ImageChunk" + } + ] + }, + "title": "Chunks", + "type": "array" + } + }, + "required": [ + "chunks" + ], + "title": "ModelInput", + "type": "object" + }, + "TensorData": { + "additionalProperties": false, + "properties": { + "data": { + "anyOf": [ + { + "items": { + "type": "integer" + }, + "type": "array" + }, + { + "items": { + "type": "number" + }, + "type": "array" + } + ], + "title": "Data" + }, + "dtype": { + "enum": [ + "int64", + "float32" + ], + "title": "Dtype", + "type": "string" + }, + "shape": { + "anyOf": [ + { + "items": { + "type": "integer" + }, + "type": "array" + }, + { + "type": "null" + } + ], + "default": null, + "title": "Shape" + } + }, + "required": [ + "data", + "dtype" + ], + "title": "TensorData", + "type": "object" + } + }, + "additionalProperties": false, + "properties": { + "forward_backward_input": { + "$ref": "#/$defs/ForwardBackwardInput" + }, + "model_id": { + "title": "Model Id", + "type": "string" + }, + "seq_id": { + "anyOf": [ + { + "type": "integer" + }, + { + "type": "null" + } + ], + "default": null, + "title": "Seq Id" + } + }, + "required": [ + "forward_backward_input", + "model_id" + ], + "title": "ForwardBackwardRequest", + "type": "object" + } + } + ], + "cookies": [], + "headers": [], + "operationId": "forward_backward", + "path": [], + "query": [], + "response": {}, + "responses": {}, + "statusCode": 200 } }, "/get_info": { "POST": { - "operationId": "get_info_get_info_post", - "parameters": [], - "responses": [ - "200", - "422" - ] + "body": [ + { + "name": "body", + "required": true, + "schema": { + "additionalProperties": false, + "properties": { + "model_id": { + "title": "Model Id", + "type": "string" + }, + "type": { + "const": "get_info", + "default": "get_info", + "title": "Type", + "type": "string" + } + }, + "required": [ + "model_id" + ], + "title": "GetInfoRequest", + "type": "object" + } + } + ], + "cookies": [], + "headers": [], + "operationId": "get_info", + "path": [], + "query": [], + "response": {}, + "responses": {}, + "statusCode": 200 } }, "/get_server_capabilities": { "GET": { - "operationId": "get_server_capabilities_get_server_capabilities_get", - "parameters": [], - "responses": [ - "200" - ] + "body": [], + "cookies": [], + "headers": [], + "operationId": "get_server_capabilities", + "path": [], + "query": [], + "response": { + "$defs": { + "SupportedModel": { + "description": "Information about a model supported by the server.", + "properties": { + "model_name": { + "anyOf": [ + { + "type": "string" + }, + { + "type": "null" + } + ], + "default": null, + "title": "Model Name" + } + }, + "title": "SupportedModel", + "type": "object" + } + }, + "description": "Response containing the server's supported models and capabilities.", + "properties": { + "supported_models": { + "items": { + "$ref": "#/$defs/SupportedModel" + }, + "title": "Supported Models", + "type": "array" + } + }, + "required": [ + "supported_models" + ], + "title": "GetServerCapabilitiesResponse", + "type": "object" + }, + "responses": {}, + "statusCode": 200 } }, "/healthz": { "GET": { - "operationId": "healthz_healthz_get", - "parameters": [], - "responses": [ - "200" - ] + "body": [], + "cookies": [], + "headers": [], + "operationId": "healthz", + "path": [], + "query": [], + "response": { + "properties": { + "status": { + "const": "ok", + "title": "Status", + "type": "string" + } + }, + "required": [ + "status" + ], + "title": "HealthResponse", + "type": "object" + }, + "responses": {}, + "statusCode": 200 } }, "/load_weights": { "POST": { - "operationId": "load_weights_load_weights_post", - "parameters": [], - "responses": [ - "200", - "422" - ] + "body": [ + { + "name": "body", + "required": true, + "schema": { + "additionalProperties": false, + "properties": { + "model_id": { + "title": "Model Id", + "type": "string" + }, + "optimizer": { + "title": "Optimizer", + "type": "boolean" + }, + "path": { + "title": "Path", + "type": "string" + }, + "seq_id": { + "anyOf": [ + { + "type": "integer" + }, + { + "type": "null" + } + ], + "default": null, + "title": "Seq Id" + }, + "type": { + "const": "load_weights", + "default": "load_weights", + "title": "Type", + "type": "string" + } + }, + "required": [ + "model_id", + "path", + "optimizer" + ], + "title": "LoadWeightsRequest", + "type": "object" + } + } + ], + "cookies": [], + "headers": [], + "operationId": "load_weights", + "path": [], + "query": [], + "response": {}, + "responses": {}, + "statusCode": 200 } }, "/models": { "GET": { - "operationId": "list_models_models_get", - "parameters": [], - "responses": [ - "200" - ] + "body": [], + "cookies": [], + "headers": [], + "operationId": "list_models", + "path": [], + "query": [], + "response": { + "type": "null" + }, + "responses": {}, + "statusCode": 200 } }, "/optim_step": { "POST": { - "operationId": "optim_step_optim_step_post", - "parameters": [], - "responses": [ - "200", - "422" - ] + "body": [ + { + "name": "body", + "required": true, + "schema": { + "$defs": { + "AdamParams": { + "additionalProperties": false, + "properties": { + "beta1": { + "default": 0.9, + "title": "Beta1", + "type": "number" + }, + "beta2": { + "default": 0.95, + "title": "Beta2", + "type": "number" + }, + "eps": { + "default": 1e-12, + "title": "Eps", + "type": "number" + }, + "grad_clip_norm": { + "default": 0.0, + "title": "Grad Clip Norm", + "type": "number" + }, + "learning_rate": { + "default": 0.0001, + "title": "Learning Rate", + "type": "number" + }, + "weight_decay": { + "default": 0.0, + "title": "Weight Decay", + "type": "number" + } + }, + "title": "AdamParams", + "type": "object" + } + }, + "additionalProperties": false, + "properties": { + "adam_params": { + "$ref": "#/$defs/AdamParams" + }, + "model_id": { + "title": "Model Id", + "type": "string" + }, + "seq_id": { + "anyOf": [ + { + "type": "integer" + }, + { + "type": "null" + } + ], + "default": null, + "title": "Seq Id" + }, + "type": { + "const": "optim_step", + "default": "optim_step", + "title": "Type", + "type": "string" + } + }, + "required": [ + "adam_params", + "model_id" + ], + "title": "OptimStepRequest", + "type": "object" + } + } + ], + "cookies": [], + "headers": [], + "operationId": "optim_step", + "path": [], + "query": [], + "response": {}, + "responses": {}, + "statusCode": 200 } }, "/retrieve_future": { "POST": { - "operationId": "retrieve_future_retrieve_future_post", - "parameters": [], - "responses": [ - "200", - "422" - ] + "body": [ + { + "name": "body", + "required": true, + "schema": { + "additionalProperties": false, + "properties": { + "allow_metadata_only": { + "default": false, + "title": "Allow Metadata Only", + "type": "boolean" + }, + "request_id": { + "title": "Request Id", + "type": "string" + } + }, + "required": [ + "request_id" + ], + "title": "FutureRetrieveRequest", + "type": "object" + } + } + ], + "cookies": [], + "headers": [], + "operationId": "retrieve_future", + "path": [], + "query": [], + "response": {}, + "responses": {}, + "statusCode": 200 } }, "/save_weights": { "POST": { - "operationId": "save_weights_save_weights_post", - "parameters": [], - "responses": [ - "200", - "422" - ] + "body": [ + { + "name": "body", + "required": true, + "schema": { + "additionalProperties": false, + "properties": { + "model_id": { + "title": "Model Id", + "type": "string" + }, + "path": { + "anyOf": [ + { + "type": "string" + }, + { + "type": "null" + } + ], + "default": null, + "title": "Path" + }, + "seq_id": { + "anyOf": [ + { + "type": "integer" + }, + { + "type": "null" + } + ], + "default": null, + "title": "Seq Id" + }, + "ttl_seconds": { + "anyOf": [ + { + "type": "integer" + }, + { + "type": "null" + } + ], + "default": null, + "title": "Ttl Seconds" + }, + "type": { + "const": "save_weights", + "default": "save_weights", + "title": "Type", + "type": "string" + } + }, + "required": [ + "model_id" + ], + "title": "SaveWeightsRequest", + "type": "object" + } + } + ], + "cookies": [], + "headers": [], + "operationId": "save_weights", + "path": [], + "query": [], + "response": {}, + "responses": {}, + "statusCode": 200 } }, "/save_weights_for_sampler": { "POST": { - "operationId": "save_weights_for_sampler_save_weights_for_sampler_post", - "parameters": [], - "responses": [ - "200", - "422" - ] + "body": [ + { + "name": "body", + "required": true, + "schema": { + "additionalProperties": false, + "properties": { + "model_id": { + "title": "Model Id", + "type": "string" + }, + "path": { + "anyOf": [ + { + "type": "string" + }, + { + "type": "null" + } + ], + "default": null, + "title": "Path" + }, + "sampling_session_seq_id": { + "anyOf": [ + { + "type": "integer" + }, + { + "type": "null" + } + ], + "default": null, + "title": "Sampling Session Seq Id" + }, + "seq_id": { + "anyOf": [ + { + "type": "integer" + }, + { + "type": "null" + } + ], + "default": null, + "title": "Seq Id" + }, + "ttl_seconds": { + "anyOf": [ + { + "type": "integer" + }, + { + "type": "null" + } + ], + "default": null, + "title": "Ttl Seconds" + }, + "type": { + "const": "save_weights_for_sampler", + "default": "save_weights_for_sampler", + "title": "Type", + "type": "string" + } + }, + "required": [ + "model_id" + ], + "title": "SaveWeightsForSamplerRequest", + "type": "object" + } + } + ], + "cookies": [], + "headers": [], + "operationId": "save_weights_for_sampler", + "path": [], + "query": [], + "response": {}, + "responses": {}, + "statusCode": 200 } }, "/session_heartbeat": { "POST": { - "operationId": "session_heartbeat_session_heartbeat_post", - "parameters": [], - "responses": [ - "200", - "422" - ] + "body": [ + { + "name": "body", + "required": true, + "schema": { + "additionalProperties": false, + "properties": { + "session_id": { + "title": "Session Id", + "type": "string" + }, + "type": { + "const": "session_heartbeat", + "default": "session_heartbeat", + "title": "Type", + "type": "string" + } + }, + "required": [ + "session_id" + ], + "title": "SessionHeartbeatRequest", + "type": "object" + } + } + ], + "cookies": [], + "headers": [], + "operationId": "session_heartbeat", + "path": [], + "query": [], + "response": { + "properties": { + "type": { + "const": "session_heartbeat", + "default": "session_heartbeat", + "title": "Type", + "type": "string" + } + }, + "title": "SessionHeartbeatResponse", + "type": "object" + }, + "responses": {}, + "statusCode": 200 } }, "/telemetry": { "POST": { - "operationId": "telemetry_telemetry_post", - "parameters": [], - "responses": [ - "200", - "422" - ] + "body": [ + { + "name": "body", + "required": true, + "schema": { + "$defs": { + "GenericEvent": { + "properties": { + "event": { + "enum": [ + "SESSION_START", + "SESSION_END", + "UNHANDLED_EXCEPTION", + "GENERIC_EVENT" + ], + "title": "Event", + "type": "string" + }, + "event_data": { + "additionalProperties": true, + "default": {}, + "title": "Event Data", + "type": "object" + }, + "event_id": { + "title": "Event Id", + "type": "string" + }, + "event_name": { + "title": "Event Name", + "type": "string" + }, + "event_session_index": { + "title": "Event Session Index", + "type": "integer" + }, + "severity": { + "enum": [ + "DEBUG", + "INFO", + "WARNING", + "ERROR", + "CRITICAL" + ], + "title": "Severity", + "type": "string" + }, + "timestamp": { + "format": "date-time", + "title": "Timestamp", + "type": "string" + } + }, + "required": [ + "event", + "event_id", + "event_name", + "event_session_index", + "severity", + "timestamp" + ], + "title": "GenericEvent", + "type": "object" + }, + "SessionEndEvent": { + "properties": { + "duration": { + "title": "Duration", + "type": "string" + }, + "event": { + "enum": [ + "SESSION_START", + "SESSION_END", + "UNHANDLED_EXCEPTION", + "GENERIC_EVENT" + ], + "title": "Event", + "type": "string" + }, + "event_id": { + "title": "Event Id", + "type": "string" + }, + "event_session_index": { + "title": "Event Session Index", + "type": "integer" + }, + "severity": { + "enum": [ + "DEBUG", + "INFO", + "WARNING", + "ERROR", + "CRITICAL" + ], + "title": "Severity", + "type": "string" + }, + "timestamp": { + "format": "date-time", + "title": "Timestamp", + "type": "string" + } + }, + "required": [ + "duration", + "event", + "event_id", + "event_session_index", + "severity", + "timestamp" + ], + "title": "SessionEndEvent", + "type": "object" + }, + "SessionStartEvent": { + "properties": { + "event": { + "enum": [ + "SESSION_START", + "SESSION_END", + "UNHANDLED_EXCEPTION", + "GENERIC_EVENT" + ], + "title": "Event", + "type": "string" + }, + "event_id": { + "title": "Event Id", + "type": "string" + }, + "event_session_index": { + "title": "Event Session Index", + "type": "integer" + }, + "severity": { + "enum": [ + "DEBUG", + "INFO", + "WARNING", + "ERROR", + "CRITICAL" + ], + "title": "Severity", + "type": "string" + }, + "timestamp": { + "format": "date-time", + "title": "Timestamp", + "type": "string" + } + }, + "required": [ + "event", + "event_id", + "event_session_index", + "severity", + "timestamp" + ], + "title": "SessionStartEvent", + "type": "object" + }, + "UnhandledExceptionEvent": { + "properties": { + "error_message": { + "title": "Error Message", + "type": "string" + }, + "error_type": { + "title": "Error Type", + "type": "string" + }, + "event": { + "enum": [ + "SESSION_START", + "SESSION_END", + "UNHANDLED_EXCEPTION", + "GENERIC_EVENT" + ], + "title": "Event", + "type": "string" + }, + "event_id": { + "title": "Event Id", + "type": "string" + }, + "event_session_index": { + "title": "Event Session Index", + "type": "integer" + }, + "severity": { + "enum": [ + "DEBUG", + "INFO", + "WARNING", + "ERROR", + "CRITICAL" + ], + "title": "Severity", + "type": "string" + }, + "timestamp": { + "format": "date-time", + "title": "Timestamp", + "type": "string" + }, + "traceback": { + "anyOf": [ + { + "type": "string" + }, + { + "type": "null" + } + ], + "default": null, + "title": "Traceback" + } + }, + "required": [ + "error_message", + "error_type", + "event", + "event_id", + "event_session_index", + "severity", + "timestamp" + ], + "title": "UnhandledExceptionEvent", + "type": "object" + } + }, + "additionalProperties": false, + "properties": { + "events": { + "items": { + "anyOf": [ + { + "$ref": "#/$defs/SessionStartEvent" + }, + { + "$ref": "#/$defs/SessionEndEvent" + }, + { + "$ref": "#/$defs/UnhandledExceptionEvent" + }, + { + "$ref": "#/$defs/GenericEvent" + } + ] + }, + "title": "Events", + "type": "array" + }, + "platform": { + "title": "Platform", + "type": "string" + }, + "sdk_version": { + "title": "Sdk Version", + "type": "string" + }, + "session_id": { + "title": "Session Id", + "type": "string" + } + }, + "required": [ + "events", + "platform", + "sdk_version", + "session_id" + ], + "title": "TelemetrySendRequest", + "type": "object" + } + } + ], + "cookies": [], + "headers": [], + "operationId": "telemetry", + "path": [], + "query": [], + "response": { + "properties": { + "status": { + "const": "accepted", + "title": "Status", + "type": "string" + } + }, + "required": [ + "status" + ], + "title": "TelemetryResponse", + "type": "object" + }, + "responses": {}, + "statusCode": 200 } }, "/training_runs": { "GET": { - "operationId": "get_training_runs_training_runs_get", - "parameters": [ + "body": [], + "cookies": [], + "headers": [], + "operationId": "get_training_runs", + "path": [], + "query": [ { - "in": "query", "name": "limit", "required": false, "schema": { - "default": 20, - "title": "Limit", "type": "integer" } }, { - "in": "query", "name": "offset", "required": false, "schema": { - "default": 0, - "title": "Offset", "type": "integer" } } ], - "responses": [ - "200", - "422" - ] + "response": { + "$defs": { + "Checkpoint": { + "properties": { + "checkpoint_id": { + "title": "Checkpoint Id", + "type": "string" + }, + "checkpoint_type": { + "enum": [ + "training", + "sampler" + ], + "title": "Checkpoint Type", + "type": "string" + }, + "expires_at": { + "anyOf": [ + { + "format": "date-time", + "type": "string" + }, + { + "type": "null" + } + ], + "default": null, + "title": "Expires At" + }, + "public": { + "default": false, + "title": "Public", + "type": "boolean" + }, + "size_bytes": { + "anyOf": [ + { + "type": "integer" + }, + { + "type": "null" + } + ], + "default": null, + "title": "Size Bytes" + }, + "time": { + "format": "date-time", + "title": "Time", + "type": "string" + }, + "tinker_path": { + "title": "Tinker Path", + "type": "string" + } + }, + "required": [ + "checkpoint_id", + "checkpoint_type", + "time", + "tinker_path" + ], + "title": "Checkpoint", + "type": "object" + }, + "Cursor": { + "properties": { + "limit": { + "title": "Limit", + "type": "integer" + }, + "offset": { + "title": "Offset", + "type": "integer" + }, + "total_count": { + "title": "Total Count", + "type": "integer" + } + }, + "required": [ + "offset", + "limit", + "total_count" + ], + "title": "Cursor", + "type": "object" + }, + "TrainingRun": { + "properties": { + "base_model": { + "title": "Base Model", + "type": "string" + }, + "corrupted": { + "default": false, + "title": "Corrupted", + "type": "boolean" + }, + "is_lora": { + "title": "Is Lora", + "type": "boolean" + }, + "last_checkpoint": { + "anyOf": [ + { + "$ref": "#/$defs/Checkpoint" + }, + { + "type": "null" + } + ], + "default": null + }, + "last_request_time": { + "format": "date-time", + "title": "Last Request Time", + "type": "string" + }, + "last_sampler_checkpoint": { + "anyOf": [ + { + "$ref": "#/$defs/Checkpoint" + }, + { + "type": "null" + } + ], + "default": null + }, + "lora_rank": { + "anyOf": [ + { + "type": "integer" + }, + { + "type": "null" + } + ], + "default": null, + "title": "Lora Rank" + }, + "model_owner": { + "title": "Model Owner", + "type": "string" + }, + "training_run_id": { + "title": "Training Run Id", + "type": "string" + }, + "user_metadata": { + "anyOf": [ + { + "additionalProperties": { + "type": "string" + }, + "type": "object" + }, + { + "type": "null" + } + ], + "default": null, + "title": "User Metadata" + } + }, + "required": [ + "training_run_id", + "base_model", + "model_owner", + "is_lora", + "last_request_time" + ], + "title": "TrainingRun", + "type": "object" + } + }, + "properties": { + "cursor": { + "$ref": "#/$defs/Cursor" + }, + "training_runs": { + "items": { + "$ref": "#/$defs/TrainingRun" + }, + "title": "Training Runs", + "type": "array" + } + }, + "required": [ + "training_runs", + "cursor" + ], + "title": "TrainingRunsResponse", + "type": "object" + }, + "responses": {}, + "statusCode": 200 } }, "/training_runs/{run_id}": { "GET": { - "operationId": "get_training_run_training_runs__run_id__get", - "parameters": [ + "body": [], + "cookies": [], + "headers": [], + "operationId": "get_training_run", + "path": [ { - "in": "path", "name": "run_id", "required": true, "schema": { - "title": "Run Id", "type": "string" } } ], - "responses": [ - "200", - "422" - ] + "query": [], + "response": { + "$defs": { + "Checkpoint": { + "properties": { + "checkpoint_id": { + "title": "Checkpoint Id", + "type": "string" + }, + "checkpoint_type": { + "enum": [ + "training", + "sampler" + ], + "title": "Checkpoint Type", + "type": "string" + }, + "expires_at": { + "anyOf": [ + { + "format": "date-time", + "type": "string" + }, + { + "type": "null" + } + ], + "default": null, + "title": "Expires At" + }, + "public": { + "default": false, + "title": "Public", + "type": "boolean" + }, + "size_bytes": { + "anyOf": [ + { + "type": "integer" + }, + { + "type": "null" + } + ], + "default": null, + "title": "Size Bytes" + }, + "time": { + "format": "date-time", + "title": "Time", + "type": "string" + }, + "tinker_path": { + "title": "Tinker Path", + "type": "string" + } + }, + "required": [ + "checkpoint_id", + "checkpoint_type", + "time", + "tinker_path" + ], + "title": "Checkpoint", + "type": "object" + } + }, + "properties": { + "base_model": { + "title": "Base Model", + "type": "string" + }, + "corrupted": { + "default": false, + "title": "Corrupted", + "type": "boolean" + }, + "is_lora": { + "title": "Is Lora", + "type": "boolean" + }, + "last_checkpoint": { + "anyOf": [ + { + "$ref": "#/$defs/Checkpoint" + }, + { + "type": "null" + } + ], + "default": null + }, + "last_request_time": { + "format": "date-time", + "title": "Last Request Time", + "type": "string" + }, + "last_sampler_checkpoint": { + "anyOf": [ + { + "$ref": "#/$defs/Checkpoint" + }, + { + "type": "null" + } + ], + "default": null + }, + "lora_rank": { + "anyOf": [ + { + "type": "integer" + }, + { + "type": "null" + } + ], + "default": null, + "title": "Lora Rank" + }, + "model_owner": { + "title": "Model Owner", + "type": "string" + }, + "training_run_id": { + "title": "Training Run Id", + "type": "string" + }, + "user_metadata": { + "anyOf": [ + { + "additionalProperties": { + "type": "string" + }, + "type": "object" + }, + { + "type": "null" + } + ], + "default": null, + "title": "User Metadata" + } + }, + "required": [ + "training_run_id", + "base_model", + "model_owner", + "is_lora", + "last_request_time" + ], + "title": "TrainingRun", + "type": "object" + }, + "responses": {}, + "statusCode": 200 } }, "/training_runs/{run_id}/checkpoints": { "GET": { - "operationId": "get_run_checkpoints_training_runs__run_id__checkpoints_get", - "parameters": [ + "body": [], + "cookies": [], + "headers": [], + "operationId": "get_run_checkpoints", + "path": [ { - "in": "path", "name": "run_id", "required": true, "schema": { - "title": "Run Id", "type": "string" } } ], - "responses": [ - "200", - "422" - ] + "query": [], + "response": { + "$defs": { + "Checkpoint": { + "properties": { + "checkpoint_id": { + "title": "Checkpoint Id", + "type": "string" + }, + "checkpoint_type": { + "enum": [ + "training", + "sampler" + ], + "title": "Checkpoint Type", + "type": "string" + }, + "expires_at": { + "anyOf": [ + { + "format": "date-time", + "type": "string" + }, + { + "type": "null" + } + ], + "default": null, + "title": "Expires At" + }, + "public": { + "default": false, + "title": "Public", + "type": "boolean" + }, + "size_bytes": { + "anyOf": [ + { + "type": "integer" + }, + { + "type": "null" + } + ], + "default": null, + "title": "Size Bytes" + }, + "time": { + "format": "date-time", + "title": "Time", + "type": "string" + }, + "tinker_path": { + "title": "Tinker Path", + "type": "string" + } + }, + "required": [ + "checkpoint_id", + "checkpoint_type", + "time", + "tinker_path" + ], + "title": "Checkpoint", + "type": "object" + }, + "Cursor": { + "properties": { + "limit": { + "title": "Limit", + "type": "integer" + }, + "offset": { + "title": "Offset", + "type": "integer" + }, + "total_count": { + "title": "Total Count", + "type": "integer" + } + }, + "required": [ + "offset", + "limit", + "total_count" + ], + "title": "Cursor", + "type": "object" + } + }, + "properties": { + "checkpoints": { + "items": { + "$ref": "#/$defs/Checkpoint" + }, + "title": "Checkpoints", + "type": "array" + }, + "cursor": { + "anyOf": [ + { + "$ref": "#/$defs/Cursor" + }, + { + "type": "null" + } + ], + "default": null + } + }, + "required": [ + "checkpoints" + ], + "title": "CheckpointsListResponse", + "type": "object" + }, + "responses": {}, + "statusCode": 200 } }, "/training_runs/{run_id}/checkpoints/{checkpoint_id}": { "DELETE": { - "operationId": "delete_run_checkpoint_training_runs__run_id__checkpoints__checkpoint_id__delete", - "parameters": [ + "body": [], + "cookies": [], + "headers": [], + "operationId": "delete_run_checkpoint", + "path": [ { - "in": "path", "name": "run_id", "required": true, "schema": { - "title": "Run Id", "type": "string" } }, { - "in": "path", "name": "checkpoint_id", "required": true, "schema": { - "title": "Checkpoint Id", "type": "string" } } ], - "responses": [ - "200", - "422" - ] + "query": [], + "response": {}, + "responses": {}, + "statusCode": 200 } }, "/training_runs/{run_id}/checkpoints/{checkpoint_id}/publish": { "POST": { - "operationId": "publish_checkpoint_training_runs__run_id__checkpoints__checkpoint_id__publish_post", - "parameters": [ + "body": [], + "cookies": [], + "headers": [], + "operationId": "publish_checkpoint", + "path": [ { - "in": "path", "name": "run_id", "required": true, "schema": { - "title": "Run Id", "type": "string" } }, { - "in": "path", "name": "checkpoint_id", "required": true, "schema": { - "title": "Checkpoint Id", "type": "string" } } ], - "responses": [ - "200", - "422" - ] + "query": [], + "response": { + "type": "null" + }, + "responses": {}, + "statusCode": 200 } }, "/twinkle/capacity_info": { "GET": { - "operationId": "get_capacity_info_twinkle_capacity_info_get", - "parameters": [], - "responses": [ - "200" - ] + "body": [], + "cookies": [], + "headers": [], + "operationId": "get_capacity_info", + "path": [], + "query": [], + "response": { + "description": "Response body for the /capacity_info endpoint.", + "properties": { + "free_loras": { + "title": "Free Loras", + "type": "integer" + }, + "max_loras": { + "title": "Max Loras", + "type": "integer" + }, + "used_loras": { + "title": "Used Loras", + "type": "integer" + } + }, + "required": [ + "max_loras", + "used_loras", + "free_loras" + ], + "title": "CapacityInfoResponse", + "type": "object" + }, + "responses": {}, + "statusCode": 200 } }, "/twinkle/checkpoint_path/{run_id}/{checkpoint_id}": { "GET": { - "operationId": "get_checkpoint_path_twinkle_checkpoint_path__run_id___checkpoint_id__get", - "parameters": [ + "body": [], + "cookies": [], + "headers": [], + "operationId": "get_checkpoint_path", + "path": [ { - "in": "path", "name": "run_id", "required": true, "schema": { - "title": "Run Id", "type": "string" } }, { - "in": "path", "name": "checkpoint_id", "required": true, "schema": { - "title": "Checkpoint Id", "type": "string" } } ], - "responses": [ - "200", - "422" - ] + "query": [], + "response": { + "description": "Response body for the /checkpoint_path endpoint.", + "properties": { + "path": { + "title": "Path", + "type": "string" + }, + "twinkle_path": { + "title": "Twinkle Path", + "type": "string" + } + }, + "required": [ + "path", + "twinkle_path" + ], + "title": "CheckpointPathResponse", + "type": "object" + }, + "responses": {}, + "statusCode": 200 } }, "/twinkle/create_session": { "POST": { - "operationId": "create_session_twinkle_create_session_post", - "parameters": [], - "responses": [ - "200", - "422" - ] + "body": [ + { + "name": "body", + "required": true, + "schema": { + "description": "Request body for POST /twinkle/create_session.", + "properties": { + "metadata": { + "anyOf": [ + { + "additionalProperties": true, + "type": "object" + }, + { + "type": "null" + } + ], + "default": null, + "title": "Metadata" + } + }, + "title": "CreateSessionRequest", + "type": "object" + } + } + ], + "cookies": [], + "headers": [], + "operationId": "create_session", + "path": [], + "query": [], + "response": { + "description": "Response body for POST /twinkle/create_session.", + "properties": { + "session_id": { + "title": "Session Id", + "type": "string" + } + }, + "required": [ + "session_id" + ], + "title": "CreateSessionResponse", + "type": "object" + }, + "responses": {}, + "statusCode": 200 } }, "/twinkle/get_server_capabilities": { "GET": { - "operationId": "get_server_capabilities_twinkle_get_server_capabilities_get", - "parameters": [], - "responses": [ - "200" - ] + "body": [], + "cookies": [], + "headers": [], + "operationId": "get_server_capabilities", + "path": [], + "query": [], + "response": { + "$defs": { + "SupportedModel": { + "description": "Information about a supported model.", + "properties": { + "model_name": { + "title": "Model Name", + "type": "string" + } + }, + "required": [ + "model_name" + ], + "title": "SupportedModel", + "type": "object" + } + }, + "description": "Response body for the /get_server_capabilities endpoint.", + "properties": { + "supported_models": { + "items": { + "$ref": "#/$defs/SupportedModel" + }, + "title": "Supported Models", + "type": "array" + } + }, + "required": [ + "supported_models" + ], + "title": "GetServerCapabilitiesResponse", + "type": "object" + }, + "responses": {}, + "statusCode": 200 } }, "/twinkle/healthz": { "GET": { - "operationId": "healthz_twinkle_healthz_get", - "parameters": [], - "responses": [ - "200" - ] + "body": [], + "cookies": [], + "headers": [], + "operationId": "healthz", + "path": [], + "query": [], + "response": { + "properties": { + "status": { + "title": "Status", + "type": "string" + } + }, + "required": [ + "status" + ], + "title": "HealthResponse", + "type": "object" + }, + "responses": {}, + "statusCode": 200 } }, "/twinkle/healthz/deep": { "GET": { - "operationId": "healthz_deep_twinkle_healthz_deep_get", - "parameters": [], - "responses": [ - "200" - ] + "body": [], + "cookies": [], + "headers": [], + "operationId": "healthz_deep", + "path": [], + "query": [], + "response": { + "pythonType": "dict" + }, + "responses": {}, + "statusCode": 200 } }, "/twinkle/session_heartbeat": { "POST": { - "operationId": "session_heartbeat_twinkle_session_heartbeat_post", - "parameters": [], - "responses": [ - "200", - "422" - ] + "body": [ + { + "name": "body", + "required": true, + "schema": { + "description": "Request body for POST /twinkle/session_heartbeat.", + "properties": { + "session_id": { + "title": "Session Id", + "type": "string" + } + }, + "required": [ + "session_id" + ], + "title": "SessionHeartbeatRequest", + "type": "object" + } + } + ], + "cookies": [], + "headers": [], + "operationId": "session_heartbeat", + "path": [], + "query": [], + "response": { + "description": "Response body for POST /twinkle/session_heartbeat.", + "properties": {}, + "title": "SessionHeartbeatResponse", + "type": "object" + }, + "responses": {}, + "statusCode": 200 } }, "/twinkle/status": { "GET": { - "operationId": "status_twinkle_status_get", - "parameters": [], - "responses": [ - "200" - ] + "body": [], + "cookies": [], + "headers": [], + "operationId": "status", + "path": [], + "query": [], + "response": { + "pythonType": "dict" + }, + "responses": {}, + "statusCode": 200 } }, "/twinkle/training_runs": { "GET": { - "operationId": "get_training_runs_twinkle_training_runs_get", - "parameters": [ + "body": [], + "cookies": [], + "headers": [], + "operationId": "get_training_runs", + "path": [], + "query": [ { - "in": "query", "name": "limit", "required": false, "schema": { - "default": 20, - "title": "Limit", "type": "integer" } }, { - "in": "query", "name": "offset", "required": false, "schema": { - "default": 0, - "title": "Offset", "type": "integer" } } ], - "responses": [ - "200", - "422" - ] + "response": { + "$defs": { + "Cursor": { + "properties": { + "limit": { + "title": "Limit", + "type": "integer" + }, + "offset": { + "title": "Offset", + "type": "integer" + }, + "total_count": { + "title": "Total Count", + "type": "integer" + } + }, + "required": [ + "limit", + "offset", + "total_count" + ], + "title": "Cursor", + "type": "object" + }, + "TrainingRun": { + "description": "Twinkle training run model.", + "properties": { + "base_model": { + "title": "Base Model", + "type": "string" + }, + "corrupted": { + "default": false, + "title": "Corrupted", + "type": "boolean" + }, + "is_lora": { + "default": false, + "title": "Is Lora", + "type": "boolean" + }, + "last_checkpoint": { + "anyOf": [ + { + "additionalProperties": true, + "type": "object" + }, + { + "type": "null" + } + ], + "default": null, + "title": "Last Checkpoint" + }, + "last_request_time": { + "anyOf": [ + { + "format": "date-time", + "type": "string" + }, + { + "type": "null" + } + ], + "default": null, + "title": "Last Request Time" + }, + "last_sampler_checkpoint": { + "anyOf": [ + { + "additionalProperties": true, + "type": "object" + }, + { + "type": "null" + } + ], + "default": null, + "title": "Last Sampler Checkpoint" + }, + "lora_rank": { + "anyOf": [ + { + "type": "integer" + }, + { + "type": "null" + } + ], + "default": null, + "title": "Lora Rank" + }, + "model_owner": { + "title": "Model Owner", + "type": "string" + }, + "save_dir": { + "anyOf": [ + { + "type": "string" + }, + { + "type": "null" + } + ], + "default": null, + "title": "Save Dir" + }, + "training_run_id": { + "title": "Training Run Id", + "type": "string" + }, + "user_metadata": { + "anyOf": [ + { + "additionalProperties": true, + "type": "object" + }, + { + "type": "null" + } + ], + "default": null, + "title": "User Metadata" + } + }, + "required": [ + "training_run_id", + "base_model", + "model_owner" + ], + "title": "TrainingRun", + "type": "object" + } + }, + "properties": { + "cursor": { + "$ref": "#/$defs/Cursor" + }, + "training_runs": { + "items": { + "$ref": "#/$defs/TrainingRun" + }, + "title": "Training Runs", + "type": "array" + } + }, + "required": [ + "training_runs", + "cursor" + ], + "title": "TrainingRunsResponse", + "type": "object" + }, + "responses": {}, + "statusCode": 200 } }, "/twinkle/training_runs/{run_id}": { "GET": { - "operationId": "get_training_run_twinkle_training_runs__run_id__get", - "parameters": [ + "body": [], + "cookies": [], + "headers": [], + "operationId": "get_training_run", + "path": [ { - "in": "path", "name": "run_id", "required": true, "schema": { - "title": "Run Id", "type": "string" } } ], - "responses": [ - "200", - "422" - ] + "query": [], + "response": { + "description": "Twinkle training run model.", + "properties": { + "base_model": { + "title": "Base Model", + "type": "string" + }, + "corrupted": { + "default": false, + "title": "Corrupted", + "type": "boolean" + }, + "is_lora": { + "default": false, + "title": "Is Lora", + "type": "boolean" + }, + "last_checkpoint": { + "anyOf": [ + { + "additionalProperties": true, + "type": "object" + }, + { + "type": "null" + } + ], + "default": null, + "title": "Last Checkpoint" + }, + "last_request_time": { + "anyOf": [ + { + "format": "date-time", + "type": "string" + }, + { + "type": "null" + } + ], + "default": null, + "title": "Last Request Time" + }, + "last_sampler_checkpoint": { + "anyOf": [ + { + "additionalProperties": true, + "type": "object" + }, + { + "type": "null" + } + ], + "default": null, + "title": "Last Sampler Checkpoint" + }, + "lora_rank": { + "anyOf": [ + { + "type": "integer" + }, + { + "type": "null" + } + ], + "default": null, + "title": "Lora Rank" + }, + "model_owner": { + "title": "Model Owner", + "type": "string" + }, + "save_dir": { + "anyOf": [ + { + "type": "string" + }, + { + "type": "null" + } + ], + "default": null, + "title": "Save Dir" + }, + "training_run_id": { + "title": "Training Run Id", + "type": "string" + }, + "user_metadata": { + "anyOf": [ + { + "additionalProperties": true, + "type": "object" + }, + { + "type": "null" + } + ], + "default": null, + "title": "User Metadata" + } + }, + "required": [ + "training_run_id", + "base_model", + "model_owner" + ], + "title": "TrainingRun", + "type": "object" + }, + "responses": {}, + "statusCode": 200 } }, "/twinkle/training_runs/{run_id}/checkpoints": { "GET": { - "operationId": "get_run_checkpoints_twinkle_training_runs__run_id__checkpoints_get", - "parameters": [ + "body": [], + "cookies": [], + "headers": [], + "operationId": "get_run_checkpoints", + "path": [ { - "in": "path", "name": "run_id", "required": true, "schema": { - "title": "Run Id", "type": "string" } } ], - "responses": [ - "200", - "422" - ] + "query": [], + "response": { + "$defs": { + "Checkpoint": { + "description": "Twinkle checkpoint model.", + "properties": { + "base_model": { + "anyOf": [ + { + "type": "string" + }, + { + "type": "null" + } + ], + "default": null, + "title": "Base Model" + }, + "checkpoint_id": { + "title": "Checkpoint Id", + "type": "string" + }, + "checkpoint_type": { + "title": "Checkpoint Type", + "type": "string" + }, + "is_lora": { + "default": false, + "title": "Is Lora", + "type": "boolean" + }, + "lora_rank": { + "anyOf": [ + { + "type": "integer" + }, + { + "type": "null" + } + ], + "default": null, + "title": "Lora Rank" + }, + "public": { + "default": false, + "title": "Public", + "type": "boolean" + }, + "size_bytes": { + "title": "Size Bytes", + "type": "integer" + }, + "time": { + "format": "date-time", + "title": "Time", + "type": "string" + }, + "train_attn": { + "anyOf": [ + { + "type": "boolean" + }, + { + "type": "null" + } + ], + "default": null, + "title": "Train Attn" + }, + "train_mlp": { + "anyOf": [ + { + "type": "boolean" + }, + { + "type": "null" + } + ], + "default": null, + "title": "Train Mlp" + }, + "train_unembed": { + "anyOf": [ + { + "type": "boolean" + }, + { + "type": "null" + } + ], + "default": null, + "title": "Train Unembed" + }, + "twinkle_path": { + "title": "Twinkle Path", + "type": "string" + }, + "user_metadata": { + "anyOf": [ + { + "additionalProperties": true, + "type": "object" + }, + { + "type": "null" + } + ], + "default": null, + "title": "User Metadata" + } + }, + "required": [ + "checkpoint_id", + "checkpoint_type", + "time", + "size_bytes", + "twinkle_path" + ], + "title": "Checkpoint", + "type": "object" + }, + "Cursor": { + "properties": { + "limit": { + "title": "Limit", + "type": "integer" + }, + "offset": { + "title": "Offset", + "type": "integer" + }, + "total_count": { + "title": "Total Count", + "type": "integer" + } + }, + "required": [ + "limit", + "offset", + "total_count" + ], + "title": "Cursor", + "type": "object" + } + }, + "properties": { + "checkpoints": { + "items": { + "$ref": "#/$defs/Checkpoint" + }, + "title": "Checkpoints", + "type": "array" + }, + "cursor": { + "anyOf": [ + { + "$ref": "#/$defs/Cursor" + }, + { + "type": "null" + } + ], + "default": null + } + }, + "required": [ + "checkpoints" + ], + "title": "CheckpointsListResponse", + "type": "object" + }, + "responses": {}, + "statusCode": 200 } }, "/twinkle/training_runs/{run_id}/checkpoints/{checkpoint_id}": { "DELETE": { - "operationId": "delete_run_checkpoint_twinkle_training_runs__run_id__checkpoints__checkpoint_id__delete", - "parameters": [ + "body": [], + "cookies": [], + "headers": [], + "operationId": "delete_run_checkpoint", + "path": [ { - "in": "path", "name": "run_id", "required": true, "schema": { - "title": "Run Id", "type": "string" } }, { - "in": "path", "name": "checkpoint_id", "required": true, "schema": { - "title": "Checkpoint Id", "type": "string" } } ], - "responses": [ - "200", - "422" - ] + "query": [], + "response": { + "properties": { + "message": { + "title": "Message", + "type": "string" + }, + "success": { + "title": "Success", + "type": "boolean" + } + }, + "required": [ + "success", + "message" + ], + "title": "DeleteCheckpointResponse", + "type": "object" + }, + "responses": {}, + "statusCode": 200 } }, "/twinkle/weights_info": { "POST": { - "operationId": "weights_info_twinkle_weights_info_post", - "parameters": [], - "responses": [ - "200", - "422" - ] + "body": [ + { + "name": "body", + "required": true, + "schema": { + "properties": { + "twinkle_path": { + "title": "Twinkle Path", + "type": "string" + } + }, + "required": [ + "twinkle_path" + ], + "title": "WeightsInfoRequest", + "type": "object" + } + } + ], + "cookies": [], + "headers": [], + "operationId": "weights_info", + "path": [], + "query": [], + "response": { + "description": "Twinkle weights info response.", + "properties": { + "base_model": { + "title": "Base Model", + "type": "string" + }, + "is_lora": { + "default": false, + "title": "Is Lora", + "type": "boolean" + }, + "lora_rank": { + "anyOf": [ + { + "type": "integer" + }, + { + "type": "null" + } + ], + "default": null, + "title": "Lora Rank" + }, + "model_owner": { + "title": "Model Owner", + "type": "string" + }, + "training_run_id": { + "title": "Training Run Id", + "type": "string" + } + }, + "required": [ + "training_run_id", + "base_model", + "model_owner" + ], + "title": "WeightsInfoResponse", + "type": "object" + }, + "responses": {}, + "statusCode": 200 } }, "/unload_model": { "POST": { - "operationId": "unload_model_unload_model_post", - "parameters": [], - "responses": [ - "200", - "422" - ] + "body": [ + { + "name": "body", + "required": true, + "schema": { + "additionalProperties": false, + "properties": { + "model_id": { + "title": "Model Id", + "type": "string" + }, + "type": { + "const": "unload_model", + "default": "unload_model", + "title": "Type", + "type": "string" + } + }, + "required": [ + "model_id" + ], + "title": "UnloadModelRequest", + "type": "object" + } + } + ], + "cookies": [], + "headers": [], + "operationId": "unload_model", + "path": [], + "query": [], + "response": {}, + "responses": {}, + "statusCode": 200 } }, "/weights_info": { "POST": { - "operationId": "weights_info_weights_info_post", - "parameters": [], - "responses": [ - "200", - "422" - ] + "body": [ + { + "name": "body", + "required": true, + "schema": { + "additionalProperties": {}, + "type": "object" + } + } + ], + "cookies": [], + "headers": [], + "operationId": "weights_info", + "path": [], + "query": [], + "response": { + "description": "Minimal information for loading public checkpoints.", + "properties": { + "base_model": { + "title": "Base Model", + "type": "string" + }, + "is_lora": { + "title": "Is Lora", + "type": "boolean" + }, + "lora_rank": { + "anyOf": [ + { + "type": "integer" + }, + { + "type": "null" + } + ], + "default": null, + "title": "Lora Rank" + }, + "train_attn": { + "anyOf": [ + { + "type": "boolean" + }, + { + "type": "null" + } + ], + "default": null, + "title": "Train Attn" + }, + "train_mlp": { + "anyOf": [ + { + "type": "boolean" + }, + { + "type": "null" + } + ], + "default": null, + "title": "Train Mlp" + }, + "train_unembed": { + "anyOf": [ + { + "type": "boolean" + }, + { + "type": "null" + } + ], + "default": null, + "title": "Train Unembed" + } + }, + "required": [ + "base_model", + "is_lora" + ], + "title": "WeightsInfoResponse", + "type": "object" + }, + "responses": {}, + "statusCode": 200 } } } @@ -580,421 +4103,3070 @@ "paths": { "/healthz": { "GET": { - "operationId": "healthz_healthz_get", - "parameters": [], - "responses": [ - "200" - ] + "body": [], + "cookies": [], + "headers": [], + "operationId": "model_healthz", + "path": [], + "query": [], + "response": { + "pythonType": "dict" + }, + "responses": {}, + "statusCode": 200 } }, "/tinker/create_model": { "POST": { - "operationId": "create_model_tinker_create_model_post", - "parameters": [], - "responses": [ - "200", - "422" - ] + "body": [ + { + "name": "body", + "required": true, + "schema": { + "$defs": { + "LoraConfig": { + "additionalProperties": false, + "properties": { + "rank": { + "title": "Rank", + "type": "integer" + }, + "seed": { + "anyOf": [ + { + "type": "integer" + }, + { + "type": "null" + } + ], + "default": null, + "title": "Seed" + }, + "train_attn": { + "default": true, + "title": "Train Attn", + "type": "boolean" + }, + "train_mlp": { + "default": true, + "title": "Train Mlp", + "type": "boolean" + }, + "train_unembed": { + "default": true, + "title": "Train Unembed", + "type": "boolean" + } + }, + "required": [ + "rank" + ], + "title": "LoraConfig", + "type": "object" + } + }, + "additionalProperties": false, + "properties": { + "base_model": { + "title": "Base Model", + "type": "string" + }, + "lora_config": { + "anyOf": [ + { + "$ref": "#/$defs/LoraConfig" + }, + { + "type": "null" + } + ], + "default": null + }, + "model_seq_id": { + "title": "Model Seq Id", + "type": "integer" + }, + "session_id": { + "title": "Session Id", + "type": "string" + }, + "type": { + "const": "create_model", + "default": "create_model", + "title": "Type", + "type": "string" + }, + "user_metadata": { + "anyOf": [ + { + "additionalProperties": true, + "type": "object" + }, + { + "type": "null" + } + ], + "default": null, + "title": "User Metadata" + } + }, + "required": [ + "session_id", + "model_seq_id", + "base_model" + ], + "title": "CreateModelRequest", + "type": "object" + } + } + ], + "cookies": [], + "headers": [], + "operationId": "create_model", + "path": [], + "query": [], + "response": { + "properties": { + "model_id": { + "anyOf": [ + { + "type": "string" + }, + { + "type": "null" + } + ], + "default": null, + "title": "Model Id" + }, + "request_id": { + "title": "Request Id", + "type": "string" + } + }, + "required": [ + "request_id" + ], + "title": "UntypedAPIFuture", + "type": "object" + }, + "responses": {}, + "statusCode": 200 } }, "/tinker/forward": { "POST": { - "operationId": "forward_tinker_forward_post", - "parameters": [], - "responses": [ - "200", - "422" - ] - } - }, - "/tinker/forward_backward": { - "POST": { - "operationId": "forward_backward_tinker_forward_backward_post", - "parameters": [], - "responses": [ - "200", - "422" - ] - } - }, - "/tinker/get_info": { - "POST": { - "operationId": "get_info_tinker_get_info_post", - "parameters": [], - "responses": [ - "200", - "422" - ] - } + "body": [ + { + "name": "body", + "required": true, + "schema": { + "$defs": { + "Datum": { + "additionalProperties": false, + "properties": { + "loss_fn_inputs": { + "additionalProperties": { + "$ref": "#/$defs/TensorData" + }, + "title": "Loss Fn Inputs", + "type": "object" + }, + "model_input": { + "$ref": "#/$defs/ModelInput" + } + }, + "required": [ + "loss_fn_inputs", + "model_input" + ], + "title": "Datum", + "type": "object" + }, + "EncodedTextChunk": { + "additionalProperties": false, + "properties": { + "tokens": { + "items": { + "type": "integer" + }, + "title": "Tokens", + "type": "array" + }, + "type": { + "const": "encoded_text", + "default": "encoded_text", + "title": "Type", + "type": "string" + } + }, + "required": [ + "tokens" + ], + "title": "EncodedTextChunk", + "type": "object" + }, + "ForwardBackwardInput": { + "additionalProperties": false, + "properties": { + "data": { + "items": { + "$ref": "#/$defs/Datum" + }, + "title": "Data", + "type": "array" + }, + "loss_fn": { + "enum": [ + "cross_entropy", + "importance_sampling", + "ppo", + "cispo", + "dro" + ], + "title": "Loss Fn", + "type": "string" + }, + "loss_fn_config": { + "anyOf": [ + { + "additionalProperties": { + "type": "number" + }, + "type": "object" + }, + { + "type": "null" + } + ], + "default": null, + "title": "Loss Fn Config" + } + }, + "required": [ + "data", + "loss_fn" + ], + "title": "ForwardBackwardInput", + "type": "object" + }, + "ImageAssetPointerChunk": { + "additionalProperties": false, + "properties": { + "expected_tokens": { + "anyOf": [ + { + "type": "integer" + }, + { + "type": "null" + } + ], + "default": null, + "title": "Expected Tokens" + }, + "format": { + "enum": [ + "png", + "jpeg" + ], + "title": "Format", + "type": "string" + }, + "location": { + "title": "Location", + "type": "string" + }, + "type": { + "const": "image_asset_pointer", + "default": "image_asset_pointer", + "title": "Type", + "type": "string" + } + }, + "required": [ + "format", + "location" + ], + "title": "ImageAssetPointerChunk", + "type": "object" + }, + "ImageChunk": { + "additionalProperties": false, + "properties": { + "data": { + "format": "binary", + "title": "Data", + "type": "string" + }, + "expected_tokens": { + "anyOf": [ + { + "type": "integer" + }, + { + "type": "null" + } + ], + "default": null, + "title": "Expected Tokens" + }, + "format": { + "enum": [ + "png", + "jpeg" + ], + "title": "Format", + "type": "string" + }, + "type": { + "const": "image", + "default": "image", + "title": "Type", + "type": "string" + } + }, + "required": [ + "data", + "format" + ], + "title": "ImageChunk", + "type": "object" + }, + "ModelInput": { + "additionalProperties": false, + "properties": { + "chunks": { + "items": { + "anyOf": [ + { + "$ref": "#/$defs/EncodedTextChunk" + }, + { + "$ref": "#/$defs/ImageAssetPointerChunk" + }, + { + "$ref": "#/$defs/ImageChunk" + } + ] + }, + "title": "Chunks", + "type": "array" + } + }, + "required": [ + "chunks" + ], + "title": "ModelInput", + "type": "object" + }, + "TensorData": { + "additionalProperties": false, + "properties": { + "data": { + "anyOf": [ + { + "items": { + "type": "integer" + }, + "type": "array" + }, + { + "items": { + "type": "number" + }, + "type": "array" + } + ], + "title": "Data" + }, + "dtype": { + "enum": [ + "int64", + "float32" + ], + "title": "Dtype", + "type": "string" + }, + "shape": { + "anyOf": [ + { + "items": { + "type": "integer" + }, + "type": "array" + }, + { + "type": "null" + } + ], + "default": null, + "title": "Shape" + } + }, + "required": [ + "data", + "dtype" + ], + "title": "TensorData", + "type": "object" + } + }, + "additionalProperties": false, + "properties": { + "forward_input": { + "$ref": "#/$defs/ForwardBackwardInput" + }, + "model_id": { + "title": "Model Id", + "type": "string" + }, + "seq_id": { + "anyOf": [ + { + "type": "integer" + }, + { + "type": "null" + } + ], + "default": null, + "title": "Seq Id" + } + }, + "required": [ + "forward_input", + "model_id" + ], + "title": "ForwardRequest", + "type": "object" + } + } + ], + "cookies": [], + "headers": [], + "operationId": "forward", + "path": [], + "query": [], + "response": { + "properties": { + "model_id": { + "anyOf": [ + { + "type": "string" + }, + { + "type": "null" + } + ], + "default": null, + "title": "Model Id" + }, + "request_id": { + "title": "Request Id", + "type": "string" + } + }, + "required": [ + "request_id" + ], + "title": "UntypedAPIFuture", + "type": "object" + }, + "responses": {}, + "statusCode": 200 + } + }, + "/tinker/forward_backward": { + "POST": { + "body": [ + { + "name": "body", + "required": true, + "schema": { + "$defs": { + "Datum": { + "additionalProperties": false, + "properties": { + "loss_fn_inputs": { + "additionalProperties": { + "$ref": "#/$defs/TensorData" + }, + "title": "Loss Fn Inputs", + "type": "object" + }, + "model_input": { + "$ref": "#/$defs/ModelInput" + } + }, + "required": [ + "loss_fn_inputs", + "model_input" + ], + "title": "Datum", + "type": "object" + }, + "EncodedTextChunk": { + "additionalProperties": false, + "properties": { + "tokens": { + "items": { + "type": "integer" + }, + "title": "Tokens", + "type": "array" + }, + "type": { + "const": "encoded_text", + "default": "encoded_text", + "title": "Type", + "type": "string" + } + }, + "required": [ + "tokens" + ], + "title": "EncodedTextChunk", + "type": "object" + }, + "ForwardBackwardInput": { + "additionalProperties": false, + "properties": { + "data": { + "items": { + "$ref": "#/$defs/Datum" + }, + "title": "Data", + "type": "array" + }, + "loss_fn": { + "enum": [ + "cross_entropy", + "importance_sampling", + "ppo", + "cispo", + "dro" + ], + "title": "Loss Fn", + "type": "string" + }, + "loss_fn_config": { + "anyOf": [ + { + "additionalProperties": { + "type": "number" + }, + "type": "object" + }, + { + "type": "null" + } + ], + "default": null, + "title": "Loss Fn Config" + } + }, + "required": [ + "data", + "loss_fn" + ], + "title": "ForwardBackwardInput", + "type": "object" + }, + "ImageAssetPointerChunk": { + "additionalProperties": false, + "properties": { + "expected_tokens": { + "anyOf": [ + { + "type": "integer" + }, + { + "type": "null" + } + ], + "default": null, + "title": "Expected Tokens" + }, + "format": { + "enum": [ + "png", + "jpeg" + ], + "title": "Format", + "type": "string" + }, + "location": { + "title": "Location", + "type": "string" + }, + "type": { + "const": "image_asset_pointer", + "default": "image_asset_pointer", + "title": "Type", + "type": "string" + } + }, + "required": [ + "format", + "location" + ], + "title": "ImageAssetPointerChunk", + "type": "object" + }, + "ImageChunk": { + "additionalProperties": false, + "properties": { + "data": { + "format": "binary", + "title": "Data", + "type": "string" + }, + "expected_tokens": { + "anyOf": [ + { + "type": "integer" + }, + { + "type": "null" + } + ], + "default": null, + "title": "Expected Tokens" + }, + "format": { + "enum": [ + "png", + "jpeg" + ], + "title": "Format", + "type": "string" + }, + "type": { + "const": "image", + "default": "image", + "title": "Type", + "type": "string" + } + }, + "required": [ + "data", + "format" + ], + "title": "ImageChunk", + "type": "object" + }, + "ModelInput": { + "additionalProperties": false, + "properties": { + "chunks": { + "items": { + "anyOf": [ + { + "$ref": "#/$defs/EncodedTextChunk" + }, + { + "$ref": "#/$defs/ImageAssetPointerChunk" + }, + { + "$ref": "#/$defs/ImageChunk" + } + ] + }, + "title": "Chunks", + "type": "array" + } + }, + "required": [ + "chunks" + ], + "title": "ModelInput", + "type": "object" + }, + "TensorData": { + "additionalProperties": false, + "properties": { + "data": { + "anyOf": [ + { + "items": { + "type": "integer" + }, + "type": "array" + }, + { + "items": { + "type": "number" + }, + "type": "array" + } + ], + "title": "Data" + }, + "dtype": { + "enum": [ + "int64", + "float32" + ], + "title": "Dtype", + "type": "string" + }, + "shape": { + "anyOf": [ + { + "items": { + "type": "integer" + }, + "type": "array" + }, + { + "type": "null" + } + ], + "default": null, + "title": "Shape" + } + }, + "required": [ + "data", + "dtype" + ], + "title": "TensorData", + "type": "object" + } + }, + "additionalProperties": false, + "properties": { + "forward_backward_input": { + "$ref": "#/$defs/ForwardBackwardInput" + }, + "model_id": { + "title": "Model Id", + "type": "string" + }, + "seq_id": { + "anyOf": [ + { + "type": "integer" + }, + { + "type": "null" + } + ], + "default": null, + "title": "Seq Id" + } + }, + "required": [ + "forward_backward_input", + "model_id" + ], + "title": "ForwardBackwardRequest", + "type": "object" + } + } + ], + "cookies": [], + "headers": [], + "operationId": "forward_backward", + "path": [], + "query": [], + "response": { + "properties": { + "model_id": { + "anyOf": [ + { + "type": "string" + }, + { + "type": "null" + } + ], + "default": null, + "title": "Model Id" + }, + "request_id": { + "title": "Request Id", + "type": "string" + } + }, + "required": [ + "request_id" + ], + "title": "UntypedAPIFuture", + "type": "object" + }, + "responses": {}, + "statusCode": 200 + } + }, + "/tinker/get_info": { + "POST": { + "body": [ + { + "name": "body", + "required": true, + "schema": { + "additionalProperties": false, + "properties": { + "model_id": { + "title": "Model Id", + "type": "string" + }, + "type": { + "const": "get_info", + "default": "get_info", + "title": "Type", + "type": "string" + } + }, + "required": [ + "model_id" + ], + "title": "GetInfoRequest", + "type": "object" + } + } + ], + "cookies": [], + "headers": [], + "operationId": "get_info", + "path": [], + "query": [], + "response": { + "$defs": { + "ModelData": { + "description": "Metadata about a model's architecture and configuration.", + "properties": { + "arch": { + "anyOf": [ + { + "type": "string" + }, + { + "type": "null" + } + ], + "default": null, + "title": "Arch" + }, + "model_name": { + "anyOf": [ + { + "type": "string" + }, + { + "type": "null" + } + ], + "default": null, + "title": "Model Name" + }, + "tokenizer_id": { + "anyOf": [ + { + "type": "string" + }, + { + "type": "null" + } + ], + "default": null, + "title": "Tokenizer Id" + } + }, + "title": "ModelData", + "type": "object" + } + }, + "description": "Response containing information about a training client's model.", + "properties": { + "is_lora": { + "anyOf": [ + { + "type": "boolean" + }, + { + "type": "null" + } + ], + "default": null, + "title": "Is Lora" + }, + "lora_rank": { + "anyOf": [ + { + "type": "integer" + }, + { + "type": "null" + } + ], + "default": null, + "title": "Lora Rank" + }, + "model_data": { + "$ref": "#/$defs/ModelData" + }, + "model_id": { + "title": "Model Id", + "type": "string" + }, + "model_name": { + "anyOf": [ + { + "type": "string" + }, + { + "type": "null" + } + ], + "default": null, + "title": "Model Name" + }, + "type": { + "anyOf": [ + { + "const": "get_info", + "type": "string" + }, + { + "type": "null" + } + ], + "default": null, + "title": "Type" + } + }, + "required": [ + "model_data", + "model_id" + ], + "title": "GetInfoResponse", + "type": "object" + }, + "responses": {}, + "statusCode": 200 + } }, "/tinker/load_weights": { "POST": { - "operationId": "load_weights_tinker_load_weights_post", - "parameters": [], - "responses": [ - "200", - "422" - ] + "body": [ + { + "name": "body", + "required": true, + "schema": { + "additionalProperties": false, + "properties": { + "model_id": { + "title": "Model Id", + "type": "string" + }, + "optimizer": { + "title": "Optimizer", + "type": "boolean" + }, + "path": { + "title": "Path", + "type": "string" + }, + "seq_id": { + "anyOf": [ + { + "type": "integer" + }, + { + "type": "null" + } + ], + "default": null, + "title": "Seq Id" + }, + "type": { + "const": "load_weights", + "default": "load_weights", + "title": "Type", + "type": "string" + } + }, + "required": [ + "model_id", + "path", + "optimizer" + ], + "title": "LoadWeightsRequest", + "type": "object" + } + } + ], + "cookies": [], + "headers": [], + "operationId": "load_weights", + "path": [], + "query": [], + "response": { + "properties": { + "model_id": { + "anyOf": [ + { + "type": "string" + }, + { + "type": "null" + } + ], + "default": null, + "title": "Model Id" + }, + "request_id": { + "title": "Request Id", + "type": "string" + } + }, + "required": [ + "request_id" + ], + "title": "UntypedAPIFuture", + "type": "object" + }, + "responses": {}, + "statusCode": 200 } }, "/tinker/optim_step": { "POST": { - "operationId": "optim_step_tinker_optim_step_post", - "parameters": [], - "responses": [ - "200", - "422" - ] + "body": [ + { + "name": "body", + "required": true, + "schema": { + "$defs": { + "AdamParams": { + "additionalProperties": false, + "properties": { + "beta1": { + "default": 0.9, + "title": "Beta1", + "type": "number" + }, + "beta2": { + "default": 0.95, + "title": "Beta2", + "type": "number" + }, + "eps": { + "default": 1e-12, + "title": "Eps", + "type": "number" + }, + "grad_clip_norm": { + "default": 0.0, + "title": "Grad Clip Norm", + "type": "number" + }, + "learning_rate": { + "default": 0.0001, + "title": "Learning Rate", + "type": "number" + }, + "weight_decay": { + "default": 0.0, + "title": "Weight Decay", + "type": "number" + } + }, + "title": "AdamParams", + "type": "object" + } + }, + "additionalProperties": false, + "properties": { + "adam_params": { + "$ref": "#/$defs/AdamParams" + }, + "model_id": { + "title": "Model Id", + "type": "string" + }, + "seq_id": { + "anyOf": [ + { + "type": "integer" + }, + { + "type": "null" + } + ], + "default": null, + "title": "Seq Id" + }, + "type": { + "const": "optim_step", + "default": "optim_step", + "title": "Type", + "type": "string" + } + }, + "required": [ + "adam_params", + "model_id" + ], + "title": "OptimStepRequest", + "type": "object" + } + } + ], + "cookies": [], + "headers": [], + "operationId": "optim_step", + "path": [], + "query": [], + "response": { + "properties": { + "model_id": { + "anyOf": [ + { + "type": "string" + }, + { + "type": "null" + } + ], + "default": null, + "title": "Model Id" + }, + "request_id": { + "title": "Request Id", + "type": "string" + } + }, + "required": [ + "request_id" + ], + "title": "UntypedAPIFuture", + "type": "object" + }, + "responses": {}, + "statusCode": 200 } }, "/tinker/save_weights": { "POST": { - "operationId": "save_weights_tinker_save_weights_post", - "parameters": [], - "responses": [ - "200", - "422" - ] + "body": [ + { + "name": "body", + "required": true, + "schema": { + "additionalProperties": false, + "properties": { + "model_id": { + "title": "Model Id", + "type": "string" + }, + "path": { + "anyOf": [ + { + "type": "string" + }, + { + "type": "null" + } + ], + "default": null, + "title": "Path" + }, + "seq_id": { + "anyOf": [ + { + "type": "integer" + }, + { + "type": "null" + } + ], + "default": null, + "title": "Seq Id" + }, + "ttl_seconds": { + "anyOf": [ + { + "type": "integer" + }, + { + "type": "null" + } + ], + "default": null, + "title": "Ttl Seconds" + }, + "type": { + "const": "save_weights", + "default": "save_weights", + "title": "Type", + "type": "string" + } + }, + "required": [ + "model_id" + ], + "title": "SaveWeightsRequest", + "type": "object" + } + } + ], + "cookies": [], + "headers": [], + "operationId": "save_weights", + "path": [], + "query": [], + "response": { + "properties": { + "model_id": { + "anyOf": [ + { + "type": "string" + }, + { + "type": "null" + } + ], + "default": null, + "title": "Model Id" + }, + "request_id": { + "title": "Request Id", + "type": "string" + } + }, + "required": [ + "request_id" + ], + "title": "UntypedAPIFuture", + "type": "object" + }, + "responses": {}, + "statusCode": 200 } }, "/tinker/save_weights_for_sampler": { "POST": { - "operationId": "save_weights_for_sampler_tinker_save_weights_for_sampler_post", - "parameters": [], - "responses": [ - "200", - "422" - ] + "body": [ + { + "name": "body", + "required": true, + "schema": { + "additionalProperties": false, + "properties": { + "model_id": { + "title": "Model Id", + "type": "string" + }, + "path": { + "anyOf": [ + { + "type": "string" + }, + { + "type": "null" + } + ], + "default": null, + "title": "Path" + }, + "sampling_session_seq_id": { + "anyOf": [ + { + "type": "integer" + }, + { + "type": "null" + } + ], + "default": null, + "title": "Sampling Session Seq Id" + }, + "seq_id": { + "anyOf": [ + { + "type": "integer" + }, + { + "type": "null" + } + ], + "default": null, + "title": "Seq Id" + }, + "ttl_seconds": { + "anyOf": [ + { + "type": "integer" + }, + { + "type": "null" + } + ], + "default": null, + "title": "Ttl Seconds" + }, + "type": { + "const": "save_weights_for_sampler", + "default": "save_weights_for_sampler", + "title": "Type", + "type": "string" + } + }, + "required": [ + "model_id" + ], + "title": "SaveWeightsForSamplerRequest", + "type": "object" + } + } + ], + "cookies": [], + "headers": [], + "operationId": "save_weights_for_sampler", + "path": [], + "query": [], + "response": { + "properties": { + "model_id": { + "anyOf": [ + { + "type": "string" + }, + { + "type": "null" + } + ], + "default": null, + "title": "Model Id" + }, + "request_id": { + "title": "Request Id", + "type": "string" + } + }, + "required": [ + "request_id" + ], + "title": "UntypedAPIFuture", + "type": "object" + }, + "responses": {}, + "statusCode": 200 } }, "/tinker/unload_model": { "POST": { - "operationId": "unload_model_tinker_unload_model_post", - "parameters": [], - "responses": [ - "200", - "422" - ] + "body": [ + { + "name": "body", + "required": true, + "schema": { + "additionalProperties": false, + "properties": { + "model_id": { + "title": "Model Id", + "type": "string" + }, + "type": { + "const": "unload_model", + "default": "unload_model", + "title": "Type", + "type": "string" + } + }, + "required": [ + "model_id" + ], + "title": "UnloadModelRequest", + "type": "object" + } + } + ], + "cookies": [], + "headers": [], + "operationId": "unload_model", + "path": [], + "query": [], + "response": { + "properties": { + "model_id": { + "anyOf": [ + { + "type": "string" + }, + { + "type": "null" + } + ], + "default": null, + "title": "Model Id" + }, + "request_id": { + "title": "Request Id", + "type": "string" + } + }, + "required": [ + "request_id" + ], + "title": "UntypedAPIFuture", + "type": "object" + }, + "responses": {}, + "statusCode": 200 } }, "/twinkle/add_adapter_to_model": { "POST": { - "operationId": "add_adapter_to_model_twinkle_add_adapter_to_model_post", - "parameters": [], - "responses": [ - "200", - "422" - ] + "body": [ + { + "name": "body", + "required": true, + "schema": { + "additionalProperties": true, + "properties": { + "adapter_name": { + "title": "Adapter Name", + "type": "string" + }, + "config": { + "anyOf": [ + { + "type": "string" + }, + { + "type": "null" + } + ], + "default": null, + "title": "Config" + }, + "save_dir": { + "anyOf": [ + { + "type": "string" + }, + { + "type": "null" + } + ], + "default": null, + "title": "Save Dir" + } + }, + "required": [ + "adapter_name" + ], + "title": "AddAdapterRequest", + "type": "object" + } + } + ], + "cookies": [], + "headers": [], + "operationId": "add_adapter_to_model", + "path": [], + "query": [], + "response": { + "description": "Response body for the /add_adapter_to_sampler endpoint.", + "properties": { + "adapter_name": { + "title": "Adapter Name", + "type": "string" + }, + "status": { + "default": "ok", + "title": "Status", + "type": "string" + } + }, + "required": [ + "adapter_name" + ], + "title": "AddAdapterResponse", + "type": "object" + }, + "responses": {}, + "statusCode": 200 } }, "/twinkle/add_metric": { "POST": { - "operationId": "add_metric_twinkle_add_metric_post", - "parameters": [], - "responses": [ - "200", - "422" - ] + "body": [ + { + "name": "body", + "required": true, + "schema": { + "additionalProperties": true, + "properties": { + "adapter_name": { + "title": "Adapter Name", + "type": "string" + }, + "is_training": { + "anyOf": [ + { + "type": "boolean" + }, + { + "type": "null" + } + ], + "default": null, + "title": "Is Training" + }, + "metric_cls": { + "title": "Metric Cls", + "type": "string" + } + }, + "required": [ + "metric_cls", + "adapter_name" + ], + "title": "AddMetricRequest", + "type": "object" + } + } + ], + "cookies": [], + "headers": [], + "operationId": "add_metric", + "path": [], + "query": [], + "response": { + "type": "null" + }, + "responses": {}, + "statusCode": 200 } }, "/twinkle/apply_patch": { "POST": { - "operationId": "apply_patch_twinkle_apply_patch_post", - "parameters": [], - "responses": [ - "200", - "422" - ] + "body": [ + { + "name": "body", + "required": true, + "schema": { + "additionalProperties": true, + "properties": { + "adapter_name": { + "title": "Adapter Name", + "type": "string" + }, + "patch_cls": { + "title": "Patch Cls", + "type": "string" + } + }, + "required": [ + "patch_cls", + "adapter_name" + ], + "title": "ApplyPatchRequest", + "type": "object" + } + } + ], + "cookies": [], + "headers": [], + "operationId": "apply_patch", + "path": [], + "query": [], + "response": { + "type": "null" + }, + "responses": {}, + "statusCode": 200 } }, "/twinkle/backward": { "POST": { - "operationId": "backward_twinkle_backward_post", - "parameters": [], - "responses": [ - "200", - "422" - ] + "body": [ + { + "name": "body", + "required": true, + "schema": { + "additionalProperties": true, + "properties": { + "adapter_name": { + "title": "Adapter Name", + "type": "string" + } + }, + "required": [ + "adapter_name" + ], + "title": "AdapterRequest", + "type": "object" + } + } + ], + "cookies": [], + "headers": [], + "operationId": "backward", + "path": [], + "query": [], + "response": { + "type": "null" + }, + "responses": {}, + "statusCode": 200 } }, "/twinkle/calculate_loss": { "POST": { - "operationId": "calculate_loss_twinkle_calculate_loss_post", - "parameters": [], - "responses": [ - "200", - "422" - ] + "body": [ + { + "name": "body", + "required": true, + "schema": { + "additionalProperties": true, + "properties": { + "adapter_name": { + "title": "Adapter Name", + "type": "string" + } + }, + "required": [ + "adapter_name" + ], + "title": "AdapterRequest", + "type": "object" + } + } + ], + "cookies": [], + "headers": [], + "operationId": "calculate_loss", + "path": [], + "query": [], + "response": { + "description": "Response for /calculate_loss endpoint (returns float).", + "properties": { + "result": { + "title": "Result", + "type": "number" + } + }, + "required": [ + "result" + ], + "title": "CalculateLossResponse", + "type": "object" + }, + "responses": {}, + "statusCode": 200 } }, "/twinkle/calculate_metric": { "POST": { - "operationId": "calculate_metric_twinkle_calculate_metric_post", - "parameters": [], - "responses": [ - "200", - "422" - ] + "body": [ + { + "name": "body", + "required": true, + "schema": { + "additionalProperties": true, + "properties": { + "adapter_name": { + "title": "Adapter Name", + "type": "string" + }, + "is_training": { + "default": true, + "title": "Is Training", + "type": "boolean" + } + }, + "required": [ + "adapter_name" + ], + "title": "CalculateMetricRequest", + "type": "object" + } + } + ], + "cookies": [], + "headers": [], + "operationId": "calculate_metric", + "path": [], + "query": [], + "response": { + "description": "Response for /calculate_metric endpoint (returns Dict).", + "properties": { + "result": { + "additionalProperties": true, + "title": "Result", + "type": "object" + } + }, + "required": [ + "result" + ], + "title": "CalculateMetricResponse", + "type": "object" + }, + "responses": {}, + "statusCode": 200 } }, "/twinkle/clip_grad_and_step": { "POST": { - "operationId": "clip_grad_and_step_twinkle_clip_grad_and_step_post", - "parameters": [], - "responses": [ - "200", - "422" - ] + "body": [ + { + "name": "body", + "required": true, + "schema": { + "additionalProperties": true, + "properties": { + "adapter_name": { + "title": "Adapter Name", + "type": "string" + }, + "max_grad_norm": { + "default": 1.0, + "title": "Max Grad Norm", + "type": "number" + }, + "norm_type": { + "default": 2, + "title": "Norm Type", + "type": "integer" + } + }, + "required": [ + "adapter_name" + ], + "title": "ClipGradAndStepRequest", + "type": "object" + } + } + ], + "cookies": [], + "headers": [], + "operationId": "clip_grad_and_step", + "path": [], + "query": [], + "response": { + "type": "null" + }, + "responses": {}, + "statusCode": 200 } }, "/twinkle/clip_grad_norm": { "POST": { - "operationId": "clip_grad_norm_twinkle_clip_grad_norm_post", - "parameters": [], - "responses": [ - "200", - "422" - ] + "body": [ + { + "name": "body", + "required": true, + "schema": { + "additionalProperties": true, + "properties": { + "adapter_name": { + "title": "Adapter Name", + "type": "string" + } + }, + "required": [ + "adapter_name" + ], + "title": "AdapterRequest", + "type": "object" + } + } + ], + "cookies": [], + "headers": [], + "operationId": "clip_grad_norm", + "path": [], + "query": [], + "response": { + "description": "Response for /clip_grad_norm endpoint (returns float as str).", + "properties": { + "result": { + "title": "Result", + "type": "string" + } + }, + "required": [ + "result" + ], + "title": "ClipGradNormResponse", + "type": "object" + }, + "responses": {}, + "statusCode": 200 } }, "/twinkle/create": { "POST": { - "operationId": "create_twinkle_create_post", - "parameters": [], - "responses": [ - "200", - "422" - ] + "body": [ + { + "name": "body", + "required": true, + "schema": { + "additionalProperties": true, + "properties": {}, + "title": "CreateRequest", + "type": "object" + } + } + ], + "cookies": [], + "headers": [], + "operationId": "create", + "path": [], + "query": [], + "response": { + "description": "Response for /create endpoint.", + "properties": { + "status": { + "default": "ok", + "title": "Status", + "type": "string" + } + }, + "title": "CreateResponse", + "type": "object" + }, + "responses": {}, + "statusCode": 200 } }, "/twinkle/forward": { "POST": { - "operationId": "forward_twinkle_forward_post", - "parameters": [], - "responses": [ - "200", - "422" - ] + "body": [ + { + "name": "body", + "required": true, + "schema": { + "additionalProperties": true, + "properties": { + "adapter_name": { + "title": "Adapter Name", + "type": "string" + }, + "inputs": { + "title": "Inputs" + } + }, + "required": [ + "inputs", + "adapter_name" + ], + "title": "ForwardRequest", + "type": "object" + } + } + ], + "cookies": [], + "headers": [], + "operationId": "forward", + "path": [], + "query": [], + "response": { + "description": "Response for /forward and /forward_only endpoints (returns ModelOutput).", + "properties": { + "result": { + "title": "Result" + } + }, + "required": [ + "result" + ], + "title": "ForwardResponse", + "type": "object" + }, + "responses": {}, + "statusCode": 200 } }, - "/twinkle/forward_from_data_plane": { + "/twinkle/forward_backward": { "POST": { - "operationId": "forward_from_data_plane_twinkle_forward_from_data_plane_post", - "parameters": [], - "responses": [ - "200", - "422" - ] + "body": [ + { + "name": "body", + "required": true, + "schema": { + "additionalProperties": true, + "properties": { + "adapter_name": { + "title": "Adapter Name", + "type": "string" + }, + "inputs": { + "title": "Inputs" + } + }, + "required": [ + "inputs", + "adapter_name" + ], + "title": "ForwardRequest", + "type": "object" + } + } + ], + "cookies": [], + "headers": [], + "operationId": "forward_backward", + "path": [], + "query": [], + "response": { + "description": "Response for /forward_backward endpoint (returns ModelOutput).", + "properties": { + "result": { + "title": "Result" + } + }, + "required": [ + "result" + ], + "title": "ForwardBackwardResponse", + "type": "object" + }, + "responses": {}, + "statusCode": 200 } }, - "/twinkle/forward_backward": { + "/twinkle/forward_backward_from_data_plane": { "POST": { - "operationId": "forward_backward_twinkle_forward_backward_post", - "parameters": [], - "responses": [ - "200", - "422" - ] + "body": [ + { + "name": "body", + "required": true, + "schema": { + "$defs": { + "DataRef": { + "description": "Opaque reference to rows stored in the server-side TransferQueue.", + "properties": { + "fields": { + "items": { + "type": "string" + }, + "title": "Fields", + "type": "array" + }, + "kind": { + "default": "data", + "title": "Kind", + "type": "string" + }, + "num_tokens": { + "default": 0, + "title": "Num Tokens", + "type": "integer" + }, + "ref_id": { + "title": "Ref Id", + "type": "string" + }, + "size": { + "title": "Size", + "type": "integer" + } + }, + "required": [ + "ref_id", + "size" + ], + "title": "DataRef", + "type": "object" + } + }, + "additionalProperties": true, + "properties": { + "adapter_name": { + "title": "Adapter Name", + "type": "string" + }, + "input_field": { + "anyOf": [ + { + "type": "string" + }, + { + "type": "null" + } + ], + "default": null, + "title": "Input Field" + }, + "input_refs": { + "items": { + "$ref": "#/$defs/DataRef" + }, + "minItems": 1, + "title": "Input Refs", + "type": "array" + }, + "kwarg_fields": { + "additionalProperties": { + "type": "string" + }, + "title": "Kwarg Fields", + "type": "object" + } + }, + "required": [ + "input_refs", + "adapter_name" + ], + "title": "DataPlaneForwardRequest", + "type": "object" + } + } + ], + "cookies": [], + "headers": [], + "operationId": "forward_backward_from_data_plane", + "path": [], + "query": [], + "response": { + "description": "Response for /forward_backward endpoint (returns ModelOutput).", + "properties": { + "result": { + "title": "Result" + } + }, + "required": [ + "result" + ], + "title": "ForwardBackwardResponse", + "type": "object" + }, + "responses": {}, + "statusCode": 200 } }, - "/twinkle/forward_backward_from_data_plane": { + "/twinkle/forward_from_data_plane": { "POST": { - "operationId": "forward_backward_from_data_plane_twinkle_forward_backward_from_data_plane_post", - "parameters": [], - "responses": [ - "200", - "422" - ] + "body": [ + { + "name": "body", + "required": true, + "schema": { + "$defs": { + "DataRef": { + "description": "Opaque reference to rows stored in the server-side TransferQueue.", + "properties": { + "fields": { + "items": { + "type": "string" + }, + "title": "Fields", + "type": "array" + }, + "kind": { + "default": "data", + "title": "Kind", + "type": "string" + }, + "num_tokens": { + "default": 0, + "title": "Num Tokens", + "type": "integer" + }, + "ref_id": { + "title": "Ref Id", + "type": "string" + }, + "size": { + "title": "Size", + "type": "integer" + } + }, + "required": [ + "ref_id", + "size" + ], + "title": "DataRef", + "type": "object" + } + }, + "additionalProperties": true, + "properties": { + "adapter_name": { + "title": "Adapter Name", + "type": "string" + }, + "input_field": { + "anyOf": [ + { + "type": "string" + }, + { + "type": "null" + } + ], + "default": null, + "title": "Input Field" + }, + "input_refs": { + "items": { + "$ref": "#/$defs/DataRef" + }, + "minItems": 1, + "title": "Input Refs", + "type": "array" + }, + "kwarg_fields": { + "additionalProperties": { + "type": "string" + }, + "title": "Kwarg Fields", + "type": "object" + } + }, + "required": [ + "input_refs", + "adapter_name" + ], + "title": "DataPlaneForwardRequest", + "type": "object" + } + } + ], + "cookies": [], + "headers": [], + "operationId": "forward_from_data_plane", + "path": [], + "query": [], + "response": { + "description": "Response for /forward and /forward_only endpoints (returns ModelOutput).", + "properties": { + "result": { + "title": "Result" + } + }, + "required": [ + "result" + ], + "title": "ForwardResponse", + "type": "object" + }, + "responses": {}, + "statusCode": 200 } }, "/twinkle/forward_only": { "POST": { - "operationId": "forward_only_twinkle_forward_only_post", - "parameters": [], - "responses": [ - "200", - "422" - ] + "body": [ + { + "name": "body", + "required": true, + "schema": { + "additionalProperties": true, + "properties": { + "adapter_name": { + "anyOf": [ + { + "type": "string" + }, + { + "type": "null" + } + ], + "default": null, + "title": "Adapter Name" + }, + "inputs": { + "title": "Inputs" + } + }, + "required": [ + "inputs" + ], + "title": "ForwardOnlyRequest", + "type": "object" + } + } + ], + "cookies": [], + "headers": [], + "operationId": "forward_only", + "path": [], + "query": [], + "response": { + "description": "Response for /forward and /forward_only endpoints (returns ModelOutput).", + "properties": { + "result": { + "title": "Result" + } + }, + "required": [ + "result" + ], + "title": "ForwardResponse", + "type": "object" + }, + "responses": {}, + "statusCode": 200 } }, "/twinkle/forward_only_from_data_plane": { "POST": { - "operationId": "forward_only_from_data_plane_twinkle_forward_only_from_data_plane_post", - "parameters": [], - "responses": [ - "200", - "422" - ] + "body": [ + { + "name": "body", + "required": true, + "schema": { + "$defs": { + "DataRef": { + "description": "Opaque reference to rows stored in the server-side TransferQueue.", + "properties": { + "fields": { + "items": { + "type": "string" + }, + "title": "Fields", + "type": "array" + }, + "kind": { + "default": "data", + "title": "Kind", + "type": "string" + }, + "num_tokens": { + "default": 0, + "title": "Num Tokens", + "type": "integer" + }, + "ref_id": { + "title": "Ref Id", + "type": "string" + }, + "size": { + "title": "Size", + "type": "integer" + } + }, + "required": [ + "ref_id", + "size" + ], + "title": "DataRef", + "type": "object" + } + }, + "additionalProperties": true, + "properties": { + "adapter_name": { + "title": "Adapter Name", + "type": "string" + }, + "input_field": { + "anyOf": [ + { + "type": "string" + }, + { + "type": "null" + } + ], + "default": null, + "title": "Input Field" + }, + "input_refs": { + "items": { + "$ref": "#/$defs/DataRef" + }, + "minItems": 1, + "title": "Input Refs", + "type": "array" + }, + "kwarg_fields": { + "additionalProperties": { + "type": "string" + }, + "title": "Kwarg Fields", + "type": "object" + }, + "output_fields": { + "additionalProperties": { + "type": "string" + }, + "title": "Output Fields", + "type": "object" + }, + "output_ref": { + "anyOf": [ + { + "$ref": "#/$defs/DataRef" + }, + { + "type": "null" + } + ], + "default": null + } + }, + "required": [ + "input_refs", + "adapter_name" + ], + "title": "DataPlaneForwardOnlyRequest", + "type": "object" + } + } + ], + "cookies": [], + "headers": [], + "operationId": "forward_only_from_data_plane", + "path": [], + "query": [], + "response": { + "description": "Response for /forward and /forward_only endpoints (returns ModelOutput).", + "properties": { + "result": { + "title": "Result" + } + }, + "required": [ + "result" + ], + "title": "ForwardResponse", + "type": "object" + }, + "responses": {}, + "statusCode": 200 } }, "/twinkle/get_state_dict": { "POST": { - "operationId": "get_state_dict_twinkle_get_state_dict_post", - "parameters": [], - "responses": [ - "200", - "422" - ] + "body": [ + { + "name": "body", + "required": true, + "schema": { + "additionalProperties": true, + "properties": { + "adapter_name": { + "title": "Adapter Name", + "type": "string" + } + }, + "required": [ + "adapter_name" + ], + "title": "GetStateDictRequest", + "type": "object" + } + } + ], + "cookies": [], + "headers": [], + "operationId": "get_state_dict", + "path": [], + "query": [], + "response": { + "description": "Response for /get_state_dict endpoint (returns Dict).", + "properties": { + "result": { + "additionalProperties": true, + "title": "Result", + "type": "object" + } + }, + "required": [ + "result" + ], + "title": "GetStateDictResponse", + "type": "object" + }, + "responses": {}, + "statusCode": 200 } }, "/twinkle/get_train_configs": { "POST": { - "operationId": "get_train_configs_twinkle_get_train_configs_post", - "parameters": [], - "responses": [ - "200", - "422" - ] + "body": [ + { + "name": "body", + "required": true, + "schema": { + "additionalProperties": true, + "properties": { + "adapter_name": { + "title": "Adapter Name", + "type": "string" + } + }, + "required": [ + "adapter_name" + ], + "title": "AdapterRequest", + "type": "object" + } + } + ], + "cookies": [], + "headers": [], + "operationId": "get_train_configs", + "path": [], + "query": [], + "response": { + "description": "Response for /get_train_configs endpoint (returns str).", + "properties": { + "result": { + "title": "Result", + "type": "string" + } + }, + "required": [ + "result" + ], + "title": "GetTrainConfigsResponse", + "type": "object" + }, + "responses": {}, + "statusCode": 200 } }, "/twinkle/load": { "POST": { - "operationId": "load_twinkle_load_post", - "parameters": [], - "responses": [ - "200", - "422" - ] + "body": [ + { + "name": "body", + "required": true, + "schema": { + "additionalProperties": true, + "properties": { + "adapter_name": { + "title": "Adapter Name", + "type": "string" + }, + "load_optimizer": { + "default": false, + "title": "Load Optimizer", + "type": "boolean" + }, + "name": { + "title": "Name", + "type": "string" + } + }, + "required": [ + "adapter_name", + "name" + ], + "title": "LoadRequest", + "type": "object" + } + } + ], + "cookies": [], + "headers": [], + "operationId": "load", + "path": [], + "query": [], + "response": { + "type": "null" + }, + "responses": {}, + "statusCode": 200 } }, "/twinkle/lr_step": { "POST": { - "operationId": "lr_step_twinkle_lr_step_post", - "parameters": [], - "responses": [ - "200", - "422" - ] + "body": [ + { + "name": "body", + "required": true, + "schema": { + "additionalProperties": true, + "properties": { + "adapter_name": { + "title": "Adapter Name", + "type": "string" + } + }, + "required": [ + "adapter_name" + ], + "title": "AdapterRequest", + "type": "object" + } + } + ], + "cookies": [], + "headers": [], + "operationId": "lr_step", + "path": [], + "query": [], + "response": { + "type": "null" + }, + "responses": {}, + "statusCode": 200 } }, "/twinkle/remove_adapter": { "POST": { - "operationId": "remove_adapter_twinkle_remove_adapter_post", - "parameters": [], - "responses": [ - "200", - "422" - ] + "body": [ + { + "name": "body", + "required": true, + "schema": { + "additionalProperties": true, + "properties": { + "adapter_name": { + "title": "Adapter Name", + "type": "string" + } + }, + "required": [ + "adapter_name" + ], + "title": "AdapterRequest", + "type": "object" + } + } + ], + "cookies": [], + "headers": [], + "operationId": "remove_adapter", + "path": [], + "query": [], + "response": { + "additionalProperties": { + "type": "string" + }, + "type": "object" + }, + "responses": {}, + "statusCode": 200 } }, "/twinkle/resume_from_checkpoint": { "POST": { - "operationId": "resume_from_checkpoint_twinkle_resume_from_checkpoint_post", - "parameters": [], - "responses": [ - "200", - "422" - ] + "body": [ + { + "name": "body", + "required": true, + "schema": { + "additionalProperties": true, + "description": "Request for /resume_from_checkpoint endpoint.", + "properties": { + "adapter_name": { + "default": "", + "title": "Adapter Name", + "type": "string" + }, + "name": { + "title": "Name", + "type": "string" + }, + "resume_only_model": { + "default": false, + "title": "Resume Only Model", + "type": "boolean" + } + }, + "required": [ + "name" + ], + "title": "ResumeFromCheckpointRequest", + "type": "object" + } + } + ], + "cookies": [], + "headers": [], + "operationId": "resume_from_checkpoint", + "path": [], + "query": [], + "response": { + "description": "Response for /resume_from_checkpoint endpoint.", + "properties": { + "result": { + "additionalProperties": true, + "title": "Result", + "type": "object" + } + }, + "required": [ + "result" + ], + "title": "TrainingProgressResponse", + "type": "object" + }, + "responses": {}, + "statusCode": 200 } }, "/twinkle/save": { "POST": { - "operationId": "save_twinkle_save_post", - "parameters": [], - "responses": [ - "200", - "422" - ] + "body": [ + { + "name": "body", + "required": true, + "schema": { + "additionalProperties": true, + "properties": { + "adapter_name": { + "title": "Adapter Name", + "type": "string" + }, + "is_sampler": { + "default": false, + "title": "Is Sampler", + "type": "boolean" + }, + "name": { + "anyOf": [ + { + "type": "string" + }, + { + "type": "null" + } + ], + "default": null, + "title": "Name" + }, + "save_optimizer": { + "default": false, + "title": "Save Optimizer", + "type": "boolean" + } + }, + "required": [ + "adapter_name" + ], + "title": "SaveRequest", + "type": "object" + } + } + ], + "cookies": [], + "headers": [], + "operationId": "save", + "path": [], + "query": [], + "response": { + "description": "Response for /save endpoint (returns twinkle path + checkpoint dir).", + "properties": { + "checkpoint_dir": { + "anyOf": [ + { + "type": "string" + }, + { + "type": "null" + } + ], + "default": null, + "title": "Checkpoint Dir" + }, + "twinkle_path": { + "title": "Twinkle Path", + "type": "string" + } + }, + "required": [ + "twinkle_path" + ], + "title": "SaveResponse", + "type": "object" + }, + "responses": {}, + "statusCode": 200 } }, "/twinkle/set_loss": { "POST": { - "operationId": "set_loss_twinkle_set_loss_post", - "parameters": [], - "responses": [ - "200", - "422" - ] + "body": [ + { + "name": "body", + "required": true, + "schema": { + "additionalProperties": true, + "properties": { + "adapter_name": { + "title": "Adapter Name", + "type": "string" + }, + "loss_cls": { + "title": "Loss Cls", + "type": "string" + } + }, + "required": [ + "loss_cls", + "adapter_name" + ], + "title": "SetLossRequest", + "type": "object" + } + } + ], + "cookies": [], + "headers": [], + "operationId": "set_loss", + "path": [], + "query": [], + "response": { + "type": "null" + }, + "responses": {}, + "statusCode": 200 } }, "/twinkle/set_lr_scheduler": { "POST": { - "operationId": "set_lr_scheduler_twinkle_set_lr_scheduler_post", - "parameters": [], - "responses": [ - "200", - "422" - ] + "body": [ + { + "name": "body", + "required": true, + "schema": { + "additionalProperties": true, + "properties": { + "adapter_name": { + "title": "Adapter Name", + "type": "string" + }, + "scheduler_cls": { + "title": "Scheduler Cls", + "type": "string" + } + }, + "required": [ + "scheduler_cls", + "adapter_name" + ], + "title": "SetLrSchedulerRequest", + "type": "object" + } + } + ], + "cookies": [], + "headers": [], + "operationId": "set_lr_scheduler", + "path": [], + "query": [], + "response": { + "type": "null" + }, + "responses": {}, + "statusCode": 200 } }, "/twinkle/set_optimizer": { "POST": { - "operationId": "set_optimizer_twinkle_set_optimizer_post", - "parameters": [], - "responses": [ - "200", - "422" - ] + "body": [ + { + "name": "body", + "required": true, + "schema": { + "additionalProperties": true, + "properties": { + "adapter_name": { + "title": "Adapter Name", + "type": "string" + }, + "optimizer_cls": { + "title": "Optimizer Cls", + "type": "string" + } + }, + "required": [ + "optimizer_cls", + "adapter_name" + ], + "title": "SetOptimizerRequest", + "type": "object" + } + } + ], + "cookies": [], + "headers": [], + "operationId": "set_optimizer", + "path": [], + "query": [], + "response": { + "type": "null" + }, + "responses": {}, + "statusCode": 200 } }, "/twinkle/set_processor": { "POST": { - "operationId": "set_processor_twinkle_set_processor_post", - "parameters": [], - "responses": [ - "200", - "422" - ] + "body": [ + { + "name": "body", + "required": true, + "schema": { + "additionalProperties": true, + "properties": { + "adapter_name": { + "title": "Adapter Name", + "type": "string" + }, + "processor_cls": { + "title": "Processor Cls", + "type": "string" + } + }, + "required": [ + "processor_cls", + "adapter_name" + ], + "title": "SetProcessorRequest", + "type": "object" + } + } + ], + "cookies": [], + "headers": [], + "operationId": "set_processor", + "path": [], + "query": [], + "response": { + "type": "null" + }, + "responses": {}, + "statusCode": 200 } }, "/twinkle/set_template": { "POST": { - "operationId": "set_template_twinkle_set_template_post", - "parameters": [], - "responses": [ - "200", - "422" - ] + "body": [ + { + "name": "body", + "required": true, + "schema": { + "additionalProperties": true, + "properties": { + "adapter_name": { + "title": "Adapter Name", + "type": "string" + }, + "template_cls": { + "title": "Template Cls", + "type": "string" + } + }, + "required": [ + "template_cls", + "adapter_name" + ], + "title": "SetTemplateRequest", + "type": "object" + } + } + ], + "cookies": [], + "headers": [], + "operationId": "set_template", + "path": [], + "query": [], + "response": { + "type": "null" + }, + "responses": {}, + "statusCode": 200 } }, "/twinkle/step": { "POST": { - "operationId": "step_twinkle_step_post", - "parameters": [], - "responses": [ - "200", - "422" - ] + "body": [ + { + "name": "body", + "required": true, + "schema": { + "additionalProperties": true, + "properties": { + "adapter_name": { + "title": "Adapter Name", + "type": "string" + } + }, + "required": [ + "adapter_name" + ], + "title": "AdapterRequest", + "type": "object" + } + } + ], + "cookies": [], + "headers": [], + "operationId": "step", + "path": [], + "query": [], + "response": { + "type": "null" + }, + "responses": {}, + "statusCode": 200 } }, "/twinkle/upload_status/{request_id}": { "GET": { - "operationId": "upload_status_twinkle_upload_status__request_id__get", - "parameters": [ + "body": [], + "cookies": [], + "headers": [], + "operationId": "upload_status", + "path": [ { - "in": "path", "name": "request_id", "required": true, "schema": { - "title": "Request Id", "type": "string" } } ], - "responses": [ - "200", - "422" - ] + "query": [], + "response": { + "description": "Response for /upload_status/{request_id} endpoint.", + "properties": { + "error": { + "anyOf": [ + { + "type": "string" + }, + { + "type": "null" + } + ], + "default": null, + "title": "Error" + }, + "request_id": { + "title": "Request Id", + "type": "string" + }, + "status": { + "title": "Status", + "type": "string" + } + }, + "required": [ + "request_id", + "status" + ], + "title": "UploadStatusResponse", + "type": "object" + }, + "responses": {}, + "statusCode": 200 } }, "/twinkle/upload_to_hub": { "POST": { - "operationId": "upload_to_hub_twinkle_upload_to_hub_post", - "parameters": [], - "responses": [ - "200", - "422" - ] + "body": [ + { + "name": "body", + "required": true, + "schema": { + "additionalProperties": true, + "properties": { + "async_upload": { + "default": false, + "title": "Async Upload", + "type": "boolean" + }, + "checkpoint_dir": { + "anyOf": [ + { + "type": "string" + }, + { + "additionalProperties": true, + "type": "object" + } + ], + "title": "Checkpoint Dir" + }, + "hub_model_id": { + "title": "Hub Model Id", + "type": "string" + }, + "hub_token": { + "anyOf": [ + { + "type": "string" + }, + { + "type": "null" + } + ], + "default": null, + "title": "Hub Token" + } + }, + "required": [ + "checkpoint_dir", + "hub_model_id" + ], + "title": "UploadToHubRequest", + "type": "object" + } + } + ], + "cookies": [], + "headers": [], + "operationId": "upload_to_hub", + "path": [], + "query": [], + "response": { + "description": "Response for /upload_to_hub endpoint.", + "properties": { + "request_id": { + "title": "Request Id", + "type": "string" + } + }, + "required": [ + "request_id" + ], + "title": "UploadToHubResponse", + "type": "object" + }, + "responses": {}, + "statusCode": 200 } }, "/twinkle/zero_grad": { "POST": { - "operationId": "zero_grad_twinkle_zero_grad_post", - "parameters": [], - "responses": [ - "200", - "422" - ] + "body": [ + { + "name": "body", + "required": true, + "schema": { + "additionalProperties": true, + "properties": { + "adapter_name": { + "title": "Adapter Name", + "type": "string" + } + }, + "required": [ + "adapter_name" + ], + "title": "AdapterRequest", + "type": "object" + } + } + ], + "cookies": [], + "headers": [], + "operationId": "zero_grad", + "path": [], + "query": [], + "response": { + "type": "null" + }, + "responses": {}, + "statusCode": 200 } } } @@ -1003,22 +7175,101 @@ "paths": { "/twinkle/call": { "POST": { - "operationId": "call_twinkle_call_post", - "parameters": [], - "responses": [ - "200", - "422" - ] + "body": [ + { + "name": "body", + "required": true, + "schema": { + "additionalProperties": true, + "properties": { + "function": { + "title": "Function", + "type": "string" + }, + "processor_id": { + "title": "Processor Id", + "type": "string" + } + }, + "required": [ + "processor_id", + "function" + ], + "title": "ProcessorCallRequest", + "type": "object" + } + } + ], + "cookies": [], + "headers": [], + "operationId": "call", + "path": [], + "query": [], + "response": { + "description": "Response body for the /call endpoint.", + "properties": { + "result": { + "title": "Result" + } + }, + "required": [ + "result" + ], + "title": "ProcessorCallResponse", + "type": "object" + }, + "responses": {}, + "statusCode": 200 } }, "/twinkle/create": { "POST": { - "operationId": "create_twinkle_create_post", - "parameters": [], - "responses": [ - "200", - "422" - ] + "body": [ + { + "name": "body", + "required": true, + "schema": { + "additionalProperties": true, + "properties": { + "class_type": { + "title": "Class Type", + "type": "string" + }, + "processor_type": { + "title": "Processor Type", + "type": "string" + } + }, + "required": [ + "processor_type", + "class_type" + ], + "title": "ProcessorCreateRequest", + "type": "object" + } + } + ], + "cookies": [], + "headers": [], + "operationId": "create", + "path": [], + "query": [], + "response": { + "description": "Response body for the /create endpoint.", + "properties": { + "processor_id": { + "title": "Processor Id", + "type": "string" + } + }, + "required": [ + "processor_id" + ], + "title": "ProcessorCreateResponse", + "type": "object" + }, + "responses": {}, + "statusCode": 200 } } } @@ -1027,91 +7278,1074 @@ "paths": { "/tinker/asample": { "POST": { - "operationId": "asample_tinker_asample_post", - "parameters": [], - "responses": [ - "200", - "422" - ] + "body": [ + { + "name": "body", + "required": true, + "schema": { + "$defs": { + "EncodedTextChunk": { + "additionalProperties": false, + "properties": { + "tokens": { + "items": { + "type": "integer" + }, + "title": "Tokens", + "type": "array" + }, + "type": { + "const": "encoded_text", + "default": "encoded_text", + "title": "Type", + "type": "string" + } + }, + "required": [ + "tokens" + ], + "title": "EncodedTextChunk", + "type": "object" + }, + "ImageAssetPointerChunk": { + "additionalProperties": false, + "properties": { + "expected_tokens": { + "anyOf": [ + { + "type": "integer" + }, + { + "type": "null" + } + ], + "default": null, + "title": "Expected Tokens" + }, + "format": { + "enum": [ + "png", + "jpeg" + ], + "title": "Format", + "type": "string" + }, + "location": { + "title": "Location", + "type": "string" + }, + "type": { + "const": "image_asset_pointer", + "default": "image_asset_pointer", + "title": "Type", + "type": "string" + } + }, + "required": [ + "format", + "location" + ], + "title": "ImageAssetPointerChunk", + "type": "object" + }, + "ImageChunk": { + "additionalProperties": false, + "properties": { + "data": { + "format": "binary", + "title": "Data", + "type": "string" + }, + "expected_tokens": { + "anyOf": [ + { + "type": "integer" + }, + { + "type": "null" + } + ], + "default": null, + "title": "Expected Tokens" + }, + "format": { + "enum": [ + "png", + "jpeg" + ], + "title": "Format", + "type": "string" + }, + "type": { + "const": "image", + "default": "image", + "title": "Type", + "type": "string" + } + }, + "required": [ + "data", + "format" + ], + "title": "ImageChunk", + "type": "object" + }, + "ModelInput": { + "additionalProperties": false, + "properties": { + "chunks": { + "items": { + "anyOf": [ + { + "$ref": "#/$defs/EncodedTextChunk" + }, + { + "$ref": "#/$defs/ImageAssetPointerChunk" + }, + { + "$ref": "#/$defs/ImageChunk" + } + ] + }, + "title": "Chunks", + "type": "array" + } + }, + "required": [ + "chunks" + ], + "title": "ModelInput", + "type": "object" + }, + "SamplingParams": { + "properties": { + "max_tokens": { + "anyOf": [ + { + "type": "integer" + }, + { + "type": "null" + } + ], + "default": null, + "title": "Max Tokens" + }, + "seed": { + "anyOf": [ + { + "type": "integer" + }, + { + "type": "null" + } + ], + "default": null, + "title": "Seed" + }, + "stop": { + "anyOf": [ + { + "type": "string" + }, + { + "items": { + "type": "string" + }, + "type": "array" + }, + { + "items": { + "type": "integer" + }, + "type": "array" + }, + { + "type": "null" + } + ], + "default": null, + "title": "Stop" + }, + "temperature": { + "default": 1, + "title": "Temperature", + "type": "number" + }, + "top_k": { + "default": -1, + "title": "Top K", + "type": "integer" + }, + "top_p": { + "default": 1, + "title": "Top P", + "type": "number" + } + }, + "title": "SamplingParams", + "type": "object" + } + }, + "additionalProperties": false, + "properties": { + "base_model": { + "anyOf": [ + { + "type": "string" + }, + { + "type": "null" + } + ], + "default": null, + "title": "Base Model" + }, + "model_path": { + "anyOf": [ + { + "type": "string" + }, + { + "type": "null" + } + ], + "default": null, + "title": "Model Path" + }, + "num_samples": { + "default": 1, + "title": "Num Samples", + "type": "integer" + }, + "prompt": { + "$ref": "#/$defs/ModelInput" + }, + "prompt_logprobs": { + "anyOf": [ + { + "type": "boolean" + }, + { + "type": "null" + } + ], + "default": null, + "title": "Prompt Logprobs" + }, + "sampling_params": { + "$ref": "#/$defs/SamplingParams" + }, + "sampling_session_id": { + "anyOf": [ + { + "type": "string" + }, + { + "type": "null" + } + ], + "default": null, + "title": "Sampling Session Id" + }, + "seq_id": { + "anyOf": [ + { + "type": "integer" + }, + { + "type": "null" + } + ], + "default": null, + "title": "Seq Id" + }, + "topk_prompt_logprobs": { + "default": 0, + "title": "Topk Prompt Logprobs", + "type": "integer" + }, + "type": { + "const": "sample", + "default": "sample", + "title": "Type", + "type": "string" + } + }, + "required": [ + "prompt", + "sampling_params" + ], + "title": "SampleRequest", + "type": "object" + } + } + ], + "cookies": [], + "headers": [], + "operationId": "asample", + "path": [], + "query": [], + "response": { + "properties": { + "model_id": { + "anyOf": [ + { + "type": "string" + }, + { + "type": "null" + } + ], + "default": null, + "title": "Model Id" + }, + "request_id": { + "title": "Request Id", + "type": "string" + } + }, + "required": [ + "request_id" + ], + "title": "UntypedAPIFuture", + "type": "object" + }, + "responses": {}, + "statusCode": 200 } }, "/twinkle/add_adapter_to_sampler": { "POST": { - "operationId": "add_adapter_to_sampler_twinkle_add_adapter_to_sampler_post", - "parameters": [], - "responses": [ - "200", - "422" - ] + "body": [ + { + "name": "body", + "required": true, + "schema": { + "additionalProperties": true, + "properties": { + "adapter_name": { + "title": "Adapter Name", + "type": "string" + }, + "config": { + "anyOf": [ + { + "type": "string" + }, + { + "type": "null" + } + ], + "default": null, + "title": "Config" + }, + "save_dir": { + "anyOf": [ + { + "type": "string" + }, + { + "type": "null" + } + ], + "default": null, + "title": "Save Dir" + } + }, + "required": [ + "adapter_name" + ], + "title": "AddAdapterRequest", + "type": "object" + } + } + ], + "cookies": [], + "headers": [], + "operationId": "add_adapter_to_sampler", + "path": [], + "query": [], + "response": { + "description": "Response body for the /add_adapter_to_sampler endpoint.", + "properties": { + "adapter_name": { + "title": "Adapter Name", + "type": "string" + }, + "status": { + "default": "ok", + "title": "Status", + "type": "string" + } + }, + "required": [ + "adapter_name" + ], + "title": "AddAdapterResponse", + "type": "object" + }, + "responses": {}, + "statusCode": 200 } }, "/twinkle/apply_patch": { "POST": { - "operationId": "apply_patch_twinkle_apply_patch_post", - "parameters": [], - "responses": [ - "200", - "422" - ] + "body": [ + { + "name": "body", + "required": true, + "schema": { + "additionalProperties": true, + "properties": { + "adapter_name": { + "title": "Adapter Name", + "type": "string" + }, + "patch_cls": { + "title": "Patch Cls", + "type": "string" + } + }, + "required": [ + "patch_cls", + "adapter_name" + ], + "title": "ApplyPatchRequest", + "type": "object" + } + } + ], + "cookies": [], + "headers": [], + "operationId": "apply_patch", + "path": [], + "query": [], + "response": { + "type": "null" + }, + "responses": {}, + "statusCode": 200 } }, "/twinkle/create": { "POST": { - "operationId": "create_twinkle_create_post", - "parameters": [], - "responses": [ - "200" - ] + "body": [], + "cookies": [], + "headers": [], + "operationId": "create", + "path": [], + "query": [], + "response": { + "description": "Response for /create endpoint.", + "properties": { + "status": { + "default": "ok", + "title": "Status", + "type": "string" + } + }, + "title": "CreateResponse", + "type": "object" + }, + "responses": {}, + "statusCode": 200 } }, "/twinkle/sample": { "POST": { - "operationId": "sample_twinkle_sample_post", - "parameters": [], - "responses": [ - "200", - "422" - ] + "body": [ + { + "name": "body", + "required": true, + "schema": { + "description": "Request body for the /sample endpoint.", + "properties": { + "adapter_name": { + "default": "", + "description": "Adapter name for LoRA inference", + "title": "Adapter Name", + "type": "string" + }, + "adapter_uri": { + "anyOf": [ + { + "type": "string" + }, + { + "type": "null" + } + ], + "default": null, + "description": "Adapter URI (twinkle:// path or local path) for LoRA inference", + "title": "Adapter Uri" + }, + "inputs": { + "description": "List of Trajectory or InputFeature dicts", + "title": "Inputs" + }, + "sampling_params": { + "anyOf": [ + { + "additionalProperties": true, + "type": "object" + }, + { + "type": "null" + } + ], + "default": null, + "description": "Sampling parameters (max_tokens, temperature, num_samples, etc.)", + "title": "Sampling Params" + } + }, + "required": [ + "inputs" + ], + "title": "SampleRequest", + "type": "object" + } + } + ], + "cookies": [], + "headers": [], + "operationId": "sample", + "path": [], + "query": [], + "response": { + "$defs": { + "SampleResponseModel": { + "description": "Mirroring twinkle.data_format.SampleResponse.", + "properties": { + "prompt_logprobs": { + "anyOf": [ + { + "items": { + "anyOf": [ + { + "type": "number" + }, + { + "type": "null" + } + ] + }, + "type": "array" + }, + { + "type": "null" + } + ], + "default": null, + "title": "Prompt Logprobs" + }, + "prompt_token_ids": { + "anyOf": [ + { + "items": { + "type": "integer" + }, + "type": "array" + }, + { + "type": "null" + } + ], + "default": null, + "description": "Token IDs of the prompt the sequences continue", + "title": "Prompt Token Ids" + }, + "sequences": { + "description": "List of sampled sequences", + "items": { + "$ref": "#/$defs/SampledSequenceModel" + }, + "title": "Sequences", + "type": "array" + }, + "topk_prompt_logprobs": { + "anyOf": [ + { + "items": { + "anyOf": [ + { + "items": { + "maxItems": 2, + "minItems": 2, + "prefixItems": [ + { + "type": "integer" + }, + { + "type": "number" + } + ], + "type": "array" + }, + "type": "array" + }, + { + "type": "null" + } + ] + }, + "type": "array" + }, + { + "type": "null" + } + ], + "default": null, + "title": "Topk Prompt Logprobs" + } + }, + "required": [ + "sequences" + ], + "title": "SampleResponseModel", + "type": "object" + }, + "SampledSequenceModel": { + "description": "A single sampled sequence, mirroring twinkle.data_format.SampledSequence.", + "properties": { + "decoded": { + "anyOf": [ + { + "type": "string" + }, + { + "type": "null" + } + ], + "default": null, + "description": "Decoded text of the sampled sequence", + "title": "Decoded" + }, + "logprobs": { + "anyOf": [ + { + "items": { + "anyOf": [ + { + "items": { + "maxItems": 2, + "minItems": 2, + "prefixItems": [ + { + "type": "integer" + }, + { + "type": "number" + } + ], + "type": "array" + }, + "type": "array" + }, + { + "type": "null" + } + ] + }, + "type": "array" + }, + { + "type": "null" + } + ], + "default": null, + "description": "Per-token log-probabilities", + "title": "Logprobs" + }, + "new_input_feature": { + "anyOf": [ + { + "additionalProperties": true, + "type": "object" + }, + { + "type": "null" + } + ], + "default": null, + "description": "Updated InputFeature after sampling (input_ids, labels, etc.)", + "title": "New Input Feature" + }, + "stop_reason": { + "description": "Stop reason: 'length' or 'stop'", + "enum": [ + "length", + "stop", + "abort", + "error" + ], + "title": "Stop Reason", + "type": "string" + }, + "tokens": { + "description": "Token IDs of the sampled sequence", + "items": { + "type": "integer" + }, + "title": "Tokens", + "type": "array" + } + }, + "required": [ + "stop_reason", + "tokens" + ], + "title": "SampledSequenceModel", + "type": "object" + } + }, + "description": "Response body for the /sample endpoint", + "properties": { + "samples": { + "description": "List of sample responses", + "items": { + "$ref": "#/$defs/SampleResponseModel" + }, + "title": "Samples", + "type": "array" + } + }, + "required": [ + "samples" + ], + "title": "SampleResponseModelList", + "type": "object" + }, + "responses": {}, + "statusCode": 200 } }, - "/twinkle/sample_to_data_plane": { + "/twinkle/sample_stream": { "POST": { - "operationId": "sample_to_data_plane_twinkle_sample_to_data_plane_post", - "parameters": [], - "responses": [ - "200", - "422" - ] + "body": [ + { + "name": "body", + "required": true, + "schema": { + "description": "Request body for the /sample endpoint.", + "properties": { + "adapter_name": { + "default": "", + "description": "Adapter name for LoRA inference", + "title": "Adapter Name", + "type": "string" + }, + "adapter_uri": { + "anyOf": [ + { + "type": "string" + }, + { + "type": "null" + } + ], + "default": null, + "description": "Adapter URI (twinkle:// path or local path) for LoRA inference", + "title": "Adapter Uri" + }, + "inputs": { + "description": "List of Trajectory or InputFeature dicts", + "title": "Inputs" + }, + "sampling_params": { + "anyOf": [ + { + "additionalProperties": true, + "type": "object" + }, + { + "type": "null" + } + ], + "default": null, + "description": "Sampling parameters (max_tokens, temperature, num_samples, etc.)", + "title": "Sampling Params" + } + }, + "required": [ + "inputs" + ], + "title": "SampleRequest", + "type": "object" + } + } + ], + "cookies": [], + "headers": [], + "operationId": "sample_stream", + "path": [], + "query": [], + "response": { + "type": "null" + }, + "responses": {}, + "statusCode": 200 } }, - "/twinkle/sample_stream": { + "/twinkle/sample_to_data_plane": { "POST": { - "operationId": "sample_stream_twinkle_sample_stream_post", - "parameters": [], - "responses": [ - "200", - "422" - ] + "body": [ + { + "name": "body", + "required": true, + "schema": { + "$defs": { + "DataRef": { + "description": "Opaque reference to rows stored in the server-side TransferQueue.", + "properties": { + "fields": { + "items": { + "type": "string" + }, + "title": "Fields", + "type": "array" + }, + "kind": { + "default": "data", + "title": "Kind", + "type": "string" + }, + "num_tokens": { + "default": 0, + "title": "Num Tokens", + "type": "integer" + }, + "ref_id": { + "title": "Ref Id", + "type": "string" + }, + "size": { + "title": "Size", + "type": "integer" + } + }, + "required": [ + "ref_id", + "size" + ], + "title": "DataRef", + "type": "object" + } + }, + "properties": { + "adapter_name": { + "default": "", + "title": "Adapter Name", + "type": "string" + }, + "adapter_uri": { + "anyOf": [ + { + "type": "string" + }, + { + "type": "null" + } + ], + "default": null, + "title": "Adapter Uri" + }, + "group_ids": { + "anyOf": [ + { + "items": { + "type": "string" + }, + "type": "array" + }, + { + "type": "null" + } + ], + "default": null, + "title": "Group Ids" + }, + "input_ref": { + "anyOf": [ + { + "$ref": "#/$defs/DataRef" + }, + { + "type": "null" + } + ], + "default": null + }, + "inputs": { + "default": null, + "title": "Inputs" + }, + "num_samples": { + "default": 1, + "title": "Num Samples", + "type": "integer" + }, + "policy_version": { + "anyOf": [ + { + "type": "integer" + }, + { + "type": "null" + } + ], + "default": null, + "title": "Policy Version" + }, + "sampling_params": { + "anyOf": [ + { + "additionalProperties": true, + "type": "object" + }, + { + "type": "null" + } + ], + "default": null, + "title": "Sampling Params" + } + }, + "title": "DataPlaneSampleRequest", + "type": "object" + } + } + ], + "cookies": [], + "headers": [], + "operationId": "sample_to_data_plane", + "path": [], + "query": [], + "response": { + "description": "Opaque reference to rows stored in the server-side TransferQueue.", + "properties": { + "fields": { + "items": { + "type": "string" + }, + "title": "Fields", + "type": "array" + }, + "kind": { + "default": "data", + "title": "Kind", + "type": "string" + }, + "num_tokens": { + "default": 0, + "title": "Num Tokens", + "type": "integer" + }, + "ref_id": { + "title": "Ref Id", + "type": "string" + }, + "size": { + "title": "Size", + "type": "integer" + } + }, + "required": [ + "ref_id", + "size" + ], + "title": "DataRef", + "type": "object" + }, + "responses": {}, + "statusCode": 200 } }, "/twinkle/set_template": { "POST": { - "operationId": "set_template_twinkle_set_template_post", - "parameters": [], - "responses": [ - "200", - "422" - ] + "body": [ + { + "name": "body", + "required": true, + "schema": { + "additionalProperties": true, + "properties": { + "adapter_name": { + "title": "Adapter Name", + "type": "string" + }, + "template_cls": { + "title": "Template Cls", + "type": "string" + } + }, + "required": [ + "template_cls", + "adapter_name" + ], + "title": "SetTemplateRequest", + "type": "object" + } + } + ], + "cookies": [], + "headers": [], + "operationId": "set_template", + "path": [], + "query": [], + "response": { + "description": "Response for /set_template endpoint.", + "properties": { + "status": { + "default": "ok", + "title": "Status", + "type": "string" + } + }, + "title": "SetTemplateResponse", + "type": "object" + }, + "responses": {}, + "statusCode": 200 } }, "/twinkle/unload_adapter_paths": { "POST": { - "operationId": "unload_adapter_paths_twinkle_unload_adapter_paths_post", - "parameters": [], - "responses": [ - "200", - "422" - ] + "body": [ + { + "name": "body", + "required": true, + "schema": { + "properties": { + "adapter_paths": { + "items": { + "type": "string" + }, + "title": "Adapter Paths", + "type": "array" + } + }, + "required": [ + "adapter_paths" + ], + "title": "UnloadAdapterPathsRequest", + "type": "object" + } + } + ], + "cookies": [], + "headers": [], + "operationId": "unload_adapter_paths", + "path": [], + "query": [], + "response": { + "additionalProperties": { + "type": "string" + }, + "type": "object" + }, + "responses": {}, + "statusCode": 200 } } } diff --git a/tests/server/contract/client_api_harness.py b/tests/server/contract/client_api_harness.py index 0496a1b07..59d4ac5af 100644 --- a/tests/server/contract/client_api_harness.py +++ b/tests/server/contract/client_api_harness.py @@ -2,10 +2,10 @@ """ Client-API contract harness. -Builds the four FastAPI apps used by the Ray Serve deployments (Gateway, Model, -Sampler, Processor) by registering their route-registration helpers against a -fresh FastAPI instance, then extracts the client-facing surface (route paths, -HTTP methods, and request/response schemas) as a stable JSON dict. +Builds the five FastAPI apps used by the Ray Serve deployments (Data Plane, +Gateway, Model, Sampler, Processor) by registering their route-registration helpers against a +fresh FastAPI instance, then extracts route paths, methods, parameters, and +recursive request/response type shapes as a stable JSON dict. Used to: - snapshot the current surface into ``client_api_baseline.json`` before the @@ -25,12 +25,18 @@ """ from __future__ import annotations +import dataclasses import json -from collections.abc import Callable +import re +import sys +import types as pytypes +from collections.abc import Callable, Mapping, Sequence +from enum import Enum from fastapi import FastAPI -from fastapi.openapi.utils import get_openapi +from fastapi.routing import APIRoute from pathlib import Path -from typing import Any +from pydantic import BaseModel +from typing import Annotated, Any, Literal, Union, get_args, get_origin, get_type_hints # ----- App build helpers --------------------------------------------------- # @@ -39,6 +45,14 @@ def _noop_self() -> None: return None +def build_data_plane_app() -> FastAPI: + from twinkle.server.data_plane.handlers import register_data_plane_routes + + app = FastAPI() + register_data_plane_routes(app, _noop_self) + return app + + def build_gateway_app() -> FastAPI: from twinkle.server.gateway.openai_handlers import _register_openai_routes from twinkle.server.gateway.tinker_handlers import _register_tinker_routes @@ -80,6 +94,7 @@ def build_processor_app() -> FastAPI: APP_BUILDERS: dict[str, Callable[[], FastAPI]] = { + 'data_plane': build_data_plane_app, 'gateway': build_gateway_app, 'model': build_model_app, 'sampler': build_sampler_app, @@ -91,41 +106,96 @@ def build_processor_app() -> FastAPI: _HTTP_METHODS = {'GET', 'POST', 'PUT', 'PATCH', 'DELETE'} -def _extract_app_surface(app: FastAPI) -> dict[str, Any]: - """Return a SLIM client-contract view of ``app``'s OpenAPI surface. - - Snapshots, per path and HTTP method, only the stable client-facing contract: - the ``operationId``, the ``parameters``, and the set of response status - codes. The full ``components.schemas`` body and per-operation ``requestBody`` - schema are intentionally NOT snapshotted — they churn on Pydantic / FastAPI - version bumps without representing a real client-contract change. Route - paths, HTTP methods, and response status codes remain frozen. - """ - spec = get_openapi( - title='contract', - version='0.0.0', - routes=app.routes, - ) +def _type_contract(annotation: Any, seen: frozenset[str] = frozenset()) -> Any: + """Build a stable field-level schema for Pydantic models and SDK dataclasses.""" + if annotation is None or annotation is type(None): + return {'type': 'null'} + if annotation is Any: + return {} + origin = get_origin(annotation) + args = get_args(annotation) + if origin is Annotated: + return _type_contract(args[0], seen) + if origin in (Union, pytypes.UnionType): + return {'anyOf': [_type_contract(arg, seen) for arg in args]} + if origin in (list, set, tuple, Sequence): + return {'type': 'array', 'items': _type_contract(args[0], seen) if args else {}} + if origin in (dict, Mapping): + return {'type': 'object', 'additionalProperties': _type_contract(args[1], seen) if len(args) > 1 else {}} + if origin is Literal: + return {'enum': list(args)} + if isinstance(annotation, type) and issubclass(annotation, Enum): + return {'enum': [item.value for item in annotation]} + if isinstance(annotation, type) and issubclass(annotation, BaseModel): + return annotation.model_json_schema() + if isinstance(annotation, type) and dataclasses.is_dataclass(annotation): + name = f'{annotation.__module__}.{annotation.__qualname__}' + if name in seen: + return {'$ref': name} + module = sys.modules.get(annotation.__module__) + try: + hints = get_type_hints(annotation, globalns=vars(module) if module else None) + except (NameError, TypeError): + hints = annotation.__annotations__ + properties = {} + required = [] + for field in dataclasses.fields(annotation): + if not field.init or field.name.startswith('_'): + continue + properties[field.name] = _type_contract(hints.get(field.name, Any), seen | {name}) + if field.default is dataclasses.MISSING and field.default_factory is dataclasses.MISSING: + required.append(field.name) + result = {'type': 'object', 'properties': properties} + if required: + result['required'] = required + return result + primitive = {str: 'string', int: 'integer', float: 'number', bool: 'boolean'} + if annotation in primitive: + return {'type': primitive[annotation]} + return {'pythonType': getattr(annotation, '__qualname__', repr(annotation))} + + +def _parameter_contract(field: Any) -> dict[str, Any]: + field_info = field.field_info + return { + 'name': field.alias, + 'required': bool(field_info.is_required()), + 'schema': _type_contract(field_info.annotation), + } + +def _extract_app_surface(app: FastAPI) -> dict[str, Any]: + """Return every route's complete request and response type shape.""" paths: dict[str, dict[str, Any]] = {} - for path, ops in (spec.get('paths') or {}).items(): - clean_ops: dict[str, Any] = {} - for method, op in ops.items(): - if method.upper() not in _HTTP_METHODS: - continue - clean_ops[method.upper()] = { - 'operationId': op.get('operationId'), - 'parameters': op.get('parameters', []), - 'responses': sorted((op.get('responses') or {}).keys()), + for route in app.routes: + if not isinstance(route, APIRoute): + continue + extra_responses = {} + for status, response in route.responses.items(): + extra_responses[str(status)] = { + 'description': response.get('description'), + 'content': response.get('content'), + 'model': _type_contract(response.get('model')) if response.get('model') else None, } - if clean_ops: - paths[path] = clean_ops - + operation = { + 'operationId': route.operation_id or route.name, + 'body': [_parameter_contract(field) for field in route.dependant.body_params], + 'path': [_parameter_contract(field) for field in route.dependant.path_params], + 'query': [_parameter_contract(field) for field in route.dependant.query_params], + 'headers': [_parameter_contract(field) for field in route.dependant.header_params], + 'cookies': [_parameter_contract(field) for field in route.dependant.cookie_params], + 'response': _type_contract(route.response_model), + 'responses': extra_responses, + 'statusCode': route.status_code or 200, + } + for method in sorted(route.methods & _HTTP_METHODS): + client_path = re.sub(r'{([^}:]+):[^}]+}', r'{\1}', route.path) + paths.setdefault(client_path, {})[method] = operation return {'paths': paths} def extract_full_surface() -> dict[str, Any]: - """Build all four apps and return a per-app contract surface dict.""" + """Build all five apps and return a per-app contract surface dict.""" surface: dict[str, Any] = {} for name, builder in APP_BUILDERS.items(): app = builder() diff --git a/tests/server/contract/test_client_api_contract.py b/tests/server/contract/test_client_api_contract.py new file mode 100644 index 000000000..ea0f743b8 --- /dev/null +++ b/tests/server/contract/test_client_api_contract.py @@ -0,0 +1,43 @@ +# Copyright (c) ModelScope Contributors. All rights reserved. +"""Zero-wire-change contract guard (T8.1 / R8 / Property 10). + +Exports the request/response surface of all five apps and compares it field-by-field with +the canonical baseline. The diff must be empty. Also asserts the load-bearing +invariants: ``schedule_task_and_wait`` still +exists and the only client-side additions are ``types/base.py`` and ``types/errors.py``. +""" +from __future__ import annotations + +import pytest + +from tests.server.contract.client_api_harness import extract_full_surface, load_baseline + + +def test_wire_surface_matches_baseline(): + current = extract_full_surface() + baseline = load_baseline() + assert set(current) == {'data_plane', 'gateway', 'model', 'processor', 'sampler'} + assert current == baseline, ( + 'Client-facing wire surface changed vs the canonical baseline; ' + 'this spec must be zero-wire-change. Diffing apps: ' + f'{[a for a in set(current) | set(baseline) if current.get(a) != baseline.get(a)]}') + + +def test_schedule_task_and_wait_not_removed(): + from twinkle.server.utils.task_queue.mixin import TaskQueueMixin + assert hasattr(TaskQueueMixin, 'schedule_task_and_wait') + + +def test_new_client_types_importable(): + # The only permitted client-side additions. + import twinkle_client.types.base as base + import twinkle_client.types.errors as errors + + for symbol in ('StrictRequest', 'ResponseModel', 'DataModel', 'backend_only'): + assert hasattr(base, symbol) + for symbol in ('ErrorPayload', 'ErrorCategory', 'QueueStateLiteral'): + assert hasattr(errors, symbol) + + +if __name__ == '__main__': + raise SystemExit(pytest.main([__file__, '-v'])) diff --git a/tests/server/contract/test_error_wire.py b/tests/server/contract/test_error_wire.py new file mode 100644 index 000000000..df6be2899 --- /dev/null +++ b/tests/server/contract/test_error_wire.py @@ -0,0 +1,39 @@ +from __future__ import annotations + +from fastapi import FastAPI +from fastapi.testclient import TestClient +from tinker.types import RequestFailedResponse + +from twinkle.server.gateway.tinker_handlers import _register_tinker_routes + + +class _State: + + async def get_future(self, request_id: str): + return { + 'status': 'failed', + 'result': { + 'error': 'backend timed out', + 'category': 'server', + 'error_code': 504, + 'request_id': request_id, + }, + } + + +class _Gateway: + state = _State() + + +def test_retrieve_future_returns_parseable_error_payload(): + app = FastAPI() + _register_tinker_routes(app, lambda: _Gateway()) + + response = TestClient(app).post('/retrieve_future', json={'request_id': 'req-1'}) + + assert response.status_code == 200 + body = response.json() + assert body['error_code'] == 504 + assert body['request_id'] == 'req-1' + parsed = RequestFailedResponse.model_validate(body) + assert parsed.category.value == 'server' diff --git a/tests/server/integration/e2e_helpers.py b/tests/server/integration/e2e_helpers.py index 9d3073bad..2eca95f0d 100644 --- a/tests/server/integration/e2e_helpers.py +++ b/tests/server/integration/e2e_helpers.py @@ -22,7 +22,9 @@ MODEL_ID = f'ms://{BASE_MODEL}' BASE_URL = os.environ.get('TWINKLE_SERVER_URL', 'http://localhost:9000') API_KEY = 'EMPTY_API_KEY' -TIMEOUT = 120 # seconds per operation before declaring hang +TIMEOUT = float(os.environ.get('TWINKLE_TEST_OPERATION_TIMEOUT', '120')) +# Per-operation hang threshold. PPU Megatron cold JIT can exceed 300s; callers may +# raise this without weakening the default CI/GPU bound. GRADIENT_ACCUMULATION_STEPS = 2 # Megatron requires GA >= 2 @@ -66,6 +68,17 @@ def log(msg: str) -> None: # Dataset Factories # ═══════════════════════════════════════════════════════════════════════════ + +def _local_arrow_dataset(path: str, data_slice): + """Load selected rows from cached Arrow without hub metadata access.""" + from datasets import Dataset as HFDataset + from twinkle.dataset import Dataset, DatasetMeta + + source = HFDataset.from_file(path) + indices = [index % len(source) for index in data_slice] + return Dataset(DatasetMeta(data=source.select(indices))) + + def create_sft_dataset(data_slice=range(100)): """Create SelfCognition SFT dataset (small slice for speed).""" from twinkle.dataloader import DataLoader @@ -83,7 +96,11 @@ def create_dpo_dataset(data_slice=range(50)): from twinkle.dataset import Dataset, DatasetMeta from twinkle.preprocessor import EmojiDPOProcessor - dataset = Dataset(DatasetMeta('ms://hjh0119/shareAI-Llama3-DPO-zh-en-emoji', data_slice=data_slice)) + local_arrow = os.environ.get('TWINKLE_TEST_DPO_ARROW') + if local_arrow: + dataset = _local_arrow_dataset(local_arrow, data_slice) + else: + dataset = Dataset(DatasetMeta('ms://hjh0119/shareAI-Llama3-DPO-zh-en-emoji', data_slice=data_slice)) dataset.set_template('Qwen3_5Template', model_id=MODEL_ID, max_length=1024) dataset.map(EmojiDPOProcessor, init_args={'system': 'You are a helpful assistant.'}) dataset.encode() @@ -97,7 +114,11 @@ def create_grpo_dataset(data_slice=range(50)): system_prompt = ('You are a helpful math assistant. Solve the problem with minimal but correct reasoning ' 'and put your final answer within \\boxed{}.') - dataset = Dataset(DatasetMeta('ms://modelscope/gsm8k', subset_name='main', split='train', data_slice=data_slice)) + local_arrow = os.environ.get('TWINKLE_TEST_GRPO_ARROW') + if local_arrow: + dataset = _local_arrow_dataset(local_arrow, data_slice) + else: + dataset = Dataset(DatasetMeta('ms://modelscope/gsm8k', subset_name='main', split='train', data_slice=data_slice)) dataset.set_template('Qwen3_5Template', model_id=MODEL_ID, max_length=2048, enable_thinking=False) dataset.map(GSM8KProcessor(system=system_prompt)) dataset.encode(add_generation_prompt=True) diff --git a/tests/server/integration/test_actor_recovery.py b/tests/server/integration/test_actor_recovery.py new file mode 100644 index 000000000..64696637c --- /dev/null +++ b/tests/server/integration/test_actor_recovery.py @@ -0,0 +1,86 @@ +# Copyright (c) ModelScope Contributors. All rights reserved. +"""Post-timeout liveness probe and health status bit (T4.2 / R3#2-3). + +Binds the real ``ModelManagement`` health methods onto a minimal harness with a +toggleable mock ``ping`` and a direct ``call_backend``. No GPU/Ray/full server. +""" +from __future__ import annotations + +import pytest +from fastapi import FastAPI + +from twinkle.server.model.app import ModelManagement +from twinkle.server.model.twinkle_handlers import _register_twinkle_routes + + +class _MockModel: + + def __init__(self) -> None: + self.alive = True + + def ping(self) -> bool: + if not self.alive: + raise RuntimeError('actor unreachable (simulated)') + return True + + +class _HealthHarness: + # Reuse the real implementations under test. + _run_model_health_probe = ModelManagement._run_model_health_probe + check_model_health = ModelManagement.check_model_health + mark_unhealthy = ModelManagement.mark_unhealthy + _probe_after_timeout = ModelManagement._probe_after_timeout + + def __init__(self, model: _MockModel) -> None: + self.model = model + self._model_unhealthy = False + + async def call_backend(self, fn, /, *args, admit: bool = True, **kwargs): + return fn(*args, **kwargs) + + +@pytest.mark.asyncio +async def test_timeout_probe_marks_unhealthy_then_recovers(): + model = _MockModel() + h = _HealthHarness(model) + + # Healthy at first. + result = await h.check_model_health() + assert result['healthy'] is True + assert h._model_unhealthy is False + + # A backend timeout fires the probe while the actor is unreachable. + model.alive = False + await h._probe_after_timeout() + assert h._model_unhealthy is True # /healthz would return 503 + + # Actor recovers; one successful probe clears the bit (no restart needed). + model.alive = True + result = await h.check_model_health() + assert result['healthy'] is True + assert h._model_unhealthy is False + + +@pytest.mark.asyncio +async def test_health_route_returns_503_when_probe_fails(): + model = _MockModel() + model.alive = False + harness = _HealthHarness(model) + app = FastAPI() + _register_twinkle_routes(app, lambda: harness) + route = next(route for route in app.routes if getattr(route, 'path', None) == '/healthz') + + response = await route.endpoint(object(), harness) + + assert response.status_code == 503 + + +@pytest.mark.asyncio +async def test_mark_unhealthy_is_cleared_by_successful_probe(): + h = _HealthHarness(_MockModel()) + h.mark_unhealthy() + assert h._model_unhealthy is True + + result = await h.check_model_health() + assert result['healthy'] is True + assert h._model_unhealthy is False diff --git a/tests/server/integration/test_blocking_boundary.py b/tests/server/integration/test_blocking_boundary.py new file mode 100644 index 000000000..6432ad791 --- /dev/null +++ b/tests/server/integration/test_blocking_boundary.py @@ -0,0 +1,210 @@ +# Copyright (c) ModelScope Contributors. All rights reserved. +"""Blocking_Call_Boundary integration tests (T3.8 / R9#2 / Property 3-4). + +These exercise the real ``TaskQueueMixin.call_backend`` through a minimal harness +that sets only the two attributes it uses (a dedicated executor and the optional +Admission_Gate), constructed exactly as ``_init_task_queue`` does. The backend is a +deliberately slow plain callable -- no GPU, Megatron, or Ray involved. +""" +from __future__ import annotations + +import asyncio +import threading +import time +from concurrent.futures import ThreadPoolExecutor + +import httpx +import pytest +from fastapi import FastAPI +from fastapi.responses import JSONResponse + +ray = pytest.importorskip('ray') + +from twinkle.server.utils.task_queue.mixin import TaskQueueMixin # noqa: E402 +from twinkle.server.utils.task_queue.types import BackendBusyError # noqa: E402 + + +class _Harness(TaskQueueMixin): + """Minimal holder exposing the real call_backend with a chosen gate setting.""" + + def __init__(self, gate_enabled: bool, *, max_workers: int | None = None) -> None: + self._backend_executor = ThreadPoolExecutor( + max_workers=max_workers, thread_name_prefix='twinkle-backend') + self._backend_probe_executor = ThreadPoolExecutor(max_workers=1, thread_name_prefix='twinkle-backend-probe') + self._backend_admission = asyncio.Lock() if gate_enabled else None + self._backend_poisoned = asyncio.Event() + + def close(self) -> None: + self._backend_executor.shutdown(wait=False, cancel_futures=True) + self._backend_probe_executor.shutdown(wait=False, cancel_futures=True) + + +@pytest.mark.asyncio +async def test_healthz_style_probe_responsive_during_slow_backend(): + """Property 3: while a slow backend call is in flight, an admit=False probe + (as /healthz uses) returns well within 5 seconds.""" + h = _Harness(gate_enabled=True) + try: + slow = asyncio.create_task(h.call_backend(lambda: time.sleep(3.0))) + await asyncio.sleep(0.05) # let the slow call take the gate + a thread + + loop = asyncio.get_running_loop() + start = loop.time() + probe = await h.call_backend(lambda: 'pong', admit=False) # no gate, like the ping probe + elapsed = loop.time() - start + + assert probe == 'pong' + assert elapsed < 5.0 + await slow + finally: + h.close() + + +@pytest.mark.asyncio +async def test_normal_gate_contention_waits_instead_of_failing(): + h = _Harness(gate_enabled=True) + try: + first = asyncio.create_task(h.call_backend(lambda: (time.sleep(0.2), 'first')[1])) + await asyncio.sleep(0.05) + second = asyncio.create_task(h.call_backend(lambda: 'second')) + assert await first == 'first' + assert await second == 'second' + finally: + h.close() + + +@pytest.mark.asyncio +async def test_cancelled_gate_waiter_does_not_steal_lock(): + h = _Harness(gate_enabled=True) + release = threading.Event() + + def wait_for_release(): + while not release.is_set(): + time.sleep(0.01) + + try: + first = asyncio.create_task(h.call_backend(wait_for_release)) + await asyncio.sleep(0.05) + waiter = asyncio.create_task(h.call_backend(lambda: 'cancelled')) + await asyncio.sleep(0.05) + waiter.cancel() + with pytest.raises(asyncio.CancelledError): + await waiter + release.set() + await first + assert await h.call_backend(lambda: 'next') == 'next' + finally: + release.set() + h.close() + + +@pytest.mark.asyncio +async def test_cancelled_queued_backend_call_releases_gate(): + h = _Harness(gate_enabled=True, max_workers=1) + release_worker = threading.Event() + occupied = h._backend_executor.submit(release_worker.wait) + try: + with pytest.raises(asyncio.TimeoutError): + await asyncio.wait_for(h.call_backend(lambda: 'never-started'), timeout=0.05) + await asyncio.sleep(0) + assert not h._backend_admission.locked() + release_worker.set() + occupied.result(timeout=5) + assert await h.call_backend(lambda: 'next') == 'next' + finally: + release_worker.set() + h.close() + + +@pytest.mark.asyncio +async def test_gate_held_by_leaked_call_fast_fails_next_task(): + """Property 4 / R2#4: a call that outlives its wait_for keeps the gate; the next + admitting call fails fast with BackendBusyError instead of entering the backend.""" + h = _Harness(gate_enabled=True) + entered = {'count': 0} + + def slow(): + time.sleep(1.5) + + def would_enter_backend(): + entered['count'] += 1 + return 'should-not-run' + + try: + with pytest.raises(asyncio.TimeoutError): + await asyncio.wait_for(h.call_backend(slow), timeout=0.3) + + # The leaked thread still holds the gate. + with pytest.raises(BackendBusyError): + await h.call_backend(would_enter_backend) + assert entered['count'] == 0 # never reached the backend + + # After the leaked thread truly finishes, the gate frees on its own. + await asyncio.sleep(1.6) + assert await h.call_backend(would_enter_backend) == 'should-not-run' + assert entered['count'] == 1 + finally: + h.close() + + +@pytest.mark.asyncio +async def test_probe_times_out_while_same_serial_actor_is_busy(): + + @ray.remote + class SerialActor: + + def slow(self): + time.sleep(1.0) + + def ping(self): + return True + + started_ray = not ray.is_initialized() + if started_ray: + ray.init(num_cpus=1, logging_level='ERROR') + actor = SerialActor.remote() + h = _Harness(gate_enabled=True) + app = FastAPI() + + @app.get('/healthz') + async def healthz(): + try: + await h.call_backend(lambda: ray.get(actor.ping.remote(), timeout=0.2), admit=False) + return {'healthy': True} + except ray.exceptions.GetTimeoutError: + return JSONResponse(status_code=503, content={'healthy': False}) + + try: + slow = asyncio.create_task(h.call_backend(lambda: ray.get(actor.slow.remote(), timeout=2))) + await asyncio.sleep(0.1) + start = time.monotonic() + async with httpx.AsyncClient(transport=httpx.ASGITransport(app=app), base_url='http://test') as client: + response = await client.get('/healthz') + assert response.status_code == 503 + assert time.monotonic() - start < 5 + await slow + finally: + h.close() + ray.kill(actor) + if started_ray: + ray.shutdown() + + +@pytest.mark.asyncio +async def test_sampler_without_gate_runs_two_calls_concurrently(): + """R9#2 case 3 / opt-in: with the gate disabled (SamplerManagement), two backend + calls are in flight at once rather than serialized.""" + h = _Harness(gate_enabled=False) + try: + loop = asyncio.get_running_loop() + start = loop.time() + results = await asyncio.gather( + h.call_backend(lambda: (time.sleep(1.0), 'a')[1]), + h.call_backend(lambda: (time.sleep(1.0), 'b')[1]), + ) + elapsed = loop.time() - start + + assert sorted(results) == ['a', 'b'] + assert elapsed < 1.8 # concurrent, not ~2.0s serialized + finally: + h.close() diff --git a/tests/server/integration/test_nccl_safe_tinker_e2e.py b/tests/server/integration/test_nccl_safe_tinker_e2e.py index aa94f99c4..155f27383 100644 --- a/tests/server/integration/test_nccl_safe_tinker_e2e.py +++ b/tests/server/integration/test_nccl_safe_tinker_e2e.py @@ -1,16 +1,15 @@ # Copyright (c) ModelScope Contributors. All rights reserved. -"""Real E2E test for NCCL-safe fault tolerance via Tinker client path. +"""Real E2E test for loud-failure semantics via the Tinker client path. -Exercises the /tinker/forward_backward endpoint through the upstream Tinker SDK. -All adversarial scenarios verify that safe_loss catches errors gracefully without -NCCL hang or model state corruption. +Exercises ``/tinker/forward_backward`` through the upstream Tinker SDK. The +invariant under test (post silent-degradation removal): a request whose loss +computation fails does NOT come back as a silent zero-loss success -- it enters a +failed terminal state -- and a subsequent valid request on the same deployment +still succeeds. Prerequisites: - 1. Ray cluster running with GPUs (2 for model DP/TP, optionally 1 for sampler) - 2. Twinkle server started with TWINKLE_FAIL_FAST=0 - -Usage (direct): - python tests/server/integration/test_nccl_safe_tinker_e2e.py + 1. Ray cluster running with GPUs (2 for model DP/TP) + 2. Twinkle server started with queue_config.execution_timeout=30 Usage (pytest, requires TWINKLE_TEST_GPU_E2E=1): TWINKLE_TEST_GPU_E2E=1 pytest tests/server/integration/test_nccl_safe_tinker_e2e.py -v @@ -18,10 +17,7 @@ from __future__ import annotations import os -import sys import time -import logging -import traceback import numpy as np import pytest @@ -31,409 +27,115 @@ reason='Set TWINKLE_TEST_GPU_E2E=1 to run real GPU E2E tests (requires running server)', ) -logging.basicConfig(level=logging.INFO, format='%(asctime)s [%(levelname)s] %(message)s') -logger = logging.getLogger(__name__) - - -def log(msg): - """Print + flush to avoid log suppression by init_tinker_client().""" - print(f'[E2E-Tinker] {msg}', flush=True) - - BASE_MODEL = 'Qwen/Qwen3.5-4B' SERVER_URL = os.environ.get('TWINKLE_SERVER_URL', 'http://localhost:9000') -TIMEOUT = 120 - - -def wait_for_server(url, timeout=300): - """Wait for Twinkle server to become ready.""" - import requests - start = time.time() - while time.time() - start < timeout: - try: - resp = requests.get(f'{url}/-/routes', timeout=5) - if resp.status_code == 200: - elapsed = int(time.time() - start) - log(f'Server is ready (waited {elapsed}s)') - return True - except Exception: - pass - time.sleep(5) - raise TimeoutError(f'Server not ready after {timeout}s') +EXECUTION_TIMEOUT = float(os.environ.get('TWINKLE_TEST_EXECUTION_TIMEOUT', '30')) +TIMEOUT = EXECUTION_TIMEOUT + 15 +# The `global_rank=` attribution is added by `nccl_safe_megatron`, which decorates +# only the Megatron backend; the transformers backend carries no such annotation +# (its former silent-degradation decorator was removed by R6#3). Gate the rank-attribution assertion +# on the backend so this file is safe under TWINKLE_TEST_BACKEND=transformers. +BACKEND = os.environ.get('TWINKLE_TEST_BACKEND', 'megatron') -def init_client(): - """Initialize Tinker client and create training client.""" +def _init_client(): os.environ['TINKER_BASE_URL'] = SERVER_URL os.environ['TWINKLE_SERVER_TOKEN'] = 'EMPTY_TOKEN' - from twinkle_client import init_tinker_client init_tinker_client() - from tinker import ServiceClient - service_client = ServiceClient() - training_client = service_client.create_lora_training_client(base_model=BASE_MODEL, rank=16) - log('Training client created successfully') - return training_client + return ServiceClient().create_lora_training_client(base_model=BASE_MODEL, rank=16) -def make_datum(seq_len=32, completion_len=16, *, bad_logprobs_len=None, include_advantages=True): - """Construct a Datum for GRPO training.""" +def _make_datum(seq_len=64, completion_len=32, *, bad_logprobs_len=None): from tinker import types - prompt_len = seq_len - completion_len input_tokens = list(range(1, seq_len + 1)) target_tokens = [0] * prompt_len + list(range(100, 100 + completion_len)) weights = [0] * prompt_len + [1] * completion_len - - if bad_logprobs_len is not None: - logprobs_values = np.random.randn(bad_logprobs_len).astype(np.float32) - padded_logprobs = [0.0] * prompt_len + logprobs_values.tolist() - else: - logprobs_values = np.random.randn(completion_len).astype(np.float32) - padded_logprobs = [0.0] * prompt_len + logprobs_values.tolist() - - loss_fn_inputs = { - 'target_tokens': target_tokens, - 'weights': weights, - 'logprobs': types.TensorData.from_numpy(np.array(padded_logprobs, dtype=np.float32)), - } - - if include_advantages: - advantage = float(np.random.randn()) - padded_advantages = [0.0] * prompt_len + [advantage] * completion_len - loss_fn_inputs['advantages'] = types.TensorData.from_numpy( - np.array(padded_advantages, dtype=np.float32)) - + n = bad_logprobs_len if bad_logprobs_len is not None else completion_len + padded_logprobs = [0.0] * prompt_len + np.random.randn(n).astype(np.float32).tolist() + advantage = float(np.random.randn()) return types.Datum( model_input=types.ModelInput.from_ints(input_tokens), - loss_fn_inputs=loss_fn_inputs, + loss_fn_inputs={ + 'target_tokens': target_tokens, + 'weights': weights, + 'logprobs': types.TensorData.from_numpy(np.array(padded_logprobs, dtype=np.float32)), + 'advantages': types.TensorData.from_numpy( + np.array([0.0] * prompt_len + [advantage] * completion_len, dtype=np.float32)), + }, ) -def run_forward_backward(training_client, datums, test_name, expect_success=True): - """Run forward_backward and return (success, result, elapsed_seconds).""" - log(f'[{test_name}] Sending {len(datums)} datums...') - start = time.time() - try: - result = training_client.forward_backward(datums, 'importance_sampling').result() - elapsed = time.time() - start - log(f'[{test_name}] Completed in {elapsed:.1f}s') - if hasattr(result, 'metrics') and result.metrics: - loss_avg = result.metrics.get('loss:avg', 'N/A') - log(f'[{test_name}] loss:avg = {loss_avg}') - return True, result, elapsed - except Exception as e: - elapsed = time.time() - start - log(f'[{test_name}] FAILED in {elapsed:.1f}s: {type(e).__name__}: {e}') - if elapsed > TIMEOUT: - log(f'[{test_name}] TIMEOUT! This suggests NCCL hang!') - return False, None, elapsed - - -def do_optim_step(training_client, test_name): - """Run optimizer step.""" - from tinker import types - try: - training_client.optim_step(types.AdamParams(learning_rate=1e-5)).result() - log(f'[{test_name}] optim_step OK') - return True - except Exception as e: - log(f'[{test_name}] optim_step FAILED: {e}') - return False - - -# ═══════════════════════════════════════════════════════════════════════════ -# Test Scenarios (19 tests) -# ═══════════════════════════════════════════════════════════════════════════ - -def test_1_normal_grpo(tc): - datums = [make_datum(seq_len=64, completion_len=32) for _ in range(4)] - ok, result, elapsed = run_forward_backward(tc, datums, 'TEST-1-NORMAL') - assert ok and elapsed < TIMEOUT - do_optim_step(tc, 'TEST-1-NORMAL') - return True - -def test_2_bad_old_logps(tc): - datums = [ - make_datum(seq_len=64, completion_len=32), - make_datum(seq_len=64, completion_len=32, bad_logprobs_len=5), - make_datum(seq_len=64, completion_len=32), - make_datum(seq_len=64, completion_len=32, bad_logprobs_len=99), - ] - ok, result, elapsed = run_forward_backward(tc, datums, 'TEST-2-BAD-LOGPS') - if not ok: - return elapsed < TIMEOUT - assert elapsed < TIMEOUT - do_optim_step(tc, 'TEST-2-BAD-LOGPS') - return True - -def test_3_recovery(tc): - datums = [make_datum(seq_len=64, completion_len=32) for _ in range(4)] - ok, _, elapsed = run_forward_backward(tc, datums, 'TEST-3-RECOVERY') - assert ok and elapsed < TIMEOUT - do_optim_step(tc, 'TEST-3-RECOVERY') - return True - -def test_4_no_advantages(tc): - datums = [make_datum(seq_len=64, completion_len=32, include_advantages=False) for _ in range(4)] - ok, _, elapsed = run_forward_backward(tc, datums, 'TEST-4-NO-ADV') - assert ok and elapsed < TIMEOUT - do_optim_step(tc, 'TEST-4-NO-ADV') - return True +def test_failure_is_terminal_then_valid_request_succeeds(): + """A malformed request fails loudly (terminal), a subsequent valid one succeeds. -def test_5_consecutive_bad(tc): - for i in range(5): - datums = [make_datum(seq_len=64, completion_len=32, bad_logprobs_len=3+i) for _ in range(4)] - _, _, elapsed = run_forward_backward(tc, datums, f'TEST-5-{i+1}') - if elapsed >= TIMEOUT: - return False - do_optim_step(tc, f'TEST-5-{i+1}') - return True - -def test_6_nan_logprobs(tc): + Replaces the former assertion "failure degraded to zero loss and training + continued". If the recovery request does not reach a terminal success, that is + recorded as evidence that R3#2-3 actor recovery and R2#3-4 admission gate are + necessary, not optional. + """ from tinker import types - datums = [] - for _ in range(4): - d = make_datum(seq_len=64, completion_len=32) - d.loss_fn_inputs['logprobs'] = types.TensorData.from_numpy( - np.array([float('nan')] * 64, dtype=np.float32)) - datums.append(d) - _, _, elapsed = run_forward_backward(tc, datums, 'TEST-6-NAN') - if elapsed >= TIMEOUT: - return False - do_optim_step(tc, 'TEST-6-NAN') - return True + from tinker._exceptions import RequestFailedError + tc = _init_client() -def test_7_inf_logprobs(tc): - from tinker import types - datums = [] - for _ in range(4): - d = make_datum(seq_len=64, completion_len=32) - inf_arr = np.full(64, float('inf'), dtype=np.float32) - inf_arr[::2] = float('-inf') - d.loss_fn_inputs['logprobs'] = types.TensorData.from_numpy(inf_arr) - datums.append(d) - _, _, elapsed = run_forward_backward(tc, datums, 'TEST-7-INF') - if elapsed >= TIMEOUT: - return False - do_optim_step(tc, 'TEST-7-INF') - return True + # Deliberately malformed: logprobs length inconsistent with the completion. + bad = [_make_datum(bad_logprobs_len=5) for _ in range(4)] + start = time.time() + with pytest.raises(RequestFailedError) as caught: + tc.forward_backward(bad, 'importance_sampling').result(timeout=TIMEOUT) + assert caught.value.category is types.RequestErrorCategory.Server + assert time.time() - start < TIMEOUT, 'malformed request must fail fast, not hang (NCCL)' + + # Recovery: megatron commits the DDP reducer inside its fused forward_backward, so + # a subsequent valid request must succeed. The tinker transformers path runs + # forward()/loss/backward() separately; a mid-iteration loss failure leaves DDP's + # reducer half-finished and poisons the next request, so the spec only guarantees a + # *terminal* response there (R6#14), not success. (tinker 0.16.1 exposes no GA knob.) + good = [_make_datum() for _ in range(4)] + if BACKEND == 'megatron': + result = tc.forward_backward(good, 'importance_sampling').result() + assert result is not None + tc.optim_step(types.AdamParams(learning_rate=1e-5)).result() + else: + try: + result = tc.forward_backward(good, 'importance_sampling').result(timeout=TIMEOUT) + assert result is not None + tc.optim_step(types.AdamParams(learning_rate=1e-5)).result() + except RequestFailedError as exc: + assert exc.category is types.RequestErrorCategory.Server -def test_8_extreme_advantages(tc): - from tinker import types - datums = [] - for i in range(4): - d = make_datum(seq_len=64, completion_len=32) - val = 1e30 if i % 2 == 0 else -1e30 - adv = np.full(64, 0.0, dtype=np.float32) - adv[32:] = val - d.loss_fn_inputs['advantages'] = types.TensorData.from_numpy(adv) - datums.append(d) - _, _, elapsed = run_forward_backward(tc, datums, 'TEST-8-EXTREME-ADV') - if elapsed >= TIMEOUT: - return False - do_optim_step(tc, 'TEST-8-EXTREME-ADV') - return True -def test_9_zero_completion(tc): +def test_partial_rank_failure_is_terminal_then_recovers(): from tinker import types - datums = [] - for _ in range(4): - d = types.Datum( - model_input=types.ModelInput.from_ints(list(range(1, 65))), - loss_fn_inputs={ - 'target_tokens': [0]*64, 'weights': [0]*64, - 'logprobs': types.TensorData.from_numpy(np.zeros(64, dtype=np.float32)), - 'advantages': types.TensorData.from_numpy(np.zeros(64, dtype=np.float32)), - }, - ) - datums.append(d) - _, _, elapsed = run_forward_backward(tc, datums, 'TEST-9-ZERO-COMPL') - if elapsed >= TIMEOUT: - return False - do_optim_step(tc, 'TEST-9-ZERO-COMPL') - return True - -def test_10_partial_advantages(tc): - datums = [ - make_datum(seq_len=64, completion_len=32, include_advantages=True), - make_datum(seq_len=64, completion_len=32, include_advantages=False), - make_datum(seq_len=64, completion_len=32, include_advantages=True), - make_datum(seq_len=64, completion_len=32, include_advantages=False), - ] - _, _, elapsed = run_forward_backward(tc, datums, 'TEST-10-PARTIAL-ADV') - if elapsed >= TIMEOUT: - return False - do_optim_step(tc, 'TEST-10-PARTIAL-ADV') - return True + from tinker._exceptions import RequestFailedError + tc = _init_client() -def test_11_mixed_seq_lengths(tc): - datums = [ - make_datum(seq_len=32, completion_len=16), - make_datum(seq_len=128, completion_len=64), - make_datum(seq_len=48, completion_len=24), - make_datum(seq_len=96, completion_len=48), - ] - _, _, elapsed = run_forward_backward(tc, datums, 'TEST-11-MIXED') - if elapsed >= TIMEOUT: - return False - do_optim_step(tc, 'TEST-11-MIXED') - return True - -def test_12_all_bad(tc): - datums = [make_datum(seq_len=64, completion_len=32, bad_logprobs_len=i) for i in range(4)] - _, _, elapsed = run_forward_backward(tc, datums, 'TEST-12-ALL-BAD') - if elapsed >= TIMEOUT: - return False - do_optim_step(tc, 'TEST-12-ALL-BAD') - return True - -def test_13_forward_only_then_train(tc): - datums_infer = [make_datum(seq_len=64, completion_len=32, include_advantages=False) for _ in range(4)] + batch = [_make_datum() for _ in range(4)] + batch[0] = _make_datum(bad_logprobs_len=5) start = time.time() - try: - tc.forward(datums_infer).result() - except Exception: - if time.time() - start >= TIMEOUT: - return False - datums_train = [make_datum(seq_len=64, completion_len=32) for _ in range(4)] - ok, _, elapsed = run_forward_backward(tc, datums_train, 'TEST-13-TRAIN') - if not ok or elapsed >= TIMEOUT: - return False - do_optim_step(tc, 'TEST-13-TRAIN') - return True - -def test_14_rapid_bad_good(tc): - for i in range(5): - bad = [make_datum(seq_len=64, completion_len=32, bad_logprobs_len=i+1) for _ in range(4)] - _, _, elapsed = run_forward_backward(tc, bad, f'TEST-14-BAD-{i+1}') - if elapsed >= TIMEOUT: - return False - do_optim_step(tc, f'TEST-14-BAD-{i+1}') - good = [make_datum(seq_len=64, completion_len=32) for _ in range(4)] - ok, _, elapsed = run_forward_backward(tc, good, f'TEST-14-GOOD-{i+1}') - if not ok or elapsed >= TIMEOUT: - return False - do_optim_step(tc, f'TEST-14-GOOD-{i+1}') - return True - -def test_15_final_health(tc): - datums = [make_datum(seq_len=64, completion_len=32) for _ in range(4)] - ok, _, elapsed = run_forward_backward(tc, datums, 'TEST-15-FINAL') - assert ok and elapsed < TIMEOUT - do_optim_step(tc, 'TEST-15-FINAL') - return True - -def test_16_large_batch(tc): - datums = [make_datum(seq_len=64, completion_len=32) for _ in range(16)] - ok, _, elapsed = run_forward_backward(tc, datums, 'TEST-16-LARGE') - if elapsed >= TIMEOUT: - return False - assert ok - do_optim_step(tc, 'TEST-16-LARGE') - return True - -def test_17_single_datum(tc): - # With dp_size=2 + nproc_per_node=2, minimum batch must be >= data_world_size - datums = [make_datum(seq_len=64, completion_len=32) for _ in range(4)] - ok, _, elapsed = run_forward_backward(tc, datums, 'TEST-17-SMALL') - if elapsed >= TIMEOUT: - return False - assert ok - do_optim_step(tc, 'TEST-17-SMALL') - return True - -def test_18_save_after_error(tc): - bad = [make_datum(seq_len=64, completion_len=32, bad_logprobs_len=2) for _ in range(4)] - _, _, elapsed = run_forward_backward(tc, bad, 'TEST-18-ERR') - if elapsed >= TIMEOUT: - return False - do_optim_step(tc, 'TEST-18-ERR') - try: - tc.save_weights_for_sampler().result() - except Exception: - pass - good = [make_datum(seq_len=64, completion_len=32) for _ in range(4)] - ok, _, elapsed = run_forward_backward(tc, good, 'TEST-18-POST') - if not ok or elapsed >= TIMEOUT: - return False - do_optim_step(tc, 'TEST-18-POST') - return True - -def test_19_consecutive_optim_steps(tc): - datums = [make_datum(seq_len=64, completion_len=32) for _ in range(4)] - ok, _, elapsed = run_forward_backward(tc, datums, 'TEST-19-BASE') - assert ok and elapsed < TIMEOUT - for i in range(3): - do_optim_step(tc, f'TEST-19-STEP-{i+1}') - datums = [make_datum(seq_len=64, completion_len=32) for _ in range(4)] - ok, _, elapsed = run_forward_backward(tc, datums, 'TEST-19-VERIFY') - if not ok or elapsed >= TIMEOUT: - return False - do_optim_step(tc, 'TEST-19-VERIFY') - return True - - -ALL_TESTS = [ - ('TEST-1: Normal GRPO Training', test_1_normal_grpo), - ('TEST-2: Bad old_logps (original bug)', test_2_bad_old_logps), - ('TEST-3: Recovery after error', test_3_recovery), - ('TEST-4: No advantages (zero loss)', test_4_no_advantages), - ('TEST-5: Consecutive bad batches', test_5_consecutive_bad), - ('TEST-6: NaN logprobs', test_6_nan_logprobs), - ('TEST-7: +Inf/-Inf logprobs', test_7_inf_logprobs), - ('TEST-8: Extreme advantages (1e30)', test_8_extreme_advantages), - ('TEST-9: Zero completion tokens', test_9_zero_completion), - ('TEST-10: Partial advantages (ragged)', test_10_partial_advantages), - ('TEST-11: Mixed sequence lengths', test_11_mixed_seq_lengths), - ('TEST-12: All datums bad (100%)', test_12_all_bad), - ('TEST-13: forward_only then train', test_13_forward_only_then_train), - ('TEST-14: Rapid bad->good alternation', test_14_rapid_bad_good), - ('TEST-15: Final health check', test_15_final_health), - ('TEST-16: Large batch (16 datums)', test_16_large_batch), - ('TEST-17: Single datum batch', test_17_single_datum), - ('TEST-18: Save after error', test_18_save_after_error), - ('TEST-19: Consecutive optim_steps', test_19_consecutive_optim_steps), -] - - -def main(): - log('=' * 60) - log('NCCL-Safe E2E Test - Tinker Client Path') - log('=' * 60) - log(f'Server URL: {SERVER_URL}') - log(f'Base Model: {BASE_MODEL}') - log(f'TWINKLE_FAIL_FAST = {os.getenv("TWINKLE_FAIL_FAST", "1 (default)")}') - - wait_for_server(SERVER_URL) - tc = init_client() - - results = [] - for name, test_fn in ALL_TESTS: - log(f'\n{"=" * 60}\n{name}\n{"=" * 60}') + with pytest.raises(RequestFailedError) as caught: + tc.forward_backward(batch, 'importance_sampling').result(timeout=TIMEOUT) + assert caught.value.category is types.RequestErrorCategory.Server + # Megatron attributes the failure to a global rank via nccl_safe_megatron; the + # transformers backend has no such annotation (R6#3 removed its old decorator). + if BACKEND == 'megatron': + assert 'global_rank=' in str(caught.value) + assert time.time() - start < TIMEOUT + + # See the recovery note above: success is required only where forward_backward + # commits the DDP reducer atomically (megatron). On transformers a terminal + # loud failure is acceptable (spec R6#14) -- the guarantee is no hang. + if BACKEND == 'megatron': + result = tc.forward_backward([_make_datum() for _ in range(4)], 'importance_sampling').result() + assert result is not None + tc.optim_step(types.AdamParams(learning_rate=1e-5)).result() + else: try: - passed = test_fn(tc) - results.append((name, 'PASS' if passed else 'FAIL')) - log(f'[{name}] {"PASS" if passed else "FAIL"}') - except Exception as e: - log(f'{name}: EXCEPTION: {e}') - traceback.print_exc() - results.append((name, 'FAIL')) - - log(f'\n{"=" * 60}\nRESULTS SUMMARY\n{"=" * 60}') - all_passed = all(s == 'PASS' for _, s in results) - for name, status in results: - log(f' [{status}] {name}') - log(f'\n{"ALL" if all_passed else "SOME"} {len(results)} TESTS {"PASSED" if all_passed else "FAILED"}!') - return 0 if all_passed else 1 - - -def test_nccl_safe_tinker_e2e(): - """Pytest-collected entry point.""" - rc = main() - assert rc == 0, 'Some Tinker NCCL-safe E2E tests failed' - - -if __name__ == '__main__': - sys.exit(main()) + result = tc.forward_backward([_make_datum() for _ in range(4)], + 'importance_sampling').result(timeout=TIMEOUT) + assert result is not None + tc.optim_step(types.AdamParams(learning_rate=1e-5)).result() + except RequestFailedError as exc: + assert exc.category is types.RequestErrorCategory.Server diff --git a/tests/server/integration/test_nccl_safe_twinkle_e2e.py b/tests/server/integration/test_nccl_safe_twinkle_e2e.py index c9cce48a6..29b89be40 100644 --- a/tests/server/integration/test_nccl_safe_twinkle_e2e.py +++ b/tests/server/integration/test_nccl_safe_twinkle_e2e.py @@ -1,16 +1,14 @@ # Copyright (c) ModelScope Contributors. All rights reserved. -"""Real E2E test for NCCL-safe fault tolerance via Twinkle client path. +"""Real E2E test for loud-failure semantics via the Twinkle-native client path. -Exercises the /twinkle/forward_backward endpoint through the Twinkle SDK -(init_twinkle_client + MultiLoraTransformersModel). This is a SEPARATE code -path from the Tinker SDK (/tinker/forward_backward). +Exercises ``/twinkle/forward_backward`` through the Twinkle client. The invariant +under test (post silent-degradation removal): a request whose loss computation fails +does NOT come back as a silent zero-loss success -- it fails loudly -- and a +subsequent valid request on the same deployment still succeeds. Prerequisites: - 1. Ray cluster running with GPUs (2 for model DP/TP, optionally 1 for sampler) - 2. Twinkle server started with TWINKLE_FAIL_FAST=0 - -Usage (direct): - python tests/server/integration/test_nccl_safe_twinkle_e2e.py + 1. Ray cluster running with GPUs (2 for model DP/TP) + 2. Twinkle server started with queue_config.execution_timeout=30 Usage (pytest, requires TWINKLE_TEST_GPU_E2E=1): TWINKLE_TEST_GPU_E2E=1 pytest tests/server/integration/test_nccl_safe_twinkle_e2e.py -v @@ -18,11 +16,7 @@ from __future__ import annotations import os -import sys import time -import logging -import traceback -from typing import Any, Dict, List import numpy as np import pytest @@ -32,314 +26,83 @@ reason='Set TWINKLE_TEST_GPU_E2E=1 to run real GPU E2E tests (requires running server)', ) -logging.basicConfig(level=logging.INFO, format='%(asctime)s [%(levelname)s] %(message)s') -logger = logging.getLogger(__name__) - - -def log(msg): - print(f'[E2E-Twinkle] {msg}', flush=True) - - BASE_MODEL = 'Qwen/Qwen3.5-4B' SERVER_URL = os.environ.get('TWINKLE_SERVER_URL', 'http://localhost:9000') -TIMEOUT = 120 -ADAPTER_NAME = 'nccl-safe-test' - - -def wait_for_server(url, timeout=300): - """Wait for Twinkle server to become ready.""" - import requests - start = time.time() - while time.time() - start < timeout: - try: - resp = requests.get(f'{url}/-/routes', timeout=5) - if resp.status_code == 200: - log(f'Server is ready (waited {int(time.time() - start)}s)') - return True - except Exception: - pass - time.sleep(5) - raise TimeoutError(f'Server not ready after {timeout}s') - - -def init_client(): - """Initialize Twinkle client and configure model for GRPO training.""" +EXECUTION_TIMEOUT = float(os.environ.get('TWINKLE_TEST_EXECUTION_TIMEOUT', '30')) +TIMEOUT = EXECUTION_TIMEOUT + 15 +ADAPTER_NAME = 'loud-failure-test' +# The `global_rank=` attribution is added by `nccl_safe_megatron`, which decorates +# only the Megatron backend; the transformers backend's forward_backward carries no +# such annotation (its former silent-degradation decorator was removed by R6#3). Gate the +# rank-attribution assertion on the backend so this file is safe to run under the +# integration-e2e SKILL's TWINKLE_TEST_BACKEND=transformers path. +BACKEND = os.environ.get('TWINKLE_TEST_BACKEND', 'megatron') + + +def _init_client(): + from peft import LoraConfig from twinkle_client import init_twinkle_client from twinkle_client.model import MultiLoraTransformersModel - from peft import LoraConfig init_twinkle_client(base_url=SERVER_URL, api_key='EMPTY_TOKEN') - model = MultiLoraTransformersModel(model_id=f'ms://{BASE_MODEL}') model.add_adapter_to_model( adapter_name=ADAPTER_NAME, config=LoraConfig(r=16, target_modules=['q_proj', 'v_proj']), - gradient_accumulation_steps=1, + # GA>=2 (repo convention, see e2e_helpers): with GA=1 every backward syncs + # DDP immediately, so a mid-iteration failure can leave the reducer + # half-finished and poison the next request. GA=2 runs accumulation steps + # under no_sync, keeping the recovery request clean. + gradient_accumulation_steps=2, ) model.set_loss('GRPOLoss', init_args={'epsilon': 0.2}) model.set_optimizer('Adam', lr=1e-5) model.set_template('Qwen3_5Template') model.set_processor('InputProcessor', padding_side='right') - log('Twinkle client + model configured successfully') return model -def make_input_features( - batch_size=4, seq_len=64, completion_len=32, *, - bad_old_logps_len=None, include_advantages=True, - nan_old_logps=False, extreme_advantages=None, all_labels_masked=False, -): - """Construct InputFeature list + old_logps + advantages for GRPO.""" +def _make_inputs(batch_size=4, seq_len=64, completion_len=32, *, bad_old_logps_len=None): prompt_len = seq_len - completion_len - input_features = [] - old_logps_list = [] - advantages_list = [] - - for i in range(batch_size): - input_ids = list(range(1, seq_len + 1)) - labels = [-100] * seq_len if all_labels_masked else ( - [-100] * prompt_len + list(range(100, 100 + completion_len))) - input_features.append({ - 'input_ids': input_ids, - 'labels': labels, + features, old_logps, advantages = [], [], [] + for _ in range(batch_size): + features.append({ + 'input_ids': list(range(1, seq_len + 1)), + 'labels': [-100] * prompt_len + list(range(100, 100 + completion_len)), 'attention_mask': [1] * seq_len, 'position_ids': list(range(seq_len)), }) + n = bad_old_logps_len if bad_old_logps_len is not None else completion_len + old_logps.append(np.random.randn(n).tolist()) + advantages.append(float(np.random.randn())) + return features, old_logps, advantages - if bad_old_logps_len is not None: - logps = np.random.randn(bad_old_logps_len).tolist() - elif nan_old_logps: - logps = [float('nan')] * completion_len - else: - logps = np.random.randn(completion_len).tolist() - old_logps_list.append(logps) - - if extreme_advantages is not None: - advantages_list.append(extreme_advantages if i % 2 == 0 else -extreme_advantages) - else: - advantages_list.append(float(np.random.randn())) - - old_logps = old_logps_list if include_advantages else None - advantages = advantages_list if include_advantages else None - return input_features, old_logps, advantages - - -def run_forward_backward(model, inputs, old_logps, advantages, test_name): - """Run forward_backward and return (success, result, elapsed_seconds).""" - log(f'[{test_name}] Sending {len(inputs)} input features...') - start = time.time() - try: - kwargs: Dict[str, Any] = {} - if old_logps is not None: - kwargs['old_logps'] = old_logps - if advantages is not None: - kwargs['advantages'] = advantages - - result = model.forward_backward(inputs=inputs, **kwargs) - elapsed = time.time() - start - log(f'[{test_name}] Completed in {elapsed:.1f}s') - if hasattr(result, 'result') and result.result is not None: - log(f'[{test_name}] result = {result.result}') - return True, result, elapsed - except Exception as e: - elapsed = time.time() - start - log(f'[{test_name}] FAILED in {elapsed:.1f}s: {type(e).__name__}: {e}') - if elapsed > TIMEOUT: - log(f'[{test_name}] TIMEOUT! This suggests NCCL hang!') - return False, None, elapsed - - -def do_optim_step(model, test_name): - """Run clip_grad_and_step.""" - try: - model.clip_grad_and_step() - log(f'[{test_name}] clip_grad_and_step OK') - return True - except Exception as e: - log(f'[{test_name}] clip_grad_and_step FAILED: {e}') - return False +def test_failure_is_terminal_then_valid_request_succeeds(): + """A malformed request fails loudly, a subsequent valid one succeeds. -# ═══════════════════════════════════════════════════════════════════════════ -# Test Scenarios (12 tests) -# ═══════════════════════════════════════════════════════════════════════════ + Replaces the former assertion "failure degraded to zero loss and training + continued". If the recovery request does not reach a terminal success, that is + recorded as evidence that R3#2-3 actor recovery and R2#3-4 admission gate are + necessary, not optional. + """ + model = _init_client() -def test_1_normal_grpo(m): - inputs, old_logps, adv = make_input_features(batch_size=4) - ok, _, elapsed = run_forward_backward(m, inputs, old_logps, adv, 'TEST-1-NORMAL') - assert ok and elapsed < TIMEOUT - do_optim_step(m, 'TEST-1-NORMAL') - return True - -def test_2_bad_old_logps(m): - inputs, old_logps, adv = make_input_features(batch_size=4, bad_old_logps_len=5) - ok, _, elapsed = run_forward_backward(m, inputs, old_logps, adv, 'TEST-2-BAD-LOGPS') - if not ok: - return elapsed < TIMEOUT - assert elapsed < TIMEOUT - do_optim_step(m, 'TEST-2-BAD-LOGPS') - return True - -def test_3_recovery(m): - inputs, old_logps, adv = make_input_features(batch_size=4) - ok, _, elapsed = run_forward_backward(m, inputs, old_logps, adv, 'TEST-3-RECOVERY') - assert ok and elapsed < TIMEOUT - do_optim_step(m, 'TEST-3-RECOVERY') - return True - -def test_4_nan_old_logps(m): - inputs, old_logps, adv = make_input_features(batch_size=4, nan_old_logps=True) - _, _, elapsed = run_forward_backward(m, inputs, old_logps, adv, 'TEST-4-NAN') - if elapsed >= TIMEOUT: - return False - do_optim_step(m, 'TEST-4-NAN') - return True - -def test_5_extreme_advantages(m): - inputs, old_logps, adv = make_input_features(batch_size=4, extreme_advantages=1e30) - _, _, elapsed = run_forward_backward(m, inputs, old_logps, adv, 'TEST-5-EXTREME') - if elapsed >= TIMEOUT: - return False - do_optim_step(m, 'TEST-5-EXTREME') - return True - -def test_6_all_labels_masked(m): - inputs, old_logps, adv = make_input_features(batch_size=4, all_labels_masked=True) - _, _, elapsed = run_forward_backward(m, inputs, old_logps, adv, 'TEST-6-MASKED') - if elapsed >= TIMEOUT: - return False - do_optim_step(m, 'TEST-6-MASKED') - return True - -def test_7_consecutive_bad(m): - for i in range(5): - inputs, old_logps, adv = make_input_features(batch_size=4, bad_old_logps_len=i+1) - _, _, elapsed = run_forward_backward(m, inputs, old_logps, adv, f'TEST-7-{i+1}') - if elapsed >= TIMEOUT: - return False - do_optim_step(m, f'TEST-7-{i+1}') - return True - -def test_8_rapid_bad_good(m): - for i in range(5): - bad_in, bad_lp, bad_adv = make_input_features(batch_size=4, bad_old_logps_len=i+1) - _, _, elapsed = run_forward_backward(m, bad_in, bad_lp, bad_adv, f'TEST-8-BAD-{i+1}') - if elapsed >= TIMEOUT: - return False - do_optim_step(m, f'TEST-8-BAD-{i+1}') - good_in, good_lp, good_adv = make_input_features(batch_size=4) - ok, _, elapsed = run_forward_backward(m, good_in, good_lp, good_adv, f'TEST-8-GOOD-{i+1}') - if not ok or elapsed >= TIMEOUT: - return False - do_optim_step(m, f'TEST-8-GOOD-{i+1}') - return True - -def test_9_final_health(m): - inputs, old_logps, adv = make_input_features(batch_size=4) - ok, _, elapsed = run_forward_backward(m, inputs, old_logps, adv, 'TEST-9-FINAL') - assert ok and elapsed < TIMEOUT - do_optim_step(m, 'TEST-9-FINAL') - return True - -def test_10_gradient_accumulation_error(m): - inputs, lp, adv = make_input_features(batch_size=4) - ok, _, elapsed = run_forward_backward(m, inputs, lp, adv, 'TEST-10-GA1') - if not ok or elapsed >= TIMEOUT: - return False - bad_in, bad_lp, bad_adv = make_input_features(batch_size=4, bad_old_logps_len=3) - _, _, elapsed = run_forward_backward(m, bad_in, bad_lp, bad_adv, 'TEST-10-GA2-BAD') - if elapsed >= TIMEOUT: - return False - inputs, lp, adv = make_input_features(batch_size=4) - ok, _, elapsed = run_forward_backward(m, inputs, lp, adv, 'TEST-10-GA3') - if not ok or elapsed >= TIMEOUT: - return False - do_optim_step(m, 'TEST-10-GA') - return True - -def test_11_forward_only_then_train(m): - inputs, _, _ = make_input_features(batch_size=4, include_advantages=False) + bad_features, bad_old_logps, bad_adv = _make_inputs(bad_old_logps_len=5) start = time.time() - try: - m.forward_only(inputs=inputs) - except Exception: - if time.time() - start >= TIMEOUT: - return False - train_in, lp, adv = make_input_features(batch_size=4) - ok, _, elapsed = run_forward_backward(m, train_in, lp, adv, 'TEST-11-TRAIN') - if not ok or elapsed >= TIMEOUT: - return False - do_optim_step(m, 'TEST-11-TRAIN') - return True - -def test_12_mixed_seq_lengths(m): - all_inputs, all_lp, all_adv = [], [], [] - for sl, cl in [(32, 16), (128, 64), (48, 24), (96, 48)]: - feats, lp, adv = make_input_features(batch_size=1, seq_len=sl, completion_len=cl) - all_inputs.extend(feats) - if lp: - all_lp.extend(lp) - if adv: - all_adv.extend(adv) - _, _, elapsed = run_forward_backward(m, all_inputs, all_lp, all_adv, 'TEST-12-MIXED') - if elapsed >= TIMEOUT: - return False - do_optim_step(m, 'TEST-12-MIXED') - return True - - -ALL_TESTS = [ - ('TEST-1: Normal GRPO Training', test_1_normal_grpo), - ('TEST-2: Bad old_logps (original bug)', test_2_bad_old_logps), - ('TEST-3: Recovery after error', test_3_recovery), - ('TEST-4: NaN old_logps', test_4_nan_old_logps), - ('TEST-5: Extreme advantages (1e30)', test_5_extreme_advantages), - ('TEST-6: All labels masked (-100)', test_6_all_labels_masked), - ('TEST-7: Consecutive bad batches', test_7_consecutive_bad), - ('TEST-8: Rapid bad->good', test_8_rapid_bad_good), - ('TEST-9: Final health check', test_9_final_health), - ('TEST-10: Gradient accumulation error', test_10_gradient_accumulation_error), - ('TEST-11: forward_only then train', test_11_forward_only_then_train), - ('TEST-12: Mixed sequence lengths', test_12_mixed_seq_lengths), -] - - -def main(): - log('=' * 60) - log('NCCL-Safe E2E Test - Twinkle Client Path') - log('=' * 60) - log(f'Server URL: {SERVER_URL}') - log(f'Base Model: {BASE_MODEL}') - log(f'TWINKLE_FAIL_FAST = {os.getenv("TWINKLE_FAIL_FAST", "1 (default)")}') - - wait_for_server(SERVER_URL) - m = init_client() - - results = [] - for name, test_fn in ALL_TESTS: - log(f'\n{"=" * 60}\n{name}\n{"=" * 60}') - try: - passed = test_fn(m) - results.append((name, 'PASS' if passed else 'FAIL')) - log(f'[{name}] {"PASS" if passed else "FAIL"}') - except Exception as e: - log(f'{name}: EXCEPTION: {e}') - traceback.print_exc() - results.append((name, 'FAIL')) - - log(f'\n{"=" * 60}\nRESULTS SUMMARY\n{"=" * 60}') - all_passed = all(s == 'PASS' for _, s in results) - for name, status in results: - log(f' [{status}] {name}') - log(f'\n{"ALL" if all_passed else "SOME"} {len(results)} TESTS {"PASSED" if all_passed else "FAILED"}!') - return 0 if all_passed else 1 - - -def test_nccl_safe_twinkle_e2e(): - """Pytest-collected entry point.""" - rc = main() - assert rc == 0, 'Some Twinkle NCCL-safe E2E tests failed' - - -if __name__ == '__main__': - sys.exit(main()) + with pytest.raises(Exception) as caught: + model.forward_backward( + inputs=bad_features, adapter_name=ADAPTER_NAME, old_logps=bad_old_logps, advantages=bad_adv) + message = str(caught.value) + # The failure must be loud and descriptive (not a silent zero-loss success): + # the deliberate old_logps/completion length mismatch surfaces on both backends. + assert 'mismatch' in message, message + # Megatron additionally attributes the failure to a global rank via nccl_safe_megatron. + if BACKEND == 'megatron': + assert 'global_rank=' in message, message + assert time.time() - start < TIMEOUT, 'malformed request must fail fast, not hang (NCCL)' + + good_features, good_old_logps, good_adv = _make_inputs() + result = model.forward_backward( + inputs=good_features, adapter_name=ADAPTER_NAME, old_logps=good_old_logps, advantages=good_adv) + assert result is not None diff --git a/tests/server/model/test_replica_lifecycle.py b/tests/server/model/test_replica_lifecycle.py index e42193c35..84b961391 100644 --- a/tests/server/model/test_replica_lifecycle.py +++ b/tests/server/model/test_replica_lifecycle.py @@ -12,12 +12,17 @@ class _CapacityState: def __init__(self) -> None: self.capacities: dict[str, int] = {} + self.last_seen: set[str] = set() async def register_replica(self, replica_id: str, max_loras: int) -> None: self.capacities[replica_id] = max_loras async def unregister_replica(self, replica_id: str) -> None: self.capacities.pop(replica_id, None) + self.last_seen.discard(replica_id) + + async def touch_replica_last_seen(self, replica_id: str) -> None: + self.last_seen.add(replica_id) async def get_capacity_info(self) -> dict[str, int]: max_loras = sum(self.capacities.values()) @@ -31,6 +36,7 @@ def _make_lifecycle_manager(state: _CapacityState, replica_id: str, max_loras: i manager.max_loras = max_loras manager._replica_registered = False manager.data_plane = SimpleNamespace(close=AsyncMock()) + manager.shutdown_task_queue = AsyncMock() return manager @@ -42,6 +48,7 @@ async def test_replica_lifecycle_updates_shared_capacity() -> None: await first._register_replica_on_startup() assert await state.get_capacity_info() == {'max_loras': 3, 'used_loras': 0, 'free_loras': 3} + assert state.last_seen == {'replica-1'} await second._register_replica_on_startup() assert await state.get_capacity_info() == {'max_loras': 6, 'used_loras': 0, 'free_loras': 6} @@ -52,9 +59,10 @@ async def test_replica_lifecycle_updates_shared_capacity() -> None: @pytest.mark.asyncio async def test_async_constructor_registers_replica_before_ready() -> None: - state = SimpleNamespace(register_replica=AsyncMock()) + state = SimpleNamespace(register_replica=AsyncMock(), touch_replica_last_seen=AsyncMock()) replica_context = SimpleNamespace(replica_id=SimpleNamespace(unique_id='replica-1')) manager = ModelManagement.__new__(ModelManagement) + manager._task_queue_config = SimpleNamespace(effective_execution_timeout=1800.0) with patch('twinkle.server.model.app.DeviceGroup', return_value=SimpleNamespace(name='group')), \ patch('twinkle.server.model.app.init_twinkle_runtime', return_value=None), \ @@ -75,4 +83,5 @@ async def test_async_constructor_registers_replica_before_ready() -> None: ) state.register_replica.assert_awaited_once_with('replica-1', 3) + state.touch_replica_last_seen.assert_awaited_once_with('replica-1') assert manager._replica_registered is True diff --git a/tests/server/model/test_tinker_handlers.py b/tests/server/model/test_tinker_handlers.py index 474ce700f..66a9ff981 100644 --- a/tests/server/model/test_tinker_handlers.py +++ b/tests/server/model/test_tinker_handlers.py @@ -73,9 +73,11 @@ def assert_resource_exists(self, adapter_name): pass async def schedule_task(self, task, **kwargs): - # Actually execute the task to test response logic return await task() + async def call_backend(self, fn, /, *args, **kwargs): + return fn(*args, **kwargs) + @pytest.mark.asyncio @patch('twinkle.server.model.tinker_handlers.create_checkpoint_manager') diff --git a/tests/server/model/test_twinkle_async_inputs.py b/tests/server/model/test_twinkle_async_inputs.py index 7a13b73c5..3221394a4 100644 --- a/tests/server/model/test_twinkle_async_inputs.py +++ b/tests/server/model/test_twinkle_async_inputs.py @@ -63,6 +63,9 @@ async def schedule_task_and_wait(self, task, **kwargs): self.scheduled.append(kwargs) return await task() + async def call_backend(self, fn, /, *args, admit=True, **kwargs): + return fn(*args, **kwargs) + @pytest.mark.asyncio async def test_forward_backward_resolves_multiple_data_refs_and_field_kwargs() -> None: diff --git a/tests/server/sampler/test_stream_guarantees.py b/tests/server/sampler/test_stream_guarantees.py new file mode 100644 index 000000000..d14895751 --- /dev/null +++ b/tests/server/sampler/test_stream_guarantees.py @@ -0,0 +1,113 @@ +from __future__ import annotations + +import asyncio +import json +import threading +from types import SimpleNamespace +from unittest.mock import AsyncMock + +import pytest +from fastapi import FastAPI + +from twinkle.server.sampler.app import SamplerManagement +from twinkle.server.sampler.twinkle_handlers import _await_generation, _register_twinkle_sampler_routes, _stream_queue + + +class _BlockingQueue: + + def __init__(self) -> None: + self.released = threading.Event() + self.get_exited = threading.Event() + self.closed = False + + def get(self): + self.released.wait() + self.get_exited.set() + return 'sentinel' + + def shutdown(self, *, force: bool) -> None: + assert force is True + self.closed = True + self.released.set() + + +@pytest.mark.asyncio +async def test_sampler_request_refreshes_replica_liveness(): + service = SamplerManagement.__new__(SamplerManagement) + service.replica_id = 'sampler-replica' + service.state = SimpleNamespace(touch_replica_last_seen=AsyncMock()) + service._ensure_sticky = AsyncMock() + service._ensure_state_cleanup_started = AsyncMock() + request = SimpleNamespace( + headers={'Authorization': 'Bearer token'}, state=SimpleNamespace(token='token')) + + assert await service._on_request_start(request) == 'token' + service.state.touch_replica_last_seen.assert_awaited_once_with('sampler-replica') + + +@pytest.mark.asyncio +async def test_stream_without_actor_returns_structured_error(): + service = SimpleNamespace( + sampler=SimpleNamespace(_actors=[]), + _on_request_start=AsyncMock(return_value='token'), + ) + app = FastAPI() + _register_twinkle_sampler_routes(app, lambda: service) + route = next(route for route in app.routes if getattr(route, 'path', None) == '/twinkle/sample_stream') + request = SimpleNamespace(state=SimpleNamespace(request_id='request')) + body = SimpleNamespace(adapter_name='', adapter_uri=None, inputs={'input_ids': [1]}, sampling_params=None) + + response = await route.endpoint(request, body, service) + chunks = [chunk async for chunk in response.body_iterator] + payload = json.loads(chunks[0]) + assert payload['category'] == 'server' + assert payload['error_code'] == 503 + assert payload['request_id'].startswith('req_') + + +@pytest.mark.asyncio +async def test_stream_timeout_returns_error_payload_and_closes_queue(): + queue = _BlockingQueue() + chunks = [ + chunk async for chunk in _stream_queue( + queue, + sentinel='sentinel', + request_id='req-stream', + total_timeout=0.05, + single_get_timeout=0.05, + ) + ] + + payload = json.loads(chunks[0]) + assert payload['category'] == 'server' + assert payload['error_code'] == 504 + assert payload['request_id'] == 'req-stream' + assert queue.closed is True + assert queue.get_exited.wait(timeout=5) + + +class _GenerationService: + + def __init__(self) -> None: + self.cancelled = False + self.sampler = SimpleNamespace( + get_generation_status=lambda _submission_id: {'status': 'running'}, + collect_generation=lambda _submission_id: [], + cancel_generation=self._cancel, + ) + + async def call_backend(self, fn, /, *args, **kwargs): + return fn(*args, **kwargs) + + def _cancel(self, _submission_id: str) -> None: + self.cancelled = True + + +@pytest.mark.asyncio +async def test_generation_poll_has_total_timeout_and_cancels(): + service = _GenerationService() + + with pytest.raises(asyncio.TimeoutError): + await _await_generation(service, 'submission', timeout=0.05) + + assert service.cancelled is True diff --git a/tests/server/sampler/test_tinker_handlers.py b/tests/server/sampler/test_tinker_handlers.py index 1e338594d..e6fe1596c 100644 --- a/tests/server/sampler/test_tinker_handlers.py +++ b/tests/server/sampler/test_tinker_handlers.py @@ -20,6 +20,7 @@ class _DummySampler: def __init__(self): self.adapter_paths = [] + self.sampling_params = [] def set_template(self, *args, **kwargs): return None @@ -29,6 +30,7 @@ def reset_prefix_cache(self): def sample(self, inputs, sampling_params=None, adapter_name='', *, adapter_path=None, **kwargs): self.adapter_paths.append(adapter_path) + self.sampling_params.append(sampling_params) return [ SampleResponse( sequences=[SampledSequence( @@ -52,6 +54,9 @@ async def _on_request_start(self, request): async def schedule_task(self, task, **kwargs): return await task() + async def call_backend(self, fn, /, *args, **kwargs): + return fn(*args, **kwargs) + @pytest.mark.asyncio async def test_tinker_asample_allows_base_model_session_without_model_path(): @@ -71,4 +76,6 @@ async def test_tinker_asample_allows_base_model_session_without_model_path(): response = await route.endpoint(request, body, management) assert isinstance(response, types.SampleResponse) + assert response.sequences[0].tokens == [1, 2] assert management.sampler.adapter_paths == [None] + assert management.sampler.sampling_params[0].logprobs == 1 diff --git a/tests/server/sampler/test_twinkle_async_rows.py b/tests/server/sampler/test_twinkle_async_rows.py index 30b7016b2..234437bed 100644 --- a/tests/server/sampler/test_twinkle_async_rows.py +++ b/tests/server/sampler/test_twinkle_async_rows.py @@ -1,5 +1,7 @@ from __future__ import annotations +from types import SimpleNamespace + import pytest from fastapi import FastAPI from starlette.requests import Request @@ -67,6 +69,7 @@ def __init__(self): self.enabled = True self.scheduled = [] self.put_rows = None + self._task_queue_config = SimpleNamespace(effective_execution_timeout=60.0) async def _on_request_start(self, _request): return 'token' @@ -75,6 +78,9 @@ async def schedule_task_and_wait(self, task, **kwargs): self.scheduled.append(kwargs) return await task() + async def call_backend(self, fn, /, *args, admit=True, **kwargs): + return fn(*args, **kwargs) + def submit_generation(self, submission_id, inputs, params, **kwargs): self.submission_id = submission_id self.inputs = inputs diff --git a/tests/server/state/test_error_payload.py b/tests/server/state/test_error_payload.py new file mode 100644 index 000000000..eb224f738 --- /dev/null +++ b/tests/server/state/test_error_payload.py @@ -0,0 +1,97 @@ +# Copyright (c) ModelScope Contributors. All rights reserved. +"""Tests for ErrorPayload construction, backfill, and tinker-SDK wire compat. + +Spec: T2.4 / R9#7 / R8#5. +""" +from __future__ import annotations + +import pytest +from pydantic import ValidationError + +from twinkle.server.utils.task_errors import error_payload_from_stored, task_error_payload +from twinkle_client.types.errors import ErrorCategory, ErrorPayload + + +def test_two_field_legacy_backfills_error_code_and_request_id(): + """A pre-spec {error, category} payload backfills to 500 + passed request_id.""" + stored = {'error': 'boom', 'category': 'Server'} + + payload = error_payload_from_stored(stored, request_id='req_42') + + assert isinstance(payload, ErrorPayload) + assert payload.error_code == 500 + assert payload.request_id == 'req_42' + assert payload.error == 'boom' + + +def test_overlong_traceback_is_trimmed_tail_kept_with_marker(): + long_tb = 'X' * 10 + ('line\n' * 40000) # well over 65536 chars + assert len(long_tb) > 65536 + + payload = task_error_payload( + 'RuntimeError: boom', request_id='req_1', error_code=500, traceback_text=long_tb) + + tb = payload['traceback'] + assert tb is not None + assert len(tb) <= 65536 + assert 'truncated' in tb # truncation marker present + assert tb.endswith('line\n') # tail preserved + + +def test_user_category_carries_no_traceback(): + payload = task_error_payload( + 'invalid field', request_id='req_2', error_code=422, + category=ErrorCategory.User, traceback_text='Traceback (most recent call last): ...') + + assert payload['category'] == ErrorCategory.User.value + assert 'traceback' not in payload + + +def test_error_category_matches_tinker_wire_values(): + from tinker.types import RequestErrorCategory + + assert {item.value for item in RequestErrorCategory} == {item.value for item in ErrorCategory} + + +def test_tinker_sdk_parses_six_field_like_two_field(): + """R8#5: tinker's RequestFailedResponse ignores extra fields, so a six-field + payload parses equal to a two-field one on the declared fields. + + tinker's RequestErrorCategory values are lowercase ('server'), so the payloads + here use that value; the point under test is that the four extra fields are + ignored, not the category spelling.""" + from tinker.types import RequestFailedResponse + + two = {'error': 'boom', 'category': 'server'} + six = task_error_payload('boom', request_id='req_9', error_code=504) + + parsed_six = RequestFailedResponse.model_validate(six) + parsed_two = RequestFailedResponse.model_validate(two) + + assert parsed_six.error == parsed_two.error + assert parsed_six.category == parsed_two.category + + +def test_legacy_title_case_category_is_normalized(): + payload = error_payload_from_stored({'error': 'boom', 'category': 'Server'}, request_id='req_10') + assert payload.category is ErrorCategory.Server + assert payload.category.value == 'server' + + +@pytest.mark.parametrize('category', [ErrorCategory.User, ErrorCategory.Unknown]) +def test_non_server_traceback_is_rejected(category): + with pytest.raises(ValidationError): + ErrorPayload( + error='bad input', + category=category, + error_code=400, + request_id='req_11', + traceback='server stack', + ) + + +def test_legacy_unknown_traceback_is_removed(): + payload = error_payload_from_stored( + {'error': 'legacy', 'category': 'Unknown', 'traceback': 'old stack'}, request_id='req_12') + assert payload.category is ErrorCategory.Unknown + assert payload.traceback is None diff --git a/tests/server/state/test_future_lifecycle.py b/tests/server/state/test_future_lifecycle.py new file mode 100644 index 000000000..55905afa1 --- /dev/null +++ b/tests/server/state/test_future_lifecycle.py @@ -0,0 +1,122 @@ +# Copyright (c) ModelScope Contributors. All rights reserved. +"""State-hygiene tests for FutureManager cleanup and the do-not-regress guard. + +Spec: T5.6 / R9#5 / R9#6 / Property 6 / Property 7. Uses the Ray-free FileBackend. +""" +from __future__ import annotations + +import time +from unittest import mock + +import pytest + +from twinkle.server.state.future_manager import FutureManager + + +@pytest.fixture +def manager(tmp_path): + from twinkle.server.state.backend.file_backend import FileBackend + backend = FileBackend(str(tmp_path / 'state.json')) + return FutureManager(backend, expiration_timeout=300.0) + + +async def _store(manager, request_id, status, *, replica_id=None, absolute_deadline=None): + await manager.store_status( + request_id, + status, + model_id='m1', + replica_id=replica_id, + absolute_deadline=absolute_deadline, + ) + + +@pytest.mark.asyncio +async def test_non_terminal_with_live_replica_is_kept(manager): + await _store(manager, 'r1', 'running', replica_id='replica-A') + removed = await manager.cleanup_expired( + cutoff_time=time.time() + 10, alive_replica_ids={'replica-A'}) + assert removed == 0 + rec = await manager.get('r1') + assert rec is not None and rec.status == 'running' + + +@pytest.mark.asyncio +async def test_non_terminal_orphan_is_failed_not_deleted(manager): + await _store(manager, 'r2', 'running', replica_id='dead-replica') + await manager.cleanup_expired(cutoff_time=time.time() + 10, alive_replica_ids={'replica-A'}) + rec = await manager.get('r2') + assert rec is not None # NOT deleted (Property 6) + assert rec.status == 'failed' + assert rec.result['category'] == 'server' + + +@pytest.mark.asyncio +async def test_non_terminal_past_absolute_deadline_is_failed(manager): + await _store(manager, 'r3', 'running', replica_id='replica-A', absolute_deadline=time.time() - 1) + await manager.cleanup_expired(cutoff_time=time.time() + 10, alive_replica_ids={'replica-A'}) + rec = await manager.get('r3') + assert rec is not None and rec.status == 'failed' + + +@pytest.mark.asyncio +async def test_legacy_record_without_deadline_uses_expiration_timeout(manager): + await _store(manager, 'legacy', 'running', replica_id=None) + with mock.patch('twinkle.server.state.future_manager.time.time', return_value=time.time() + 301): + await manager.cleanup_expired(cutoff_time=time.time() + 10, alive_replica_ids=set()) + rec = await manager.get('legacy') + assert rec is not None and rec.status == 'failed' + + +@pytest.mark.asyncio +async def test_terminal_expired_is_deleted(manager): + await _store(manager, 'r4', 'completed', replica_id='replica-A') + removed = await manager.cleanup_expired( + cutoff_time=time.time() + 10, alive_replica_ids={'replica-A'}) + assert removed == 1 + assert await manager.get('r4') is None + + +@pytest.mark.asyncio +async def test_terminal_to_terminal_different_is_refused_and_warns(manager): + await _store(manager, 'r5', 'failed', replica_id='replica-A') + with mock.patch('twinkle.server.state.future_manager.logger') as log: + await manager.store_status('r5', 'completed', model_id='m1') + rec = await manager.get('r5') + assert rec.status == 'failed' # not overwritten + assert log.warning.called + + +@pytest.mark.asyncio +async def test_terminal_to_terminal_same_is_dropped_without_warning(manager): + await _store(manager, 'r6', 'completed', replica_id='replica-A') + with mock.patch('twinkle.server.state.future_manager.logger') as log: + await manager.store_status('r6', 'completed', model_id='m1') + rec = await manager.get('r6') + assert rec.status == 'completed' + assert not log.warning.called + + +@pytest.mark.asyncio +async def test_replica_id_and_deadline_set_at_creation_not_overwritten(manager): + deadline = time.time() + 100 + await _store(manager, 'r7', 'pending', replica_id='replica-A', absolute_deadline=deadline) + await manager.store_status( + 'r7', 'running', model_id='m1', replica_id='replica-B', absolute_deadline=time.time() + 999) + rec = await manager.get('r7') + assert rec.replica_id == 'replica-A' + assert rec.absolute_deadline == deadline + + +@pytest.mark.asyncio +async def test_stored_timestamps_align_with_wall_clock_regardless_of_host_tz(manager): + """Writer (_now_iso), reader (_parse_timestamp) and time.time() must agree. + + A record written now must parse to within a second of time.time() on any host, + not skewed by the host's UTC offset (the former naive-local / read-as-UTC bug). + """ + before = time.time() + await _store(manager, 'r8', 'running', replica_id='replica-A') + after = time.time() + rec = await manager.get('r8') + parsed = manager._parse_timestamp(rec.created_at) + assert before - 1 <= parsed <= after + 1 diff --git a/tests/server/state/test_managers.py b/tests/server/state/test_managers.py index 1b9016be1..63fc24beb 100644 --- a/tests/server/state/test_managers.py +++ b/tests/server/state/test_managers.py @@ -172,6 +172,12 @@ async def test_replica_registration(self, manager): assert info['used_loras'] == 0 assert info['free_loras'] == 5 + @pytest.mark.asyncio + async def test_liveness_only_replica_is_alive(self, manager): + await manager.touch_replica_last_seen('sampler-replica') + alive = await manager.get_alive_replica_ids(liveness_threshold=60) + assert 'sampler-replica' in alive + @pytest.mark.asyncio async def test_capacity_info_after_add(self, manager): await manager.register_replica('r1', max_loras=3) diff --git a/tests/server/static/__init__.py b/tests/server/static/__init__.py new file mode 100644 index 000000000..85b3e739d --- /dev/null +++ b/tests/server/static/__init__.py @@ -0,0 +1 @@ +# Copyright (c) ModelScope Contributors. All rights reserved. diff --git a/tests/server/static/backend_call_exemptions.py b/tests/server/static/backend_call_exemptions.py new file mode 100644 index 000000000..4dcd3bf01 --- /dev/null +++ b/tests/server/static/backend_call_exemptions.py @@ -0,0 +1,23 @@ +# Copyright (c) ModelScope Contributors. All rights reserved. +"""Shared exemption list for the "no direct backend call" static checks. + +This file is the SINGLE source of allowed Blocking_Backend_Call bypasses. It is +consumed by this spec's check (``test_no_direct_backend_call.py``) and is intended +to be consumed unchanged by the ``server-request-lifecycle`` spec's equivalent +check -- there must be exactly one physical copy, not one per spec (R2#8). + +Each entry is ``(module_relpath, function_name)`` where ``module_relpath`` is +relative to ``src/twinkle/server`` and ``function_name`` is the innermost enclosing +function of the exempted call. + +The only allowed exemption is the ray ``Queue.get`` inside ``sample_stream``'s +``_stream_generator``: it bridges the sampler actor's process boundary and is bounded +by the dedicated double-timeout of R4#10-11 (T5.5), not by ``call_backend``. No +``remote_function`` call is exempt. +""" +from __future__ import annotations + +# (module_relpath under src/twinkle/server, innermost enclosing function name) +BACKEND_CALL_EXEMPTIONS: frozenset[tuple[str, str]] = frozenset({ + ('sampler/twinkle_handlers.py', '_stream_queue'), +}) diff --git a/tests/server/static/test_no_degraded_path.py b/tests/server/static/test_no_degraded_path.py new file mode 100644 index 000000000..765f760a9 --- /dev/null +++ b/tests/server/static/test_no_degraded_path.py @@ -0,0 +1,68 @@ +# Copyright (c) ModelScope Contributors. All rights reserved. +"""Static check: no silent-degradation symbols remain (T7.8 / R6#6 / R9#4 / Property 9). + +One wildcard search covering eight symbols; each must occur zero times in its scope. +The symbols are matched as identifiers (word boundaries) so that ``nccl_safe_megatron`` +(the retained decorator), the ``twinkle.utils.nccl_safe`` module path, and unrelated +test names like ``test_zero_loss_...`` are not counted. +""" +from __future__ import annotations + +import pathlib +import re + +import twinkle + +_TWINKLE_SRC = pathlib.Path(twinkle.__file__).resolve().parent +_REPO_ROOT = _TWINKLE_SRC.parent.parent # .../src/twinkle -> repo root +_TESTS = _REPO_ROOT / 'tests' +_COOKBOOK = _REPO_ROOT / 'cookbook' +_SELF = pathlib.Path(__file__).resolve() + +# symbol -> compiled identifier pattern. +_IDENT = { + 'safe_loss': re.compile(r'(? str | None: + if not isinstance(node, ast.Attribute): + return None + owner = node.value + if isinstance(owner, ast.Attribute) and owner.attr in ('model', 'sampler'): + return owner.attr + return None + + +def _getattr_backend_method(node: ast.AST) -> str | None: + if not isinstance(node, ast.Call) or not isinstance(node.func, ast.Name) or node.func.id != 'getattr': + return None + if not node.args: + return None + owner = node.args[0] + if isinstance(owner, ast.Attribute) and owner.attr in ('model', 'sampler'): + return owner.attr + return None + + +class _Collector(ast.NodeVisitor): + + def __init__(self, relpath: str) -> None: + self.relpath = relpath + self.func_stack: list[str] = [] + self.backend_aliases: set[str] = set() + self.offenders: list[tuple[str, str, int, str]] = [] + + def _visit_func(self, node: ast.AST) -> None: + self.func_stack.append(node.name) + self.generic_visit(node) + self.func_stack.pop() + + visit_FunctionDef = _visit_func + visit_AsyncFunctionDef = _visit_func + + def visit_Assign(self, node: ast.Assign) -> None: + if _getattr_backend_method(node.value) is not None: + self.backend_aliases.update(target.id for target in node.targets if isinstance(target, ast.Name)) + self.generic_visit(node) + + def visit_Call(self, node: ast.Call) -> None: + owner = _backend_method(node.func) + label = ast.unparse(node.func) if owner is not None else None + if isinstance(node.func, ast.Name) and node.func.id in self.backend_aliases: + owner = 'alias' + label = node.func.id + if isinstance(node.func, ast.Attribute) and node.func.attr in ('to_thread', 'run_in_executor') and node.args: + escaped_owner = _backend_method(node.args[0]) or _getattr_backend_method(node.args[0]) + if escaped_owner is not None: + owner = escaped_owner + label = f'{ast.unparse(node.func)}({ast.unparse(node.args[0])})' + if owner is not None: + enclosing = self.func_stack[-1] if self.func_stack else '' + if (self.relpath, enclosing) not in BACKEND_CALL_EXEMPTIONS: + self.offenders.append((self.relpath, enclosing, node.lineno, label or owner)) + self.generic_visit(node) + + +def test_no_direct_backend_call_in_server(): + offenders: list[tuple[str, str, int, str]] = [] + for path in _SERVER_ROOT.rglob('*.py'): + relpath = str(path.relative_to(_SERVER_ROOT)) + collector = _Collector(relpath) + collector.visit(ast.parse(path.read_text(), filename=str(path))) + offenders.extend(collector.offenders) + + assert not offenders, ( + 'Direct backend calls must go through call_backend (or be listed in ' + f'backend_call_exemptions): {offenders}') + + +def test_exemptions_are_read_from_shared_file(): + assert ('sampler/twinkle_handlers.py', '_stream_queue') in BACKEND_CALL_EXEMPTIONS + + +def test_checker_detects_indirect_backend_calls(): + source = """ +async def route(self): + await asyncio.to_thread(self.model.save) + unload = getattr(self.sampler, 'unload_adapter_paths') + unload([]) +""" + collector = _Collector('example.py') + collector.visit(ast.parse(source)) + assert len(collector.offenders) == 2 + + +def test_checker_allows_call_backend(): + source = """ +async def route(self): + unload = getattr(self.sampler, 'unload_adapter_paths') + await self.call_backend(unload, []) +""" + collector = _Collector('example.py') + collector.visit(ast.parse(source)) + assert collector.offenders == [] diff --git a/tests/server/utils/task_queue/test_config.py b/tests/server/utils/task_queue/test_config.py index 3126c507c..878bfa12e 100644 --- a/tests/server/utils/task_queue/test_config.py +++ b/tests/server/utils/task_queue/test_config.py @@ -22,6 +22,7 @@ 'tps_limit': 16000.0, 'window_seconds': 1.0, 'queue_timeout': 300.0, + 'execution_timeout': 1800.0, 'token_cleanup_interval': 60.0, 'max_input_tokens': 16000, } @@ -113,3 +114,15 @@ def test_extra_field_rejected() -> None: """``extra='forbid'`` rejects unknown keys.""" with pytest.raises(ValidationError): TaskQueueConfig(unknown_field=1) + + +def test_zero_execution_timeout_uses_finite_fallback() -> None: + assert TaskQueueConfig(execution_timeout=0).effective_execution_timeout == 3600 + + +def test_absolute_future_ttl_uses_conservative_backend_bound() -> None: + config = TaskQueueConfig(queue_timeout=10, execution_timeout=20) + assert config.absolute_future_ttl(collect_width=2) == 2 * (10 + 2 * 3600) + + config = TaskQueueConfig(queue_timeout=10, execution_timeout=5000) + assert config.absolute_future_ttl(collect_width=2) == 2 * (10 + 2 * 5000) diff --git a/tests/server/utils/test_task_errors.py b/tests/server/utils/test_task_errors.py index 0c4986649..832ca1096 100644 --- a/tests/server/utils/test_task_errors.py +++ b/tests/server/utils/test_task_errors.py @@ -1,10 +1,37 @@ -from twinkle.server.utils.task_errors import task_error_payload +from twinkle.server.utils.task_errors import error_payload_from_stored, task_error_payload +from twinkle_client.types.errors import ErrorCategory -def test_task_error_payload_keeps_lora_traceback(): - error = 'Traceback...\nRuntimeError: No lora available for tenant session-default. Max loras: 3\n' +def test_task_error_payload_builds_error_payload_dict(): + error = 'RuntimeError: No lora available for tenant session-default. Max loras: 3' - assert task_error_payload(error) == { - 'error': error, - 'category': 'Server', - } + payload = task_error_payload(error, request_id='req_1', error_code=500) + + assert payload['error'] == error + assert payload['category'] == ErrorCategory.Server.value + assert payload['error_code'] == 500 + assert payload['request_id'] == 'req_1' + assert 'traceback' not in payload + + +def test_task_error_payload_user_category_drops_traceback(): + payload = task_error_payload( + 'bad input', request_id='req_2', error_code=400, category=ErrorCategory.User, traceback_text='Traceback...') + + assert payload['category'] == ErrorCategory.User.value + assert 'traceback' not in payload + + +def test_error_summary_is_single_line(): + payload = task_error_payload( + 'RuntimeError: boom\n File "/server/path.py", line 1', request_id='req-lines') + assert payload['error'] == 'RuntimeError: boom' + + +def test_error_payload_from_stored_backfills_two_field_legacy(): + stored = {'error': 'boom', 'category': 'Server'} + + payload = error_payload_from_stored(stored, request_id='req_3') + + assert payload.error_code == 500 + assert payload.request_id == 'req_3' diff --git a/tests/server/utils/test_task_queue_mixin.py b/tests/server/utils/test_task_queue_mixin.py index f0bdbf963..04200285a 100644 --- a/tests/server/utils/test_task_queue_mixin.py +++ b/tests/server/utils/test_task_queue_mixin.py @@ -4,6 +4,7 @@ from twinkle.server.utils.task_queue.config import TaskQueueConfig from twinkle.server.utils.task_queue.mixin import TaskQueueMixin +from twinkle.server.utils.task_queue.types import UserTaskError from twinkle.server.utils.task_queue.worker import ComputeWorker @@ -57,7 +58,7 @@ async def test_preflight_rejects_batch_without_per_dp_multiple(): assert result == {'request_id': 'req1', 'model_id': 'model1'} _, kwargs = queue.state.records[-1] - assert kwargs['result']['category'] == 'User' + assert kwargs['result']['category'] == 'user' assert 'Batch size 2 must be divisible by 4' in kwargs['result']['error'] @@ -93,6 +94,7 @@ async def work(): await asyncio.sleep(0) assert [args[1] for args, _ in queue.state.records] == ['running', 'completed'] + assert queue.state.records[0][1]['absolute_deadline'] > 0 assert queue.state.records[-1][1]['result'] == {'ok': True} @@ -122,6 +124,7 @@ async def work(): @pytest.mark.asyncio async def test_polling_schedule_task_still_persists_its_result(): queue = _DummyQueue() + queue.replica_id = 'replica-1' queue.enable_compute_worker() result = {'value': 42} @@ -142,6 +145,9 @@ async def work(): finally: await queue._compute_worker.stop() + pending = next(kwargs for args, kwargs in queue.state.records if args[1] == 'pending') + assert pending['replica_id'] == 'replica-1' + assert pending['absolute_deadline'] > 0 assert completed[-1]['result'] is result @@ -167,6 +173,28 @@ async def work(): assert queue.state.records == [] +@pytest.mark.asyncio +async def test_user_task_error_is_stored_as_user_failure(): + queue = _DummyQueue() + queue.enable_compute_worker() + + async def work(): + raise UserTaskError('invalid request') + + try: + await queue.schedule_task(work, model_id='model1', token='token1') + for _ in range(100): + failed = [kwargs for args, kwargs in queue.state.records if args[1] == 'failed'] + if failed: + break + await asyncio.sleep(0) + finally: + await queue._compute_worker.stop() + + assert failed[-1]['result']['category'] == 'user' + assert 'traceback' not in failed[-1]['result'] + + @pytest.mark.asyncio async def test_schedule_task_and_wait_reports_preflight_failure_without_persisting_it(): queue = _DummyQueue() diff --git a/tests/twinkle_client/test_types_contract.py b/tests/twinkle_client/test_types_contract.py new file mode 100644 index 000000000..3b048c614 --- /dev/null +++ b/tests/twinkle_client/test_types_contract.py @@ -0,0 +1,105 @@ +# Copyright (c) ModelScope Contributors. All rights reserved. +"""Contract-base consistency and naming-disambiguation tests. + +- T6.2 / R7#9: ``QueueStateLiteral`` value set equals the server ``QueueState`` enum. +- T6.3 / R7#7: naming disambiguation guard. + +The two SDKs already share public names. The contract freezes that legacy set and +rejects new collisions while requiring explicit aliases when both SDKs are imported +in one module. +""" +from __future__ import annotations + +import ast +import pathlib +import typing + +import twinkle +from twinkle.server.utils.task_queue.types import QueueState +from twinkle_client.types.errors import QueueStateLiteral + +_TWINKLE_SRC = pathlib.Path(twinkle.__file__).resolve().parent +_LEGACY_PUBLIC_NAME_OVERLAP = frozenset({ + 'Checkpoint', + 'CheckpointsListResponse', + 'CreateModelRequest', + 'CreateSessionRequest', + 'CreateSessionResponse', + 'Cursor', + 'ForwardRequest', + 'GetServerCapabilitiesResponse', + 'HealthResponse', + 'LoraConfig', + 'SampleRequest', + 'SessionHeartbeatRequest', + 'SessionHeartbeatResponse', + 'SupportedModel', + 'TrainingRun', + 'TrainingRunsResponse', + 'WeightsInfoResponse', + 'checkpoint', +}) + + +def test_queue_state_literal_matches_server_enum(): + literal_values = set(typing.get_args(QueueStateLiteral)) + enum_values = {state.value for state in QueueState} + assert literal_values == enum_values, ( + f'QueueStateLiteral {literal_values} != QueueState {enum_values}') + + +def _origin(module: str | None) -> str | None: + """Classify an import's source module as 'tinker', 'twinkle_client', or None.""" + if not module: + return None + if module == 'tinker' or module.startswith('tinker.'): + return 'tinker' + if module == 'twinkle_client' or module.startswith('twinkle_client.'): + return 'twinkle_client' + return None + + +def _binding_collisions(tree: ast.AST) -> set[str]: + """Return local names bound to BOTH a tinker and a twinkle_client import.""" + tinker_names: set[str] = set() + twinkle_names: set[str] = set() + for node in ast.walk(tree): + if isinstance(node, ast.ImportFrom): + origin = _origin(node.module) + if origin is None: + continue + for alias in node.names: + bound = alias.asname or alias.name + (tinker_names if origin == 'tinker' else twinkle_names).add(bound) + elif isinstance(node, ast.Import): + for alias in node.names: + origin = _origin(alias.name) + if origin is None: + continue + bound = alias.asname or alias.name.split('.')[0] + (tinker_names if origin == 'tinker' else twinkle_names).add(bound) + return tinker_names & twinkle_names + + +def test_public_name_overlap_does_not_grow(): + import tinker.types + import twinkle_client.types + + overlap = { + name + for name in set(dir(tinker.types)) & set(dir(twinkle_client.types)) + if not name.startswith('_') + } + assert overlap == _LEGACY_PUBLIC_NAME_OVERLAP + + +def test_no_tinker_twinkle_same_name_binding(): + offenders: dict[str, set[str]] = {} + for path in _TWINKLE_SRC.rglob('*.py'): + tree = ast.parse(path.read_text(), filename=str(path)) + collisions = _binding_collisions(tree) + if collisions: + offenders[str(path.relative_to(_TWINKLE_SRC))] = collisions + assert not offenders, ( + 'tinker and twinkle_client types bound to the same local name (alias tinker ' + f'to disambiguate): {offenders}') diff --git a/tests/utils/test_nccl_safe.py b/tests/utils/test_nccl_safe.py new file mode 100644 index 000000000..62b0d8e15 --- /dev/null +++ b/tests/utils/test_nccl_safe.py @@ -0,0 +1,18 @@ +from unittest.mock import patch + +import pytest + +from twinkle.utils.nccl_safe import nccl_safe_megatron + + +def test_nccl_failure_preserves_type_and_adds_rank_context(): + + @nccl_safe_megatron + def fail(_self): + raise ValueError('bad shape') + + with patch('twinkle.utils.nccl_safe._global_rank', return_value=3): + with pytest.raises(ValueError) as caught: + fail(object()) + + assert 'global_rank=3' in ''.join(getattr(caught.value, '__notes__', caught.value.args))