View a markdown version of this page

Creazione di risorse per l'apprendimento per rinforzo in più turni - Amazon SageMaker AI

Le traduzioni sono generate tramite traduzione automatica. In caso di conflitto tra il contenuto di una traduzione e la versione originale in Inglese, quest'ultima prevarrà.

Creazione di risorse per l'apprendimento per rinforzo in più turni

Formato rapido del set di dati

Il set di dati di allenamento è una raccolta di istruzioni che l' SageMaker intelligenza artificiale invia al tuo agente durante l'allenamento. Ogni richiesta dà il via a un lancio: l'agente la elabora, interviene in uno o più turni e restituisce una ricompensa. La qualità e la struttura del set di dati influiscono direttamente su ciò che il modello apprende.

Formati file supportati

Formato Estensione Note
Apache Parquet .parquet Consigliato per set di dati di grandi dimensioni: archiviazione efficiente e caricamento rapido
JSON Lines .jsonl Un oggetto JSON per riga: facile da creare e leggibile dall'uomo
JSON .json Matrice di oggetti JSON
CSV .csv Comma-separated valori con una riga di intestazione

Schema del set di dati

Rilevamento rapido delle colonne

Il servizio RFT rileva la colonna dei prompt utilizzando le seguenti regole, nell'ordine:

  • Se prompt esiste una colonna denominata, viene utilizzata quella colonna.

  • Altrimenti, viene utilizzata la prima colonna del set di dati.

Assegna sempre un nome alla colonna del prompt prompt per evitare ambiguità. Puoi includere colonne aggiuntive per i tuoi scopi di tracciamento, ma solo la colonna prompt viene letta dal servizio RFT.

Come vengono utilizzati i prompt

Il servizio RFT legge la colonna dei prompt e passa il valore della stringa direttamente all'agente così com'è. Non analizza, convalida o trasforma il contenuto. Il formato da utilizzare dipende interamente da ciò che l'agente si aspetta: un agente semplice potrebbe utilizzare testo semplice, mentre uno più sofisticato potrebbe aspettarsi una stringa JSON contenente la cronologia delle conversazioni, la configurazione degli strumenti e le specifiche dei premi.

Protezione dei dati

Poiché il servizio RFT trasmette le istruzioni senza ispezione, è tua responsabilità proteggere i contenuti sensibili. Prendi in considerazione la possibilità di codificare o crittografare i dati richiesti prima di archiviarli e di gestire la decodifica o la decrittografia nel tuo agente.

Approcci comuni:

  • Codifica Base64: semplice offuscamento per dati non sensibili

  • Crittografia: per dati sensibili o proprietari (ad esempio, AES con chiavi gestite dal tuo agente)

Esempio 1: set di dati Q & A semplice (testo normale)

Per attività di formazione semplici con istruzioni in testo semplice.

Caso d'uso: risposta a domande di base, istruzioni semplici seguite

Parquet (Python)

import pyarrow as pa import pyarrow.parquet as pq data = { "prompt": [ "What is 2 + 2?", "Explain the concept of machine learning.", "Write a Python function to reverse a string.", "What is the capital of France?", "How does photosynthesis work?", ] } table = pa.table(data) pq.write_table(table, "training_data.parquet")

Linee JSON (.jsonl)

{"prompt": "What is 2 + 2?"} {"prompt": "Explain the concept of machine learning."} {"prompt": "Write a Python function to reverse a string."}

Esempio 2: con Tool Use Search/Reasoning

Per attività che richiedono l'accesso a strumenti esterni (ad esempio motori di ricerca) durante il ragionamento dei modelli.

Caso d'uso: domande e Fact-based risposte con ricerca sul web, ragionamento basato sul recupero

Struttura:

prompt (column) = JSON string (recommend encoded/encrypted) containing: ├── data_source: Dataset origin identifier ├── prompt: Conversation messages [system, user] ├── ability: Task category (e.g., "fact-reasoning") ├── env_class: "search" ├── reward_spec: Ground truth answer for evaluation └── extra_info: Tool configuration and metadata

Riga di esempio:

import pyarrow as pa import pyarrow.parquet as pq import json task_data = { "data_source": "searchR1_nq", "prompt": [ { "role": "system", "content": "You are a helpful and harmless assistant." }, { "role": "user", "content": "Answer the given question. You must conduct reasoning inside <think> and </think> first every time you get new information. After reasoning, if you find you lack some knowledge, you can call a search engine by <search> query </search> and it will return the top searched results between <information> and </information>. You can search as many times as you want. If you find no further external knowledge needed, you can directly provide the answer inside <answer> and </answer>, without detailed illustrations. For example, <answer> Beijing </answer>. Question: total number of death row inmates in the us?" } ], "ability": "fact-reasoning", "env_class": "search", "reward_spec": { "ground_truth": { "target": [ "2,718" ] }, "style": "rule" }, "extra_info": { "index": 0, "question": "total number of death row inmates in the us?", "split": "train", "need_tools_kwargs": true, "tools_kwargs": { "search": { "create_kwargs": { "question": "total number of death row inmates in the us?", "ground_truth": { "target": [ "2,718" ] }, "data_source": "searchR1_nq" } } } } } # Recommend: encode or encrypt before storing data = {"prompt": [json.dumps(task_data)]} table = pa.table(data) pq.write_table(table, "search_training_data.parquet")

