Passer à la navigation

Interface de Génération

Afficher en Markdown

Ce document explique l’interface de génération de jetons et les différents backends pour le framework NeMo RL. Le système de génération est conçu avec une interface unifiée qui permet à différents backends (comme VLLM, Hugging Face, SGLang et TRT-LLM) de fournir des capacités de génération de jetons tout en respectant la même API.

Interface de Génération

Le cœur du système de génération est défini dans interfaces.py, qui établit une interface abstraite que tous les backends de génération doivent implémenter. Cela garantit la cohérence entre différentes implémentations et facilite le remplacement des backends sans modifier le code appelant.

Composants Clés

  1. GenerationConfig : Un TypedDict qui définit la configuration pour la génération :

    class GenerationConfig(TypedDict):
    """Configuration for generation."""
    backend: str # The backend to use (e.g., "vllm", "hf")
    max_new_tokens: int # Maximum number of tokens to generate
    temperature: float # Sampling temperature
    top_p: float # Top-p sampling parameter
    top_k: int # Top-k sampling parameter
    model_name: str # Name or path of the model
  2. GenerationDatumSpec : Un TypedDict qui définit le format des données d’entrée :

    class GenerationDatumSpec(TypedDict):
    input_ids: torch.Tensor # Input token IDs
    attention_mask: torch.Tensor # Attention mask
    __extra__: Any # Additional data specific to the backend
  3. GenerationOutputSpec : Un TypedDict qui définit le format des données de sortie :

    class GenerationOutputSpec(TypedDict):
    output_ids: torch.Tensor
    generation_lengths: torch.Tensor # Length of just the generated response part
    unpadded_sequence_lengths: torch.Tensor # Length of full valid sequence (input + generated response)
    logprobs: torch.Tensor
    __extra__: Any # Additional output data specific to the backend
  4. GenerationInterface : Une classe de base abstraite que tous les backends de génération doivent implémenter :

    class GenerationInterface(ABC):
    """Abstract base class defining the interface for RL policies."""
    @abstractmethod
    def generate(
    self, data: BatchedDataDict["GenerationDatumSpec"], greedy: bool
    ) -> BatchedDataDict["GenerationOutputSpec"]:
    pass
    @abstractmethod
    def prepare_for_generation(self, *args, **kwargs):
    pass
    @abstractmethod
    def finish_generation(self, *args, **kwargs):
    pass

Un principe de conception clé pour les backends de génération est qu’ils traitent les jetons directement, sans impliquer le tokenizer. En garantissant que seuls les jetons sont échangés, nous éliminons le risque d’incohérences provenant de différentes versions ou spécifications de tokenizer entre les frameworks d’entraînement et de génération.

Backend VLLM

Le backend VLLM (models/generation/vllm/vllm_generation.py) implémente la GenerationInterface pour fournir une génération de texte efficace à l’aide de la bibliothèque VLLM, optimisée pour les grands modèles de langage.

Classe VllmGeneration

La classe VllmGeneration est l’implémentation principale de la GenerationInterface pour VLLM. Elle effectue les fonctions suivantes :

  1. Configure les workers VLLM dans un environnement distribué à l’aide de Ray.
  2. Gère le cycle de vie de ces workers (initialisation, génération, arrêt).
  3. Distribue les entrées aux workers et collecte les sorties.
  4. Gère les mises à jour de poids et la synchronisation.

VllmGenerationWorker

Le VllmGenerationWorker est un acteur Ray qui :

  1. Initialise et gère une instance de modèle VLLM.
  2. Effectue la génération réelle sur un GPU.
  3. Supporte les mises à jour dynamiques de poids via des handles IPC.
  4. Implémente des mécanismes de sommeil/réveil pour une utilisation efficace des ressources.

Extensions VLLM Personnalisées

La classe UpdatableVllmInternalWorker dans vllm_backend.py étend le worker VLLM avec des capacités supplémentaires :

  1. Signalement des ID de périphériques pour permettre le mappage des workers sur des GPU spécifiques.
  2. Mise à jour des poids à partir de handles IPC pour un partage de poids efficace.
  3. Vérification de la mise à jour correcte des poids.

Exemple d’Utilisation

Pour utiliser un backend de génération :

from nemo_rl.algorithms.utils import get_tokenizer
from nemo_rl.distributed.virtual_cluster import RayVirtualCluster
from nemo_rl.distributed.batched_data_dict import BatchedDataDict
from nemo_rl.models.generation.interfaces import configure_generation_config
from nemo_rl.models.generation.vllm import VllmGeneration, VllmConfig
# Set up the configuration
config = VllmConfig(
model_name="Qwen/Qwen2.5-1.5B",
max_new_tokens=100,
temperature=0.7,
top_p=1,
top_k=None,
backend="vllm",
vllm_cfg={
"tensor_parallel_size": 1,
"gpu_memory_utilization": 0.8,
"max_model_len": 2048,
}
)
# Configure config with tokenizer
tokenizer = get_tokenizer(config["model_name"])
config = configure_generation_config(config, tokenizer)
# Initialize the cluster and generation backend
cluster = RayVirtualCluster(...)
generator = VllmGeneration(cluster, config)
# Prepare input data
input_data = BatchedDataDict(...)
# Generate text
generator.prepare_for_generation()
output = generator.generate(input_data, greedy=False)
generator.finish_generation()

Étendre avec de Nouveaux Backends

Pour ajouter un nouveau backend de génération :

  1. Créez une nouvelle classe qui implémente GenerationInterface.
  2. Implémentez les méthodes requises : generate, prepare_for_generation et finish_generation.
  3. Assurez-vous que votre implémentation fonctionne avec les structures standard GenerationConfig et GenerationDatumSpec.
  4. Enregistrez votre backend avec le système (si nécessaire) pour le rendre accessible.

Cette conception modulaire permet une extension facile avec de nouveaux backends tout en maintenant une interface cohérente pour le reste du système.