In questo tutorial, realizziamo un'intera pipeline di addestramento con GRPO per insegnare a Gemma-3 a risolvere problemi matematici del dataset GSM8K utilizzando LoRA per mantenere leggero il carico computazionale, nonché funzioni di ricompensa personalizzate per valutare sia la correttezza formattata che l'accuratezza matematica delle risposte. Partiamo dal preparare l'ambiente di sviluppo, autenticandoci su Hugging Face e caricando il modello Gemma-3. Successivamente, formattiamo gli esempi di GSM8K in modi che richiedano un ragionamento strutturato e forniscano una risposta numerica definitiva. Utilizziamo la libreria Tunix, insieme a JAX, Flax e TensorFlow, per creare il workflow end-to-end. Alla fine di questo tutorial, potrai addestrare solo i pesi dei tuoi adattatori mantenendo comunque una configurazione leggera adatta a una singola GPU.

Funzioni di Ricompensa per il Formattaggio

Le funzioni di ricompensa vengono definite per valutare il rispetto delle specifiche di formattazione e la correttezza matematica delle risposte generate. Questo include l'utilizzo di espressioni regolari per verificare il riconoscimento di tag di apertura e chiusura come <reasoning> e </reasoning>, nonché l'accuratezza matematica del risultato, confrontandolo con la risposta corretta.

Quello che otteniamo è una combinazione diversificata di quattro criteri:

    • matchformatexactly: Punteggio di 3.0 se il formato richiesto è esattamente rispettato, altrimenti 0.0.
    • matchformatapproximately: Punteggio basato sull'accuratezza del numero di tag corretti.
    • check_answer: Valutazione della corrispondenza tra la risposta generata e quella corretta, inclusa anche una tolleranza relativa per le approssimazioni decimali.
    • check_numbers: Confronto diretto dei numeri ottenuti, valutando se corrispondono esattamente a quelli desiderati.

Queste funzioni vengono utilizzate in successione per valutare le uscite del modello, dando una misura complessiva del successo dell'addestramento del modello.

Esempio di Processamento del Dataset

Creato un modello di sistema, iniziamo a formattare il dataset GSM8K seguendo uno schema che specifica chiaramente quale parte del testo corrisponde al ragionamento e quale alla risposta. Questo processo si esegue mediante:

    • Un template TPL che include i marcatori di testa e di coda necessari.
    • Una funzione extracthashanswer, utilizzata per estrarre la risposta finale.
    • Una versione mappata e batchata del dataset, utile per l'addestramento parallelo.

I dati vengono divisi in parte di addestramento e di test, con rispettive statistiche generate per fornire un panorama chiaro di quante informazioni sono utilizzate in ogni fase.

Configurazione e Definizione delle Variabili

Un passo fondamentale consiste nella gestione delle variabili di ambiente e nella configurazione del token Hugging Face. Questo permette al sistema di accedere ai modelli necessari, come Gemma-3. Si specificano inoltre:

    • Il numero massimo di parole per prompt.
    • Le opzioni di generazione, ad esempio temperatura, topp, topk.
    • L'architettura della griglia del dispositivo, adatta sia a TPU che a GPU.
    • Parametri di addestramento come learning rate, peso del decadimento, iterazioni e limiti per le prove.

Linguaggio Tecnico e Strumenti Utilizzati

Questo articolo si concentra sull'implementazione in Python e sull'uso di numerose librerie, come:

    • JAX – per gestire il calcolo parallelo e l'ottimizzazione.
    • Flax – strumenti per la creazione e gestione del modello.
    • Qwix – modelli ausiliari per migliorare l'architettura.
    • Grain – per la gestione efficiente del flusso di dati.
    • Safetensors – salvataggio leggero e sicuro dei pesi del modello.
    • TensorFlow – in alcuni aspetti, ma escluso dall'acceleratore.

La configurazione completa dell'ambiente di sviluppo viene eseguita tramite comandi di installazione e gestione del sistema, fornendo una soluzione flessibile e scalabile a seconda della disponibilità hardware.