Documentation

Utilitaires

Détection de langue

laya.detect_language est laya.lang.analyse.

Les noms, types, valeurs par défaut et le code restent en anglais ; le reste est traduit (les entrées non encore traduites s'affichent dans l'original anglais).

analyse

analyse(state: Union[str, bytes, Mapping, list, None]) -> Dict[str, object]

Résultat complet de détection pour un état.

Renvoie script, script_profile, language (au mieux, peut être None), is_english, non_latin_fraction et mixed_segment (la ligne ou le champ qui a rendu non anglais un état majoritairement anglais, sinon None).

Ce sont les valeurs de chaîne qui sont lues. Quand un état en a plusieurs, une seule valeur non anglaise suffit : réunir chaque valeur dans une seule fenêtre laissait une longue note anglaise remplir les 4000 caractères, ou mettre en minorité un court message allemand, et ce message était alors envoyé au checkpoint anglais (#384). Le parcours par segments s’arrête toujours à 4000 caractères, ce qui garde un énorme champ peu coûteux ; une valeur qu’il n’a pas atteinte est lue à part ensuite.

Paramètres

stateUnion[str, bytes, Mapping, list, None]

detect_script

detect_script(text: str) -> str

Système d’écriture dominant de text : « latin », « han », « devanagari », ... ou « unknown » s’il n’y a aucune lettre.

Paramètres

textstr

is_english

is_english(state: Union[str, bytes, Mapping, list, None]) -> bool

True quand on peut s’attendre à ce que le checkpoint anglais lise cet état.

Paramètres

stateUnion[str, bytes, Mapping, list, None]

E-mail

clean_email_body

clean_email_body(body: str, max_chars: int = 3000) -> str

Supprime l’historique d’e-mail cité, les signatures et les avis de non-responsabilité pour garder l’entrée ciblée.

max_chars est la longueur à laquelle le résultat est coupé, 3000 caractères sauf si on l’augmente -- voir email_state, qui prend le même budget et le transmet.

Paramètres

bodystr
max_charsint= 3000

email_state

email_state(
    subject: str,
    body: str,
    sender: Optional[str] = None,
    clean: bool = True,
    max_chars: int = 3000,
    extra,
) -> Dict

Construit un dictionnaire d’état propre pour la classification d’e-mails.

max_chars est le budget auquel clean_email_body coupe le corps, et il vaut la peine de l’augmenter pour un long message : à la valeur par défaut le corps s’arrête après 3000 caractères, donc une demande qui arrive dans les derniers paragraphes n’atteint jamais le modèle -- y compris via predict_long, qui parcourt un état en fenêtres précisément pour pouvoir lire au-delà d’une fenêtre. Ignoré quand clean=False, qui fait passer le corps entier.

Tout autre mot-clé devient un champ de l’état, donc il est lu par le modèle ; une faute de frappe ici est une mutation de l’entrée, pas une erreur.

Paramètres

subjectstr
bodystr
senderOptional[str]= None
cleanbool= True
max_charsint= 3000
extra

Préréglages de questions

triage_questions

triage_questions() -> Dict

Questions prédéfinies pour le triage des tickets de support client.

email_questions

email_questions(categories: Optional[Dict[str, str]] = None) -> Dict

Questions prédéfinies pour le triage des e-mails entrants et le filtrage des menaces.

Paramètres

categoriesOptional[Dict[str, str]]= None

guard_questions

guard_questions() -> Dict

Questions prédéfinies pour les garde-fous d’entrée de LLM en temps réel.

moderation_questions

moderation_questions() -> Dict

Questions prédéfinies pour la sécurité des contenus et la modération.

router_questions

router_questions() -> Dict

Questions prédéfinies pour le routage intelligent des modèles.

Présélection

shortlist_choice

shortlist_choice(
    state: Any,
    criteria: Any,
    embed_fn: Callable[[Sequence[str]], Any],
    k: int = DEFAULT_SHORTLIST_K,
    DEFAULT_SHORTLIST_K,
    instructions: Optional[str] = None,
    return_scores: bool = False,
) -> Any

Renvoie les k meilleures étiquettes choice pour state.

embed_fn associe une liste de chaînes à un tableau de forme (len(texts), dim). Elle est appelée une fois, d’abord avec le texte de la requête puis avec une chaîne par option dans l’ordre des criteria. Les chaînes d’option correspondent à render_options pour une question choice.

Quand k est au moins le nombre d’étiquettes, chaque étiquette est renvoyée dans son ordre d’origine et embed_fn n’est pas appelée.

Les égalités conservent l’étiquette antérieure. Le classement est un cosinus signé, pas un plancher de similarité : une étiquette qui score 0 -- aucun signal, ou un vecteur non fini traité comme tel -- dépasse bien une étiquette antérieure qui a scoré négatif, et k écarte d’abord les étiquettes négatives.

Avec return_scores=True le retour est la paire (labels, scores), où scores contient le cosinus signé par étiquette conservée dans l’ordre de classement -- les mêmes valeurs que predict_shortlist rapporte dans ses métadonnées shortlist. scores est None quand rien n’a été abandonné, exactement comme dans ces métadonnées.

Paramètres

stateAny
criteriaAny
embed_fnCallable[[Sequence[str]], Any]
kint= DEFAULT_SHORTLIST_K
DEFAULT_SHORTLIST_K
instructionsOptional[str]= None
return_scoresbool= False

predict_shortlist

predict_shortlist(
    agent: Any,
    state: Any,
    questions: Dict[str, Dict[str, Any]],
    embed_fn: Callable[[Sequence[str]], Any],
    k: int = DEFAULT_SHORTLIST_K,
    DEFAULT_SHORTLIST_K,
    predict_kwargs: Any,
) -> Dict[str, Any]

Présélectionne chaque question choice, puis appelle predict ou system_one une fois.

Les questions non choice sont transmises inchangées. Une choice dont le nombre d’étiquettes est <= k est transmise inchangée et n’appelle pas embed_fn. Le dict questions de l’appelant n’est pas modifié.

Le dict renvoyé est le résultat du modèle plus une entrée shortlist. Les probabilités d’une choice présélectionnée ne portent que sur les étiquettes conservées. shortlist[qid] contient labels, scores, k, n et passthrough. labels est l’ordre de classement qu’a produit une présélection, ou l’ordre des criteria lui-même quand passthrough est défini et qu’aucun classement n’a tourné ; scores est le cosinus signé de chaque étiquette conservée dans cet ordre -- négatifs compris, jamais ramené à 0 -- ou None quand rien n’a été abandonné.

Les arguments nommés supplémentaires sont transmis à predict / system_one (par exemple model= sur un Router).

Paramètres

agentAny
stateAny
questionsDict[str, Dict[str, Any]]
embed_fnCallable[[Sequence[str]], Any]
kint= DEFAULT_SHORTLIST_K
DEFAULT_SHORTLIST_K
predict_kwargsAny

embed_fn_from_agent

embed_fn_from_agent(
    agent: Any,
    max_length: int = 512,
    batch_size: int = 32,
) -> Callable[[Sequence[str]], np.ndarray]

Met en commun par moyenne (mean-pool) l’encodeur du checkpoint déjà chargé sur agent.

Le callable encode une liste de chaînes avec agent.tok et agent.model.encoder. Il n’exécute pas la tête de décision et ne télécharge pas de poids. Un bi-encodeur dédié passé comme embed_fn présélectionnera en général mieux ; cet utilitaire est pour les appelants qui n’ont que le checkpoint Laya en mémoire.

Les positions de padding sont exclues de la moyenne. L’indicateur train/eval de l’encodeur est laissé tel que l’appelant l’a fixé (un Agent chargé est déjà en eval). Chaque appel utilise l’agent.device courant, y compris après un repli sur CPU.

Paramètres

agentAny
max_lengthint= 512
batch_sizeint= 32

cached_embed_fn

cached_embed_fn(
    embed_fn: Callable[[Sequence[str]], Any],
    maxsize: int = 4096,
) -> Callable[[Sequence[str]], np.ndarray]

Met en cache la sortie de embed_fn par chaîne d’entrée, sous une borne LRU.

predict_shortlist encode la requête plus chaque texte d’option à chaque appel. Quand le même ensemble d’options est présélectionné à chaque requête -- une liste fixe d’intentions ou d’étiquettes, comme dans l’exemple BANKING77 du README -- les lignes d’option ne changent pas entre les appels, mais elles sont ré-encodées à chaque fois. Envelopper l’embedder une fois ::

embed_fn = cached_embed_fn(embed_fn_from_agent(agent))

laisse le premier appel inchangé et réduit chaque appel répété au seul encodage de la nouvelle requête.

Les recherches sont des correspondances exactes de chaînes. Les textes absents du cache sont dédupliqués et encodés en un seul appel embed_fn, donc un cache froid coûte le même nombre d’appels par lots que la fonction non enveloppée. Les lignes sont stockées en float32 ; le cache contient au plus maxsize chaînes puis évince l’entrée la moins récemment utilisée, bornant la mémoire à environ maxsize * dim * 4 octets. Rien n’est mis en cache quand embed_fn lève une exception ou renvoie une mauvaise forme.

Le wrapper peut être partagé entre threads sans risque : le verrou ne couvre que les lectures et écritures du cache, jamais l’appel d’embedding. Le callable renvoyé porte cache_info() -- un dict avec size, maxsize, hits et misses -- et cache_clear(). Vide le cache si le modèle ou les poids derrière embed_fn changent.

Paramètres

embed_fnCallable[[Sequence[str]], Any]
maxsizeint= 4096

Abstention

check_min_confidence

check_min_confidence(v: Any)

Valide le seuil d’abstention optionnel min_confidence (#361, #394).

Soit un nombre réel dans [0.0, 1.0] (un seuil pour chaque réponse ; les booléens sont rejetés même si isinstance(True, int)), soit un mapping par bucket (voir :func:check_min_confidence_map) pour que le seuil puisse différer selon le nombre d’options. Renvoie la valeur sous sa forme validée -- un float pour le cas scalaire, un dict[str, float] pour le cas mapping --, que les fonctions de gate ci-dessous acceptent toutes deux.

Paramètres

vAny

check_min_confidence_map

check_min_confidence_map(m: Dict[Any, Any]) -> Dict[str, float]

Valide un mapping de seuils d’abstention par bucket (#394).

Les clés sont des chaînes de bucket de nombre d’options selon l’orthographe de common.temp_bucket -- "choice:2", "choice:3-5", "score:6-10", "noul:2" et ainsi de suite --, plus un "default" optionnel utilisé pour tout bucket que le mapping ne nomme pas. Les valeurs sont des floats dans [0.0, 1.0]. Un seuil de confiance ne se transfère pas d’un nombre d’options à l’autre (#394) ; cela permet à l’appelant de filtrer chaque bucket au niveau que sa calibration gagne réellement. Ajuste-en un avec :func:laya.calibrate.fit_abstention_thresholds.

Paramètres

mDict[Any, Any]

resolve_min_confidence

resolve_min_confidence(
    answer: Dict[str, Any],
    thresholds: Dict[str, float],
    default: float = 0.0,
) -> float

Le seuil auquel est filtré le bucket de nombre d’options de cette réponse, sous un mapping par bucket.

Retombe sur l’entrée "default" du mapping, puis sur default (0.0 -- ne rien filtrer), pour un bucket que le mapping ne nomme pas, afin qu’un bucket non configuré ne s’abstienne jamais par surprise.

Paramètres

answerDict[str, Any]
thresholdsDict[str, float]
defaultfloat= 0.0

flag_low_confidence

flag_low_confidence(results: List[Dict[str, Any]], min_confidence: float) -> None

Marqueur d’abstention optionnel (#361) : marque les réponses dont la confiance tombe sous min_confidence.

Lit answer_confidence (max(p), la quantité que décrivent les chiffres de calibration et celle qui ne dérive pas avec le nombre d’options), en retombant sur confidence si answer_confidence est absent. La réponse et la confiance brutes restent intactes ; low_confidence: True est ajouté quand la réponse tombe sous le seuil, et retiré si une réponse précédemment marquée le dépasse désormais (par exemple quand un dict de résultat est réutilisé ou réévalué avec un seuil différent).

min_confidence est soit un float (un seuil pour chaque réponse), soit un mapping par bucket (#394), auquel cas chaque réponse est filtrée au seuil de son propre bucket de nombre d’options via :func:resolve_min_confidence.

Paramètres

resultsList[Dict[str, Any]]
min_confidencefloat

apply_confidence_gate

apply_confidence_gate(
    results: List[Dict[str, Any]],
    min_confidence: Optional[float] = None,
) -> None

Rapporte l’état du gate de confiance, sur les réponses auxquelles un gate a réellement été appliqué.

Un gate est une politique, et une politique dont l’application ne peut pas être observée n’en est pas une. Avec min_confidence défini, ceci écrit abstention -- l’un des :data:GATE_STATES -- sur chaque réponse, plus abstention_threshold, afin qu’un appelant puisse répondre à trois questions qu’il ne peut autrement pas trancher :

  • quelle fraction des décisions s’est abstenue, au lieu de l’inférer de ce que low_confidence se trouvait être défini ;
  • sur combien de réponses le gate n’a pas pu décider, ce qu’un booléen ne peut pas exprimer du tout ;
  • quel seuil a produit ces résultats -- flag_low_confidence consomme le seuil et le laisse tomber, donc sans ceci une exécution par lots avec des seuils par classe ne peut pas être re-découpée.

GATE_UNEVALUATED est le cas qu’un booléen ne peut pas exprimer : le gate s’est exécuté et la réponse ne portait pas de confiance utilisable, donc le gate n’a pas pu décider. Rapporter cela comme une réussite est le même mensonge que de le rapporter comme un flag.

Avec min_confidence non défini, ceci n’écrit rien. Aucun abstention, aucun abstention_threshold, aucun flag. C’est tout le contrat : un appel sans gate renvoie exactement le payload qu’il renvoyait avant, et la présence du champ -- pas une quatrième valeur lue à l’intérieur -- est ce qui dit à l’appelant que le gate s’est exécuté. Appelle-le inconditionnellement, une fois par appel, à la place d’un garde if min_confidence is not None: : ce garde est ce qui laisse un chemin ne rapportant rien du tout, ce qui est l’état que cette fonction existe pour distinguer.

Le flag lui-même reste celui de :func:flag_low_confidence -- ceci délègue au lieu de réimplémenter la règle, donc le booléen et l’état rapporté ne peuvent pas diverger.

Un min_confidence d’exactement 0.0 a bien été défini, donc des états sont rapportés, et :func:flag_low_confidence traite 0.0 comme un no-op parce que rien ne peut tomber en dessous. Chaque réponse portant une confiance utilisable se lit donc comme passed, et l’écho du seuil est ce qui distingue cela d’une vraie réussite à un vrai seuil.

Paramètres

resultsList[Dict[str, Any]]
min_confidenceOptional[float]= None

GATE_STATES

GATE_STATES = (GATE_PASSED, GATE_ABSTAINED, GATE_UNEVALUATED)

Calibration et entraînement

answer_confidence

answer_confidence(p: np.ndarray, k: int) -> float

Masse de probabilité sur la réponse rapportée : max(p).

C’est la quantité qu’ajuste la mise à l’échelle des températures, et la quantité sur laquelle est calculé chaque chiffre de calibration de ce dépôt -- les deux harnais de benchmark prennent conf = max(probs) avant d’appeler ece_score. La section gating du README s’appuie sur la propriété qui l’accompagne : parmi les réponses renvoyées à la confiance c, environ c sont justes. Cette propriété est conditionnelle, et la condition n’est pas remplie par défaut -- elle ne tient qu’après que les températures ont été ajustées et validées sur des données réservées pour ce checkpoint et ce nombre d’options. Les checkpoints livrés sont trop confiants : choice:11+ est un aiguiseur ~10x qui renvoie une masse ponctuelle à 1.0, donc un seuil qui leur est appliqué sélectionne en dessous de l’exactitude du modèle (issue #394).

confidence_from_probs ci-dessous rapporte une quantité différente sur une échelle différente et ne porte pas une telle garantie, donc les deux ne doivent pas être comparées au même seuil.

Paramètres

pnp.ndarray
kint

confidence_from_probs

confidence_from_probs(p: np.ndarray, k: int) -> float

Confiance par entropie de Shannon normalisée : 1 - H(p) / log(k).

À quel point toute la distribution est concentrée. Utile, mais pas calibrée : ce n’est pas ce qu’ajuste la mise à l’échelle des températures, ni ce que mesure l’ECE rapporté. Voir answer_confidence.

Paramètres

pnp.ndarray
kint

ece_score

ece_score(conf: np.ndarray, correct: np.ndarray, bins: int = 15) -> float

Erreur de calibration attendue (ECE) sur les intervalles de confiance.

Paramètres

confnp.ndarray
correctnp.ndarray
binsint= 15

fit_temperatures

fit_temperatures = fit_temperature_map

fit_one_temperature

fit_one_temperature(pairs: Sequence, min_n: Optional[int] = None) -> float

Ajuste un scalaire T unique par NLL + LBFGS sur log T.

Le résultat est clamp_temperature de l’échelle optimisée, donc il se situe dans [TEMP_MIN, TEMP_MAX] (ou vaut le neutre 1.0 quand la valeur n’est pas un nombre). Renvoie 1.0 quand moins de min_n paires sont fournies. min_n vaut par défaut MIN_BUCKET_N (le plancher par bucket). Les ajustements au niveau du type passent MIN_TYPE_N, qui est plus bas, donc un jeu de données qui ne remplit aucun bucket obtient quand même un scalaire au lieu de rester à 1.0.

Paramètres

pairsSequence
min_nOptional[int]= None

fit_temperature_map

fit_temperature_map(
    records: Iterable,
    compute_ece: bool = False,
    seed: int = 0,
) -> Dict[str, Any]

Ajuste des scalaires au niveau du type et des températures par bucket.

MIN_BUCKET_N est le plancher par bucket : les buckets plus petits sont omis de temperature_by_options et le scalaire au niveau du type les couvre. MIN_TYPE_N est le plancher distinct et plus bas, propre à ce scalaire.

compute_ece=False (le défaut, et le chemin que stocke Agent.fit_temperatures) ajuste sur chaque enregistrement et ne renvoie aucune clé report. seed est ignoré sur ce chemin.

compute_ece=True met de côté ECE_HOLDOUT_FRAC de chaque bucket, stratifié par temp_bucket, en utilisant seed pour que les mêmes enregistrements se répartissent toujours de la même façon. Les températures sont ajustées sur le reste uniquement et l’ECE n’est scorée que sur les enregistrements réservés. report["n"] est le nombre d’enregistrements fournis ; report["n_eval"] est le compte réservé sur lequel repose l’ECE. Un bucket qui tomberait sous MIN_BUCKET_N après la mise de côté est ajusté sur tous ses enregistrements, écarté de l’ensemble d’évaluation et nommé dans report["buckets_excluded_from_eval"] au lieu d’être abandonné. n_by_bucket compte toujours l’entrée complète, y compris quand l’ajustement lui-même a utilisé un sous-ensemble.

Paramètres

recordsIterable
compute_ecebool= False
seedint= 0

fit_abstention_thresholds

fit_abstention_thresholds(
    records: Iterable,
    temperature: Sequence[float],
    temperature_by_options: Dict[str, float],
    binning_map: Optional[Dict[str, Dict[str, Any]]] = None,
    target_error: float = 0.10,
    min_bucket_n: int = MIN_ABSTAIN_BUCKET_N,
    MIN_ABSTAIN_BUCKET_N,
    conservative: bool = True,
) -> Dict[str, float]

Ajuste un seuil d’abstention par temp_bucket pour qu’un gate garde une erreur cible dans chaque bucket.

Un seul min_confidence ne se transfère pas d’un nombre d’options à l’autre (#394) : la confiance calibrée d’une réponse à 2 options et d’une à 12 options vivent sur des échelles différentes, donc un seuil s’abstient trop ou pas assez selon la question. Ceci ajuste à la place un seuil par bucket, indexé exactement comme temperature_by_options (common.temp_bucket, p. ex. "choice:3-5"), et le résultat est une table de min_confidence que :func:laya.confidence.check_min_confidence / :func:laya.confidence.apply_confidence_gate acceptent directement.

records sont les mêmes tuples (qtype, logits, target[, k]) que consomme fit_temperature_map (records_from_labeled les construit). La confiance est le max(p) calibré -- les logits sont mis à l’échelle d’abord par la temperature / temperature_by_options ajustées, donc les seuils et les nombres que le runtime rapporte sont sur la même échelle. target_error est l’erreur tolérée parmi les réponses acceptées ; min_bucket_n omet les buckets trop petits pour être ajustés, et conservative ajoute une marge d’un échantillon. Les seuils sont des coupes empiriques sur le jeu de calibration, pas une garantie formelle de couverture -- valide sur des données réservées (fit_temperature_map(..., compute_ece=True) donne une partition réservée) pour un gate de production.

Passe binning_map quand l’agent qui servira ces seuils en a une installée -- par Agent.fit_binning, ou par un payload de calibration qui porte binning_map --, parce que le runtime recalibre answer_confidence à travers cette table avant que quoi que ce soit ne la lise, donc un seuil ajusté sans elle est une coupe sur une échelle que le gate ne voit jamais. Les seuils sont alors sur l’échelle avec binning, et l’ordre dans lequel les deux ont été ajustés cesse d’importer. Mesuré sur 1,200 enregistrements synthétiques à 12 options avec target_error=0.10 : le seuil ajusté sans table garde 9.8% d’erreur sur 50% de couverture sur des confiances sans binning, et admet 94.5% des réponses à 25.6% d’erreur une fois que le même nombre est comparé à des confiances avec binning.

Paramètres

recordsIterable
temperatureSequence[float]
temperature_by_optionsDict[str, float]
binning_mapOptional[Dict[str, Dict[str, Any]]]= None
target_errorfloat= 0.10
min_bucket_nint= MIN_ABSTAIN_BUCKET_N
MIN_ABSTAIN_BUCKET_N
conservativebool= True

fit_binning_map

fit_binning_map(
    records: Iterable,
    temperature: Sequence[float],
    temperature_by_options: Dict[str, float],
    bins: int = 15,
    min_bucket_n: int = MIN_BINNING_BUCKET_N,
    MIN_BINNING_BUCKET_N,
) -> Dict[str, Dict[str, Any]]

Ajuste une table de recalibrage par binning d’histogramme par temp_bucket pour answer_confidence.

La mise à l’échelle des températures applique un scalaire par bucket ; elle ne peut pas corriger un bucket dont la courbe de fiabilité n’est pas un simple renforcement/atténuement (le choice:11+ pathologique que porte le checkpoint anglais livré en est un). Le binning par histogramme est l’alternative non paramétrique : découpe les confiances calibrées d’un bucket en bins bins de largeur égale sur [0, 1], et mappe chaque confiance qui tombe dans un bin vers l’exactitude empirique de ce bin. Il ne requiert aucune hypothèse de monotonie et aucune dépendance supplémentaire (NumPy uniquement ; la régression isotonique tirerait scikit-learn).

records sont les mêmes tuples (qtype, logits, target[, k]) que consomme fit_temperature_map ; la confiance est le max(p) calibré (logits mis à l’échelle par la temperature / temperature_by_options ajustées d’abord), donc une table de binning se compose par-dessus une table de température au lieu de la remplacer. Renvoie {bucket: {"bins": N, "values": [recalibrated confidence per bin]}} ; les buckets sous min_bucket_n sont omis. Applique-la avec :func:apply_binning_map. Un bin vide (une plage de confiance que le jeu de calibration n’a jamais produite) se mappe sur son propre point médian, c’est-à-dire laisse cette région inchangée, donc une valeur inédite n’est jamais recalibrée vers un 0 fabriqué.

Paramètres

recordsIterable
temperatureSequence[float]
temperature_by_optionsDict[str, float]
binsint= 15
min_bucket_nint= MIN_BINNING_BUCKET_N
MIN_BINNING_BUCKET_N

apply_binning_map

apply_binning_map(
    confidence: float,
    bucket: str,
    binning_map: Dict[str, Dict[str, Any]],
) -> float

Recalibre une answer_confidence pour son bucket de nombre d’options (common.temp_bucket).

Renvoie la confiance inchangée quand la table n’a pas d’entrée pour le bucket, donc un bucket pour lequel la table n’a pas été ajustée passe plutôt que d’être forcé vers une valeur fausse.

Paramètres

confidencefloat
bucketstr
binning_mapDict[str, Dict[str, Any]]

fit_binning

fit_binning(
    records,
    min_bucket_n: int = MIN_BINNING_BUCKET_N,
    MIN_BINNING_BUCKET_N,
) -> Dict[str, Any]

Ajuste une table de binning par histogramme par-dessus les températures ajustées de cet agent et la stocke.

records sont les mêmes tuples (qtype, logits, target[, k]) que ceux consommés par fit_temperatures. La table est indexée exactement comme temperature_by_options, se compose par-dessus les températures actuelles, et save_calibration l’écrit sous le nom binning_map.

Paramètres

records
min_bucket_nint= MIN_BINNING_BUCKET_N
MIN_BINNING_BUCKET_N

render_options

render_options(q: Dict) -> List[str]

Rend les textes d’option dans l’ordre des indices d’étiquette. L’ordre sémantique de Noul est toujours [false, true].

Paramètres

qDict

proper_reward

proper_reward(
    q: torch.Tensor,
    target: torch.Tensor,
    qtype: torch.Tensor,
    mask: torch.Tensor,
    w_sph: float = 0.5,
    w_rps: float = 1.0,
    log_floor: float = -9.21,
) -> torch.Tensor

Récompense de règle de score strictement propre : log score + spherical score + ranked probability score.

q: [..., N, K] distributions rapportées target: [N, K] (distributions cibles one-hot ou douces)

Paramètres

qtorch.Tensor
targettorch.Tensor
qtypetorch.Tensor
masktorch.Tensor
w_sphfloat= 0.5
w_rpsfloat= 1.0
log_floorfloat= -9.21

td_lambda_targets

td_lambda_targets(p_true: torch.Tensor, batch: Dict, lam: float = 1.0) -> torch.Tensor

Cibles TD(lambda) pour des trajectoires de conversation multi-tours.

Paramètres

p_truetorch.Tensor
batchDict
lamfloat= 1.0

QTYPES

QTYPES = {"choice": 0, "score": 1, "noul": 2}

QTYPE_NAMES

QTYPE_NAMES = {v: k for k, v in QTYPES.items()}