Documentação

Auxiliares

Deteção de idioma

laya.detect_language é laya.lang.analyse.

Os nomes, tipos, valores predefinidos e código mantêm-se em inglês; o resto é traduzido (as entradas ainda não traduzidas são mostradas no original em inglês).

analyse

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

Resultado completo de deteção para um estado.

Devolve script, script_profile, language (o melhor possível, pode ser None), is_english, non_latin_fraction e mixed_segment (a linha ou o campo que fez um estado maioritariamente inglês deixar de o ser; caso contrário, None).

O que é lido são os valores de cadeia. Quando um estado tem vários, basta um que não seja inglês: juntar cada valor numa só janela deixava uma nota longa em inglês preencher os 4000 caracteres, ou superar uma mensagem curta em alemão, e essa mensagem era então enviada para o checkpoint de inglês (#384). A análise por segmentos continua a parar aos 4000 caracteres, o que é o que mantém barato um campo enorme; um valor que não alcançou é lido à parte depois.

Parâmetros

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

detect_script

detect_script(text: str) -> str

Sistema de escrita dominante de text: 'latin', 'han', 'devanagari', ... ou 'unknown' se não houver letras.

Parâmetros

textstr

is_english

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

True quando é de esperar que o checkpoint de inglês consiga ler este estado.

Parâmetros

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

Email

clean_email_body

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

Remove o histórico de correio citado, as assinaturas e os avisos legais para manter a entrada focada.

max_chars é o comprimento para o qual o resultado é cortado, 3000 caracteres a menos que seja aumentado -- consulta email_state, que usa o mesmo orçamento e o transmite.

Parâmetros

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

Constrói um dicionário de estado limpo para a classificação de correio.

max_chars é o orçamento para o qual clean_email_body corta o corpo, e vale a pena aumentá-lo para uma mensagem longa: com o valor predefinido o corpo para aos 3000 caracteres, pelo que um pedido que chega nos últimos parágrafos nunca alcança o modelo -- incluindo através de predict_long, que analisa um estado em janelas precisamente para poder ler para além de uma janela. É ignorado quando clean=False, que deixa passar o corpo inteiro.

Qualquer outra palavra-chave torna-se um campo do estado, pelo que é lida pelo modelo; um erro de escrita aqui é uma mutação da entrada, não um erro.

Parâmetros

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

Predefinições de perguntas

triage_questions

triage_questions() -> Dict

Perguntas predefinidas para a triagem de tickets de apoio ao cliente.

email_questions

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

Perguntas predefinidas para a triagem de correio recebido e a filtragem de ameaças.

Parâmetros

categoriesOptional[Dict[str, str]]= None

guard_questions

guard_questions() -> Dict

Perguntas predefinidas para guardrails de entrada de LLM em tempo real.

moderation_questions

moderation_questions() -> Dict

Perguntas predefinidas para a segurança de conteúdo e a moderação.

router_questions

router_questions() -> Dict

Perguntas predefinidas para o encaminhamento inteligente de modelos.

Seleção prévia

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

Devolve as k etiquetas choice principais para state.

embed_fn mapeia uma lista de cadeias para um array de forma (len(texts), dim). É chamado uma vez, primeiro com o texto de consulta e depois com uma cadeia por opção em ordem de criteria. As cadeias de opção correspondem a render_options para uma pergunta choice.

Quando k é pelo menos o número de etiquetas, cada etiqueta é devolvida na sua ordem original e embed_fn não é chamado.

Os empates conservam a etiqueta anterior. A classificação é um cosseno com sinal, não um piso de similaridade: uma etiqueta que pontua 0 -- nenhum sinal, ou um vetor não finito tratado como tal -- supera uma etiqueta anterior que pontuou negativo, e k descarta as etiquetas negativas primeiro.

Com return_scores=True o retorno é o par (labels, scores), em que scores contém o cosseno com sinal de cada etiqueta conservada em ordem de classificação -- os mesmos valores que predict_shortlist reporta nos seus metadados shortlist. scores é None quando nada foi descartado, exatamente como nesses metadados.

Parâmetros

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é-seleciona cada pergunta choice e depois chama uma vez predict ou system_one.

As perguntas que não são choice são reencaminhadas sem alterações. Uma choice cujo número de etiquetas é <= k é reencaminhada sem alterações e não chama embed_fn. O dict questions de quem chama não é modificado.

O dict devolvido é o resultado do modelo mais uma entrada shortlist. As probabilidades de uma choice pré-selecionada são apenas sobre as etiquetas conservadas. shortlist[qid] contém labels, scores, k, n e passthrough. labels é a ordem de classificação que uma pré-seleção produziu, ou a própria ordem de criteria quando passthrough está definido e nenhuma classificação correu; scores é o cosseno com sinal de cada etiqueta conservada nessa ordem -- negativo incluído, nunca limitado a 0 -- ou None quando nada foi descartado.

Os argumentos de palavra-chave adicionais são reencaminhados para predict / system_one (por exemplo model= num Router).

Parâmetros

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]

