File size: 1,951 Bytes
034506e
 
 
 
 
 
 
 
 
012abcf
 
034506e
 
 
 
71db45a
034506e
71db45a
034506e
 
 
 
 
 
012abcf
 
 
 
034506e
71db45a
 
 
 
c490062
71db45a
 
 
 
 
034506e
71db45a
034506e
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
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
64
65
66
67
68
69
"""Registry of translation engines.

Add a new engine by implementing ``TranslationProvider`` and appending an
instance to ``_PROVIDERS`` below. The API automatically exposes whichever ones
report themselves available.
"""

from __future__ import annotations

import os

from providers.base import ProviderInfo, TranslationProvider
from providers.gemini import GeminiProvider
from providers.groq import GroqProvider, GroqQwenProvider
from providers.madlad import MADLADProvider
from providers.nllb import NLLBProvider
from providers.ollama import OllamaProvider
from providers.gemma import SwahiliGemmaProvider

_PROVIDERS: list[TranslationProvider] = [
    NLLBProvider(),
    NLLBProvider(
        provider_id="nllb_1_3b",
        name="NLLB-200 (1.3B)",
        model_dir=os.environ.get(
            "CT2_MODEL_DIR_1_3B", "models/nllb-200-distilled-1.3B-int8"
        ),
        hf_model=os.environ.get("HF_MODEL_1_3B", "facebook/nllb-200-distilled-1.3B"),
    ),
    NLLBProvider(
        provider_id="afrinllb",
        name="AfriNLLB (600M)",
        model_dir=os.environ.get(
            "CT2_MODEL_DIR_AFRINLLB", "/data/models/afrinllb-600m-int8"
        ),
        hf_model=os.environ.get(
            "HF_MODEL_AFRINLLB", "AfriNLP/AfriNLLB-12enc-12dec-full-ft-kd"
        ),
    ),
    MADLADProvider(),
    SwahiliGemmaProvider(),
    OllamaProvider(),
    GeminiProvider(),
    GroqQwenProvider(),
    GroqProvider(),
]

_BY_ID = {p.id: p for p in _PROVIDERS}


def all_infos() -> list[ProviderInfo]:
    return [p.info() for p in _PROVIDERS]


def available_infos() -> list[ProviderInfo]:
    return [p.info() for p in _PROVIDERS if p.is_available()]


def get(provider_id: str) -> TranslationProvider | None:
    return _BY_ID.get(provider_id)


def default_id() -> str | None:
    """First available engine, preferring local/private ones."""
    for p in _PROVIDERS:
        if p.is_available():
            return p.id
    return None