Esempio 3: generazione SQL (Multi-Turn con contesto complesso)

Per attività di generazione di codice che richiedono schemi di database, ragionamento in più fasi e feedback sull'esecuzione SQL.

Caso d'uso: Text-to-SQL generazione di codice con verifica dell'esecuzione

Struttura:

prompt (column) = JSON string (recommend encoded/encrypted) containing: ├── input_seq: Human-readable task description ├── prompt: Conversation messages [system, user] ├── env_class: "text2sql" ├── reward_spec: Ground truth SQL and evaluation config ├── instance_id: Unique task identifier ├── schema: Database schema definition ├── question: Natural language question └── extra_info: Additional metadata

Riga di esempio:

import pyarrow as pa import pyarrow.parquet as pq import json task_data = { "input_seq": "Task Overview:\nYou are a data science expert. Below, you are provided with a database schema\nand a natural language question. Your task is to understand the schema and\ngenerate a valid SQL query to answer the question.\n\nDatabase Engine: SQLite\n\nDatabase Schema:\nCREATE TABLE countries (\n country_id INTEGER PRIMARY KEY,\n english_name TEXT,\n population INTEGER\n);\n\nCREATE TABLE country_metrics (\n metric_id INTEGER PRIMARY KEY,\n country_id INTEGER,\n metric_type TEXT,\n year INTEGER,\n value REAL\n);\n\nQuestion: List all countries with their current population and average\npopulation over the last five years.", "prompt": [ { "role": "system", "content": "Task Overview:\nYou are a data science expert. Your task is to understand the schema and generate\na valid SQL query to answer the question within limited turns.\n\nInstructions:\n- Make sure you only output the information asked in the question.\n- Think through the steps before generating the final SQL query.\n\nFormat:\n- Conduct thinking inside <think>...</think> blocks.\n- You can use SQL tool written within <sql>your sql</sql> to explore or verify.\n- SQL tool output will be shown inside <observation>...</observation>.\n- Provide the final SQL query inside <solution>...</solution>." }, { "role": "user", "content": "Database Schema:\nCREATE TABLE countries (\n country_id INTEGER PRIMARY KEY,\n english_name TEXT,\n population INTEGER\n);\n\nCREATE TABLE country_metrics (\n metric_id INTEGER PRIMARY KEY,\n country_id INTEGER,\n metric_type TEXT,\n year INTEGER,\n value REAL\n);\n\nQuestion: List all countries with their current population and average\npopulation over the last five years." } ], "env_class": "text2sql", "instance_id": "sql_task_001", "reward_spec": { "ground_truth": "SELECT c.english_name, c.population, AVG(m.value) as avg_pop\nFROM countries c\nJOIN country_metrics m ON c.country_id = m.country_id\nWHERE m.metric_type = 'Population' AND m.year > strftime('%Y', 'now') - 5\nGROUP BY c.country_id;", "style": "rule" }, "schema": "CREATE TABLE countries (...); CREATE TABLE country_metrics (...);", "question": "List all countries with their current population...", "extra_info": { "split": "train", "difficulty": "medium" } } # Recommend: encode or encrypt before storing data = {"prompt": [json.dumps(task_data)]} table = pa.table(data) pq.write_table(table, "sql_training_data.parquet")

Best practice

Dimensione del set di dati

Si consiglia un numero minimo di esempi almeno pari atraining_batch_size. 10 volte più della dimensione del batch per motivi di diversità.

Qualità rapida

  • Contesto completo: include tutte le informazioni necessarie affinché il modello generi risposte utili

  • Struttura coerente: mantieni una formattazione coerente in tutti i prompt

  • Evita i duplicati: istruzioni uniche forniscono un segnale di allenamento migliore

  • Istruzioni chiare: per le attività relative all'utilizzo degli strumenti, fornite istruzioni in formato esplicito

Protezione dei dati

  • Codifica o crittografa i contenuti immediati per proteggere i dati sensibili

  • Gestisci le chiavi di decrittografia in modo sicuro sul tuo server di implementazione

  • Il servizio RFT trasmette le istruzioni senza ispezioni, quindi la protezione è una tua responsabilità

Progettazione della funzione di ricompensa