Faz mean-pool do encoder do checkpoint já carregado em agent.

O callable gera embeddings de uma lista de cadeias com agent.tok e agent.model.encoder. Não executa a cabeça de decisão nem transfere pesos. Um bi-encoder dedicado passado como embed_fn normalmente fará uma melhor pré-seleção; este helper é para quem chama e só tem o checkpoint da Laya em memória.

As posições de padding são excluídas da média. O indicador train/eval do encoder é deixado como quem chama o definiu (um Agent carregado já está em eval). Cada chamada usa o agent.device atual, mesmo depois de recair sobre a CPU.

Parâmetros

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]

Guarda em cache a saída de embed_fn por cadeia de entrada, sob um limite LRU.

predict_shortlist gera o embedding da consulta mais o de cada texto de opção em cada chamada. Quando o mesmo conjunto de opções é pré-selecionado em todos os pedidos -- uma lista fixa de intenções ou etiquetas, como no exemplo BANKING77 do README -- as linhas de opção não mudam entre chamadas, mas os seus embeddings são novamente gerados em cada uma. Envolver o embedder uma vez::

embed_fn = cached_embed_fn(embed_fn_from_agent(agent))

deixa a primeira chamada inalterada e reduz cada chamada repetida a gerar o embedding apenas da nova consulta.

As pesquisas são correspondências exatas de cadeia. Os textos em falta na cache são deduplicados e os seus embeddings são gerados numa única chamada a embed_fn, pelo que uma cache fria custa o mesmo número de chamadas em lote que a função sem envolver. As linhas são guardadas como float32; a cache guarda no máximo maxsize cadeias e depois despeja a entrada usada menos recentemente, limitando a memória a cerca de maxsize * dim * 4 bytes. Nada é guardado em cache quando embed_fn lança uma exceção ou devolve uma forma inválida.

O wrapper é seguro para partilhar entre threads: o lock cobre apenas as leituras e escritas da cache, nunca a chamada de embedding. O callable devolvido transporta cache_info() -- um dict com size, maxsize, hits e misses -- e cache_clear(). Limpa a cache se o modelo ou os pesos por trás de embed_fn mudarem.

Parâmetros

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

Abstenção

check_min_confidence

check_min_confidence(v: Any)

