Passer à la navigation

nemo_rl.models.policy.lm_policy

Afficher en Markdown

Contenu du Module

Classes

NomDescription
PolicyClasse de base abstraite définissant l’interface pour les politiques RL.

Données

PathLike

API

nemo_rl.models.policy.lm_policy.PathLike

Valeur: None

class nemo_rl.models.policy.lm_policy.Policy(cluster: nemo_rl.distributed.virtual_cluster.RayVirtualCluster, config: nemo_rl.models.policy.PolicyConfig, tokenizer: transformers.PreTrainedTokenizerBase, name_prefix: str = 'lm_policy', workers_per_node: typing.Optional[typing.Union[int, list[int]]] = None, init_optimizer: bool = True, weights_path: typing.Optional[nemo_rl.models.policy.lm_policy.PathLike] = None, optimizer_path: typing.Optional[nemo_rl.models.policy.lm_policy.PathLike] = None, init_reference_model: bool = True, processor: typing.Optional[transformers.AutoProcessor] = None)

Bases: nemo_rl.models.policy.interfaces.ColocatablePolicyInterface, nemo_rl.models.generation.interfaces.GenerationInterface

ip: str,
port: int,
) -> list[ray.ObjectRef]

Initialiser la communication collective.

get_logprobs(data: nemo_rl.distributed.batched_data_dict.BatchedDataDict[nemo_rl.models.generation.interfaces.GenerationDatumSpec]) -> nemo_rl.distributed.batched_data_dict.BatchedDataDict[nemo_rl.models.policy.interfaces.LogprobOutputSpec]

Obtenir les logprobs du modèle pour un dictionnaire de données.

Retourne : Un BatchedDataDict avec la clé “logprobs” et la forme [batch_size, sequence_length]. Nous utilisons la convention que le logprob du premier token est 0 afin de maintenir la longueur de séquence. Le logprob du token d’entrée i est spécifié à la position i dans le tenseur de logprobs de sortie.

(The rest of the document follows the same translation pattern)