La progettazione delle funzioni di ricompensa è fondamentale per fornire segnali di apprendimento efficaci in sistemi di agenti complessi e in più fasi. Nella progettazione di funzioni di ricompensa per RL a turni multipli, tenete conto delle seguenti linee guida.

  • Inizia con ricompense basate sui risultati. Valuta innanzitutto il risultato finale per stabilire una linea di base pulita e affidabile prima di aggiungere ricompense intermedie o modellare le ricompense.

  • Prendi in considerazione le ricompense continue rispetto alle ricompense binarie. I premi continui possono fornire segnali di credito parziale più chiari, ma sono facili da giocare. Le ricompense binarie sono preferite quando è difficile definire un credito parziale o quando è necessaria una linea di base pulita.

  • Usa Shaping Rewards con attenzione. Le ricompense modellate possono guidare l'apprendimento, ma dovrebbero essere usate con parsimonia, perché una forma troppo forte o disallineata può insegnare scorciatoie.

  • Proteggetevi dalla pirateria informatica basata sulle ricompense. Rendi le ricompense difficili da sfruttare e verifica che il modello stia risolvendo il vero problema anziché aggirare la regola del punteggio.

  • Convalida prima dell'allenamento. Prova la funzione di ricompensa su traiettorie reali prima di allenarti per catturare bug, scappatoie o segnali fuorvianti.

  • Monitora le metriche comportamentali, non solo i premi. Tieni traccia di metriche come la percentuale di completamento, il numero di turni, l'uso degli strumenti e il divario tra gli adattamenti per assicurarti che il modello migliori nel modo previsto.

Processo di progettazione dei premi

  1. Definisci come si presenta il successo e determina se può essere valutato automaticamente.

  2. Valuta il modello base per stabilire una percentuale di successo di base.

  3. Progetta livelli di ricompensa: premi positivi per il successo, zero premi per il fallimento e premi negativi per il comportamento degenerato.

  4. Gestisci i casi limite in modo esplicito, inclusi timeout, errori di ambiente, output non validi e risposte vuote.

  5. Controlla ogni componente della ricompensa per eventuali violazioni delle ricompense.

  6. Convalida su traiettorie reali prima dell'allenamento.

  7. Monitora insieme alle metriche comportamentali durante l'allenamento.

  8. Iterate in base ai risultati iniziali.

In pratica, una funzione di ricompensa prende come input l'intera cronologia dei messaggi di un episodio e restituisce due output: una ricompensa scalare (un punteggio a virgola mobile che misura la qualità della traiettoria, con valori più alti che indicano prestazioni migliori) e un dizionario di metriche per la registrazione, il debug e il monitoraggio.

Esempio: funzione di ricompensa degli agenti di ricerca

L'esempio seguente mostra una funzione di ricompensa per un agente che risponde alle domande utilizzando la ricerca. Dimostra la valutazione dei risultati, la definizione del formato e il controllo della correttezza delle risposte.

class TextAnswerReward: """Reward function to check text answer against gold answers. formula: format_coef * (correct_format - 1) + correct_answer """ gold_answers: list[str] format_coef: float = 0.1 async def __call__(self, history: list[Message]) -> tuple[float, dict[str, float]]: """Grade the completed episode by checking the final assistant message.""" final_message = None for msg in reversed(history): if msg.get("role") == "assistant": final_message = msg break if final_message is None: return 0.0, {"format": 0.0, "correct": 0.0} content = get_text_content(final_message) correct_format = float(self._extract_answer(content) is not None) correct_answer = float(self._check_answer(content)) reward = self.format_coef * (correct_format - 1) + correct_answer return reward, {"format": correct_format, "correct": correct_answer} def _extract_answer(self, text: str) -> str | None: if "Answer:" not in text: return None parts = text.split("Answer:") if len(parts) != 2: return None return parts[1].strip() def _check_answer(self, text: str) -> bool: model_answer = self._extract_answer(text) if model_answer is None or len(self.gold_answers) == 0: return False for gold in self.gold_answers: if normalize_answer(model_answer) == normalize_answer(gold): return True return False

Questa funzione di ricompensa include le seguenti scelte progettuali chiave:

  • La correttezza domina. Una risposta corretta ottiene sempre un punteggio più alto di una risposta errata, indipendentemente dal formato.

  • Il formato è un piccolo segnale di modellazione. Il coefficiente di formato (0,1) è pari al 10% della ricompensa del risultato, abbastanza piccolo da non consentire al modello di trarre vantaggio dalla sola conformità del formato, ma abbastanza grande da orientarlo verso output analizzabili.

  • Un formato errato con una risposta sbagliata è leggermente penalizzato. Il punteggio -0,1 crea un piccolo gradiente che si allontana dagli output completamente non strutturati, senza sovraccaricare il segnale di apprendimento.

  • Nessuna risposta viene considerata errata se il formato è errato. Se il modello non produce mai un messaggio di assistente, la funzione restituisce 0,0, distinguendolo dalla penalità attiva di -0,1 per una risposta attuale ma non valida.