Valida o limiar de abstenção opcional min_confidence (#361, #394).

Ou um número real em [0.0, 1.0] (um limiar para cada resposta; os booleanos são rejeitados embora isinstance(True, int)), ou um mapeamento por cubo (consulta :func:check_min_confidence_map) de modo que o limiar possa variar com o número de opções. Devolve o valor na sua forma validada -- um float para o caso escalar, um dict[str, float] para o caso de mapeamento -- que as funções de gating abaixo aceitam em ambos os casos.

Parâmetros

vAny

check_min_confidence_map

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

Valida um mapa de limiares de abstenção por cubo (#394).

As chaves são strings de cubo de número de opções na grafia de common.temp_bucket -- "choice:2", "choice:3-5", "score:6-10", "noul:2" e assim sucessivamente -- além de um "default" opcional usado para qualquer cubo que o mapa não nomeie. Os valores são números reais em [0.0, 1.0]. Um limiar de confiança não se transfere entre números de opções (#394); isto permite que quem chama filtre cada cubo ao nível que a sua calibração realmente merece. Ajuste um com :func:laya.calibrate.fit_abstention_thresholds.

Parâmetros

mDict[Any, Any]

resolve_min_confidence

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

O limiar a que o cubo de número de opções desta resposta é filtrado, sob um mapa por cubo.

Recorre à entrada "default" do mapa e depois a default (0.0 -- não filtra nada), para um cubo que o mapa não nomeie, para que um cubo não configurado nunca se abstenha de forma inesperada.

Parâmetros

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

flag_low_confidence

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

Marcador de abstenção opcional (#361): assinala as respostas cuja confiança cai abaixo de min_confidence.

Lê answer_confidence (max(p), a quantidade que as cifras de calibração descrevem e a que não varia com o número de opções), e recorre a confidence se answer_confidence estiver ausente. A resposta e a confiança em bruto mantêm-se intactas; é adicionado low_confidence: True quando a resposta cai abaixo do limiar, e removido se uma resposta antes assinalada agora o ultrapassar (p. ex. quando um dict de resultado é reutilizado ou reavaliado com um limiar diferente).

min_confidence é ou um float (um limiar para cada resposta) ou um mapeamento por cubo (#394), caso em que cada resposta é filtrada no limiar do seu próprio cubo de número de opções via :func:resolve_min_confidence.

Parâmetros

resultsList[Dict[str, Any]]
min_confidencefloat

apply_confidence_gate

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

Relata o estado do gating de confiança, nas respostas às quais um gating foi efetivamente aplicado.

Um gating é uma política, e uma política cuja aplicação não pode ser observada não é uma. Com min_confidence definido, isto escreve abstention -- um dos :data:GATE_STATES -- em cada resposta, mais abstention_threshold, para que quem chama possa responder a três perguntas que de outra forma não consegue:

  • que fração das decisões se absteve, em vez de a inferir do facto de low_confidence ter sido definida por acaso;
  • quantas respostas o gating não conseguiu decidir, o que um booleano não consegue expressar de modo algum;
  • que limiar produziu estes resultados -- flag_low_confidence consome o limiar e descarta-o, pelo que sem isto uma execução em lote com limiares por classe não pode ser subdividida de novo.

GATE_UNEVALUATED é o caso que um booleano não consegue expressar: o gating correu e a resposta não trazia confiança utilizável, pelo que o gating não pôde decidir. Comunicar isso como aprovado é a mesma mentira que comunicá-lo como um marcador.

Com min_confidence não definido, isto não escreve nada. Nem abstention, nem abstention_threshold, nem marcador. É todo o contrato: uma chamada sem gating devolve exatamente o payload que devolvia antes, e a presença do campo -- não um quarto valor lido dele -- é o que diz a quem chama que o gating correu. Chame-o incondicionalmente, uma vez por chamada, no lugar de uma guarda if min_confidence is not None:: essa guarda é o que deixa um caminho a comunicar nada, que é o estado que esta função existe para distinguir.

O marcador em si continua a ser o de :func:flag_low_confidence -- isto delega em vez de reimplementar a regra, para que o booleano e o estado comunicado não possam divergir.

Um min_confidence exatamente 0.0 foi definido, pelo que os estados são comunicados, e :func:flag_low_confidence trata 0.0 como uma operação sem efeito, porque nada pode cair abaixo dele. Toda a resposta que traz uma confiança utilizável, portanto, lê passed, e o eco do limiar é o que distingue isso de uma aprovação real num limiar real.

Parâmetros

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

GATE_STATES

GATE_STATES = (GATE_PASSED, GATE_ABSTAINED, GATE_UNEVALUATED)

Calibração e treino

answer_confidence

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

Massa de probabilidade sobre a resposta que é reportada: max(p).

Esta é a quantidade que o escalonamento de temperatura ajusta, e a quantidade sobre a qual é calculada cada cifra de calibração neste repositório -- ambos os harnesses de benchmark tomam conf = max(probs) antes de chamar ece_score. A secção de gating do README apoia-se na propriedade que a acompanha: das respostas devolvidas com confiança c, cerca de c são corretas. Essa propriedade é condicional, e a condição não é satisfeita por predefinição -- só se verifica depois de as temperaturas terem sido ajustadas e validadas sobre dados reservados para este checkpoint e este número de opções. Os checkpoints distribuídos são demasiado confiantes: choice:11+ é um aguçador de ~10x que devolve uma massa pontual em 1.0, pelo que um limiar aplicado a eles seleciona abaixo da exatidão do modelo (issue #394).

confidence_from_probs mais abaixo reporta uma quantidade diferente numa escala diferente e não dá tal garantia, pelo que as duas não devem ser comparadas contra o mesmo limiar.

Parâmetros

pnp.ndarray
kint

confidence_from_probs

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

Confiança por entropia de Shannon normalizada: 1 - H(p) / log(k).

Quão concentrada está toda a distribuição. É útil, mas não está calibrada: não é o que o escalonamento de temperatura ajusta nem o que o ECE reportado mede. Consulta answer_confidence.

Parâmetros

pnp.ndarray
kint

ece_score

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

Erro de calibração esperado (ECE) ao longo de bins de confiança.

Parâmetros

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

Ajusta um único escalar T por NLL + LBFGS sobre log T.

O resultado é clamp_temperature da escala otimizada, pelo que fica em [TEMP_MIN, TEMP_MAX] (ou é o neutro 1.0 quando o valor não é um número). Devolve 1.0 quando são dados menos de min_n pares. min_n assume por predefinição MIN_BUCKET_N (o mínimo por cubo). Os ajustes ao nível do tipo passam MIN_TYPE_N, que é mais baixo, pelo que um conjunto de dados que não preenche nenhum cubo obtém um escalar em vez de ficar em 1.0.

Parâmetros

pairsSequence
min_nOptional[int]= None

fit_temperature_map

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

Ajusta escalares ao nível do tipo e temperaturas por cubo.

MIN_BUCKET_N é o mínimo por cubo: os cubos mais pequenos são omitidos de temperature_by_options e o escalar ao nível do tipo cobre-os. MIN_TYPE_N é o mínimo separado e mais baixo apenas para esse escalar.

compute_ece=False (o valor predefinido, e o caminho que Agent.fit_temperatures guarda) ajusta sobre todos os registos e não devolve nenhuma chave report. seed é ignorado neste caminho.

compute_ece=True reserva ECE_HOLDOUT_FRAC de cada cubo, estratificado por temp_bucket, usando seed para que os mesmos registos se dividam sempre da mesma forma. As temperaturas são ajustadas apenas sobre o resto e o ECE é pontuado apenas sobre os registos reservados. report["n"] é o número de registos recebidos; report["n_eval"] é a contagem reservada em que o ECE assenta. Um cubo que ficaria abaixo de MIN_BUCKET_N após a reserva é ajustado sobre todos os seus registos, deixado fora do conjunto de avaliação e nomeado em report["buckets_excluded_from_eval"] em vez de ser descartado. n_by_bucket conta sempre a entrada completa, mesmo quando o próprio ajuste usou um subconjunto.

Parâmetros

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]

Ajusta um limiar de abstenção por temp_bucket para que um gating mantenha um erro alvo em cada cubo.

Um único min_confidence não se transfere entre números de opções (#394): as confianças calibradas de uma resposta de 2 opções e de uma de 12 opções vivem em escalas diferentes, pelo que um único corte se abstém a mais ou a menos consoante a pergunta. Isto ajusta um corte por cubo, com chaves exatamente como temperature_by_options (common.temp_bucket, p. ex. "choice:3-5"), e o resultado é um mapa de min_confidence que :func:laya.confidence.check_min_confidence / :func:laya.confidence.apply_confidence_gate aceitam diretamente.

records são as mesmas tuplas (qtype, logits, target[, k]) que fit_temperature_map consome (records_from_labeled constrói-as). A confiança é o max(p) calibrado -- os logits são escalados primeiro pela temperature / temperature_by_options ajustada, pelo que os limiares e os números que o runtime comunica estão na mesma escala. target_error é o erro tolerado entre as respostas aceites; min_bucket_n omite cubos demasiado pequenos para ajustar, e conservative adiciona uma margem de uma amostra. Os limiares são cortes empíricos sobre o conjunto de calibração, não uma garantia formal de cobertura -- valide em dados reservados (fit_temperature_map(..., compute_ece=True) fornece uma divisão reservada) para um gating de produção.

Passe binning_map quando o agent que servirá estes limiares tiver um instalado -- por Agent.fit_binning, ou por um payload de calibração que transporte binning_map -- porque o runtime recalibra answer_confidence através desse mapa antes de qualquer coisa o ler, pelo que um corte ajustado sem ele é um corte numa escala que o gating nunca vê. Os limiares ficam então na escala de binning, e a ordem em que os dois foram ajustados deixa de importar. Medido em 1,200 registos sintéticos de 12 opções com target_error=0.10: o corte ajustado sem um mapa mantém 9.8% de erro sobre 50% de cobertura em confianças não binarizadas, e admite 94.5% das respostas com 25.6% de erro quando o mesmo número é comparado com as binarizadas.

Parâmetros

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]]

Ajusta um mapa de recalibração por binning de histograma por temp_bucket para answer_confidence.

O escalonamento de temperatura aplica um escalar por cubo; não consegue corrigir um cubo cuja curva de fiabilidade não é um simples aguçamento/amaciamento (o choice:11+ patológico que o checkpoint inglês distribuído transporta é um caso). O binning por histograma é a alternativa não paramétrica: divida as confianças calibradas de um cubo em bins bins de largura igual em [0, 1] e mapeie toda a confiança que cai num bin para a exatidão empírica desse bin. Não precisa de nenhuma suposição de monotonicidade nem de dependência extra (apenas NumPy; a regressão isotónica exigiria o scikit-learn).

records são as mesmas tuplas (qtype, logits, target[, k]) que fit_temperature_map consome; a confiança é o max(p) calibrado (logits escalados primeiro pela temperature / temperature_by_options ajustada), pelo que um mapa de binning se compõe sobre um mapa de temperatura em vez de o substituir. Devolve {bucket: {"bins": N, "values": [recalibrated confidence per bin]}}; os cubos abaixo de min_bucket_n são omitidos. Aplique-o com :func:apply_binning_map. Um bin vazio (um intervalo de confiança que o conjunto de calibração nunca produziu) mapeia para o seu próprio ponto médio, ou seja, deixa essa região inalterada, pelo que um valor nunca visto jamais é recalibrado para um 0 fabricado.

Parâmetros

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

Recalibra um answer_confidence para o seu bucket de número de opções (common.temp_bucket).

Devolve a confiança inalterada quando o mapa não tem entrada para o cubo, para que um cubo para o qual o mapa não foi ajustado passe adiante em vez de ser forçado a um valor errado.

Parâmetros

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]

Ajusta um mapa de binning por histograma sobre as temperaturas ajustadas deste agent e guarda-o.

records são as mesmas tuplas (qtype, logits, target[, k]) que fit_temperatures consome. As chaves do mapa são exatamente como as de temperature_by_options, compõe-se sobre as temperaturas atuais e é escrito por save_calibration como binning_map.

Parâmetros

records
min_bucket_nint= MIN_BINNING_BUCKET_N
MIN_BINNING_BUCKET_N

render_options

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

Renderiza os textos de opção por ordem de índice de etiqueta. A ordem semântica de Noul é sempre [false, true].

Parâmetros

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

Recompensa de regra de pontuação estritamente proper: log score + spherical score + ranked probability score.

q: [..., N, K] distribuições reportadas target: [N, K] (distribuições-alvo one-hot ou suaves)

Parâmetros

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

Alvos TD(lambda) para trajetórias de conversa com vários turnos.

Parâmetros

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()}