File size: 2,148 Bytes
d10ad42
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
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
45
46
47
48
49
50
51
52
53
54
55
56
57
58
59
60
61
62
63
"""Server settings: every value comes from an environment variable (see serve.env), CLI flags override.

Context and batch limits default to what fits the GPU: a 27B model in bf16 needs ~54 GB for weights,
so a 96 GB RTX PRO 6000 gets 32k-token requests / 16k-token batches, while 288 GB cards get 131k.
"""

import os
from dataclasses import dataclass
from pathlib import Path

DEFAULT_CHECKPOINT = str(Path(__file__).resolve().parents[2])


@dataclass(frozen=True)
class Settings:
    checkpoint: str
    model_name: str
    host: str
    port: int
    api_key: str | None
    max_context: int
    token_budget: int
    batch_size: int
    queue_seconds: float
    warmup: bool
    device: str | None


def gpu_limits() -> tuple[int, int]:
    """(max context tokens, batch token budget) sized to the visible GPU."""
    import torch

    if not torch.cuda.is_available():
        return 8192, 8192
    gib = torch.cuda.get_device_properties(0).total_memory / 2**30
    if gib >= 200:
        return 131072, 131072
    if gib >= 80:
        return 32768, 16384
    return 8192, 8192


def load() -> Settings:
    checkpoint = os.getenv("MJ_CHECKPOINT", DEFAULT_CHECKPOINT)
    key = os.getenv("MJ_API_KEY") or None
    key_file = os.getenv("MJ_API_KEY_FILE")
    if key is None and key_file and Path(key_file).is_file():
        key = Path(key_file).read_text().strip() or None
    auto_context, auto_budget = gpu_limits() if not (os.getenv("MJ_MAX_CONTEXT") and os.getenv("MJ_TOKEN_BUDGET")) else (0, 0)
    return Settings(
        checkpoint=checkpoint,
        model_name=os.getenv("MJ_MODEL_NAME") or Path(checkpoint).resolve().name,
        host=os.getenv("MJ_HOST", "127.0.0.1"),
        port=int(os.getenv("MJ_PORT", "8000")),
        api_key=key,
        max_context=int(os.getenv("MJ_MAX_CONTEXT") or auto_context),
        token_budget=int(os.getenv("MJ_TOKEN_BUDGET") or auto_budget),
        batch_size=int(os.getenv("MJ_BATCH_SIZE", "64")),
        queue_seconds=float(os.getenv("MJ_QUEUE_SECONDS", "60")),
        warmup=os.getenv("MJ_WARMUP", "1") not in ("0", "false", "no"),
        device=os.getenv("MJ_DEVICE") or None,
    )