File size: 1,364 Bytes
d53adc9
 
 
 
 
 
 
 
 
cb4d90c
d53adc9
 
 
 
 
 
cb4d90c
 
 
 
d53adc9
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
1
2
3
4
5
6
7
8
9
10
11
12
13
14
15
16
17
18
19
20
21
22
23
24
25
26
27
28
29
30
31
32
33
34
35
36
37
38
39
40
41
42
43
44
# Generated by ml.integrations.export.runtime_packager.write_remote_code_bundle.
# Exported for HuggingFace trust_remote_code loading.
# This file is intentionally self-contained.

"""Projection helpers for the native Sophia Hybrid configuration."""

from __future__ import annotations

from dataclasses import fields
from typing import Protocol, TypeVar


class RuntimeModelArgsSource(Protocol):
    max_seq_len: int


RuntimeModelArgsT = TypeVar("RuntimeModelArgsT")


def build_runtime_model_args(
    config: RuntimeModelArgsSource,
    *,
    model_args_cls: type[RuntimeModelArgsT],
    runtime_max_seq_len: int | None = None,
) -> RuntimeModelArgsT:
    max_seq_len = int(config.max_seq_len)
    if runtime_max_seq_len is not None:
        runtime_limit = int(runtime_max_seq_len)
        if not 0 < runtime_limit <= max_seq_len:
            raise ValueError(
                "runtime_max_seq_len must be in (0, config.max_seq_len], "
                f"got {runtime_limit} with max {max_seq_len}"
            )
        max_seq_len = runtime_limit
    kwargs: dict[str, object] = {}
    for arg_field in fields(model_args_cls):
        if hasattr(config, arg_field.name):
            kwargs[arg_field.name] = getattr(config, arg_field.name)
    kwargs["max_seq_len"] = max_seq_len
    return model_args_cls(**kwargs)


__all__ = ["build_runtime_model_args"]