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) -> strSistema 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]) -> boolTrue quando é de esperar que o checkpoint de inglês consiga ler este estado.
Parâmetros
stateUnion[str, bytes, Mapping, list, None]
clean_email_body
clean_email_body(body: str, max_chars: int = 3000) -> strRemove 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
bodystrmax_charsint=3000
email_state
email_state(
subject: str,
body: str,
sender: Optional[str] = None,
clean: bool = True,
max_chars: int = 3000,
extra,
) -> DictConstró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
subjectstrbodystrsenderOptional[str]=Nonecleanbool=Truemax_charsint=3000extra
Predefinições de perguntas
triage_questions
triage_questions() -> DictPerguntas predefinidas para a triagem de tickets de apoio ao cliente.
email_questions
email_questions(categories: Optional[Dict[str, str]] = None) -> DictPerguntas predefinidas para a triagem de correio recebido e a filtragem de ameaças.
Parâmetros
categoriesOptional[Dict[str, str]]=None
guard_questions
guard_questions() -> DictPerguntas predefinidas para guardrails de entrada de LLM em tempo real.
moderation_questions
moderation_questions() -> DictPerguntas predefinidas para a segurança de conteúdo e a moderação.
router_questions
router_questions() -> DictPerguntas 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,
) -> AnyDevolve 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
stateAnycriteriaAnyembed_fnCallable[[Sequence[str]], Any]kint=DEFAULT_SHORTLIST_KDEFAULT_SHORTLIST_KinstructionsOptional[str]=Nonereturn_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
agentAnystateAnyquestionsDict[str, Dict[str, Any]]embed_fnCallable[[Sequence[str]], Any]kint=DEFAULT_SHORTLIST_KDEFAULT_SHORTLIST_Kpredict_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
agentAnymax_lengthint=512batch_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,
) -> floatO 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) -> NoneMarcador 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,
) -> NoneRelata 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_confidenceter 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_confidenceconsome 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) -> floatMassa 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.ndarraykint
confidence_from_probs
confidence_from_probs(p: np.ndarray, k: int) -> floatConfianç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.ndarraykint
ece_score
ece_score(conf: np.ndarray, correct: np.ndarray, bins: int = 15) -> floatErro de calibração esperado (ECE) ao longo de bins de confiança.
Parâmetros
confnp.ndarraycorrectnp.ndarraybinsint=15
fit_temperatures
fit_temperatures = fit_temperature_mapfit_one_temperature
fit_one_temperature(pairs: Sequence, min_n: Optional[int] = None) -> floatAjusta 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
pairsSequencemin_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
recordsIterablecompute_ecebool=Falseseedint=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
recordsIterabletemperatureSequence[float]temperature_by_optionsDict[str, float]binning_mapOptional[Dict[str, Dict[str, Any]]]=Nonetarget_errorfloat=0.10min_bucket_nint=MIN_ABSTAIN_BUCKET_NMIN_ABSTAIN_BUCKET_Nconservativebool=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
recordsIterabletemperatureSequence[float]temperature_by_optionsDict[str, float]binsint=15min_bucket_nint=MIN_BINNING_BUCKET_NMIN_BINNING_BUCKET_N
apply_binning_map
apply_binning_map(
confidence: float,
bucket: str,
binning_map: Dict[str, Dict[str, Any]],
) -> floatRecalibra 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
confidencefloatbucketstrbinning_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
recordsmin_bucket_nint=MIN_BINNING_BUCKET_NMIN_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.TensorRecompensa 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.Tensortargettorch.Tensorqtypetorch.Tensormasktorch.Tensorw_sphfloat=0.5w_rpsfloat=1.0log_floorfloat=-9.21
td_lambda_targets
td_lambda_targets(p_true: torch.Tensor, batch: Dict, lam: float = 1.0) -> torch.TensorAlvos TD(lambda) para trajetórias de conversa com vários turnos.
Parâmetros
p_truetorch.TensorbatchDictlamfloat=1.0
QTYPES
QTYPES = {"choice": 0, "score": 1, "noul": 2}QTYPE_NAMES
QTYPE_NAMES = {v: k for k, v in QTYPES.items()}