Los modelos de lenguaje grande se componen de miles de millones de parámetros (pesos). Para cada palabra que genera, el modelo tiene que realizar cálculos computacionalmente costosos en todos estos parámetros.
Los modelos de lenguaje grande aceptan una oración o secuencia de tokens y generan una distribución de probabilidad del siguiente token más probable.
Por lo tanto, normalmente la decodificación norte tokens (o generar norte palabras del modelo) requiere ejecutar el modelo norte número de veces. En cada iteración, el nuevo token se agrega a la oración de entrada y se pasa nuevamente al modelo. Esto puede resultar costoso.
Además, la estrategia de decodificación puede influir en la calidad de las palabras generadas. Generar tokens de forma sencilla, simplemente tomando el token con mayor probabilidad en la distribución de salida, puede resultar en texto repetitivo. El muestreo aleatorio de la distribución puede provocar una deriva involuntaria.
Por lo tanto, se requiere una estrategia de decodificación sólida para garantizar ambos:
- Resultados de alta calidad
- Tiempo de inferencia rápido
Ambos requisitos pueden abordarse utilizando una combinación de un modelo de lenguaje grande y pequeño, siempre que los modelos amateur y experto sean similares (por ejemplo, la misma arquitectura pero diferentes tamaños).
- Modelo objetivo/grande: LM principal con mayor número de parámetros (por ejemplo, OPT-13B)
- Modelo aficionado/pequeño: Versión más pequeña de Main LM con menos parámetros (por ejemplo, OPT-125M)
Especulativo y contrastivo La decodificación aprovecha los LLM grandes y pequeños para lograr una generación de texto confiable y eficiente.
Decodificación contrastiva es una estrategia que explota el hecho de que las fallas en los LLM grandes (como la repetición, la incoherencia) son aún más pronunciadas en los LLM pequeños. Por lo tanto, esta estrategia optimiza los tokens con la mayor diferencia de probabilidad entre el modelo pequeño y grande.
Para una única predicción, la decodificación contrastiva genera dos distribuciones de probabilidad:
- q = probabilidades logit para el modelo amateur
- pag = probabilidades logit para el modelo experto
El siguiente token se elige según los siguientes criterios:
- Descartar todos los tokens que no tengan una probabilidad suficientemente alta según el modelo experto (descartar p(x) < alfa * máx(p))
- De los tokens restantes, seleccione el que tenga la mayor diferencia entre las probabilidades de registro del modelo grande y del modelo pequeño, máx(p(x) – q(x)).
Implementación de decodificación contrastiva
from transformers import AutoTokenizer, AutoModelForCausalLM
import torch# Load models and tokenizer
tokenizer = AutoTokenizer.from_pretrained('gpt2')
amateur_lm = AutoModelForCausalLM.from_pretrained('gpt2')
expert_lm = AutoModelForCausalLM.from_pretrained('gpt2-large')
def contrastive_decoding(prompt, max_length=50):
input_ids = tokenizer(prompt, return_tensors="pt").input_ids
while input_ids.shape[1] < max_length:
# Generate amateur model output
amateur_outputs = amateur_lm(input_ids, return_dict=True)
amateur_logits = torch.softmax(amateur_outputs.logits[:, -1, :], dim=-1)
log_probs_amateur = torch.log(amateur_logits)
# Generate expert model output
expert_outputs = expert_lm(input_ids, return_dict=True)
expert_logits = torch.softmax(expert_outputs.logits[:, -1, :], dim=-1)
log_probs_exp = torch.log(expert_logits)
log_probs_diff = log_probs_exp - log_probs_amateur
# Set an alpha threshold to eliminate less confident tokens in expert
alpha = 0.1
candidate_exp_prob = torch.max(expert_logits)
# Mask tokens below threshold for expert model
V_head = expert_logits < alpha * candidate_exp_prob
# Select the next token from the log-probabilities difference, ignoring masked values
token = torch.argmax(log_probs_diff.masked_fill(V_head, -torch.inf)).unsqueeze(0)
# Append token and accumulate generated text
input_ids = torch.cat([input_ids, token.unsqueeze(1)], dim=-1)
return tokenizer.batch_decode(input_ids)
prompt = "Large Language Models are"
generated_text = contrastive_decoding(prompt, max_length=25)
print(generated_text)
Decodificación especulativa se basa en el principio de que el modelo más pequeño debe tomar muestras de la misma distribución que el modelo más grande. Por lo tanto, esta estrategia apunta a aceptar tantas predicciones del modelo más pequeño como sea posible, siempre que se alineen con la distribución del modelo más grande.
El modelo más pequeño genera norte fichas en secuencia, como posibles conjeturas. Sin embargo, todos norte Las secuencias se introducen en el modelo experto más grande como un solo lote, que es más rápido que la generación secuencial.
Esto da como resultado un caché para cada modelo, con norte distribuciones de probabilidad en cada caché.
- q = probabilidades logit para el modelo amateur
- pag = probabilidades logit para el modelo experto
A continuación, los tokens muestreados del modelo amateur se aceptan o rechazan según las siguientes condiciones:
- Si la probabilidad del token es mayor en la distribución experta (p) que en la distribución amateur (q), o p(x) > q(x), aceptar token
- Si la probabilidad del token es menor en la distribución experta (p) que en la distribución amateur (q), o p(x)
rechazar token con probabilidad 1 – p(x) / q(x)
Si se rechaza un token, se toma una muestra del siguiente token de la distribución experta o de la distribución ajustada. Además, el modelo aficionado y experto restablece el caché y vuelve a generar norte conjeturas y distribuciones de probabilidad pag y q.
Implementación de decodificación especulativa
from transformers import AutoTokenizer, AutoModelForCausalLM
import torch# Load models and tokenizer
tokenizer = AutoTokenizer.from_pretrained('gpt2')
amateur_lm = AutoModelForCausalLM.from_pretrained('gpt2')
expert_lm = AutoModelForCausalLM.from_pretrained('gpt2-large')
# Sample next token from output distribution
def sample_from_distribution(logits):
sampled_index = torch.multinomial(logits, 1)
return sampled_index
def generate_cache(input_ids, n_tokens):
# Store logits at each step for amateur and expert models
amateur_logits_per_step = []
generated_tokens = []
batch_input_ids = []
with torch.no_grad():
for _ in range(n_tokens):
# Generate amateur model output
amateur_outputs = amateur_lm(input_ids, return_dict=True)
amateur_logits = torch.softmax(amateur_outputs.logits[:, -1, :], dim=-1)
amateur_logits_per_step.append(amateur_logits)
# Sampling from amateur logits
next_token = sample_from_distribution(amateur_logits)
generated_tokens.append(next_token)
# Append to input_ids for next generation step
input_ids = torch.cat([input_ids, next_token], dim=-1)
batch_input_ids.append(input_ids.squeeze(0))
# Feed IDs to expert model as batch
batched_input_ids = torch.nn.utils.rnn.pad_sequence(batch_input_ids, batch_first=True, padding_value=0 )
expert_outputs = expert_lm(batched_input_ids, return_dict=True)
expert_logits = torch.softmax(expert_outputs.logits[:, -1, :], dim=-1)
return amateur_logits_per_step, expert_logits, torch.cat(generated_tokens, dim=-1)
def speculative_decoding(prompt, n_tokens=5, max_length=50):
input_ids = tokenizer(prompt, return_tensors="pt").input_ids
while input_ids.shape[1] < max_length:
amateur_logits_per_step, expert_logits, generated_ids = generate_cache(
input_ids, n_tokens
)
accepted = 0
for n in range(n_tokens):
token = generated_ids[:, n][0]
r = torch.rand(1).item()
# Extract probabilities
p_x = expert_logits[n][token].item()
q_x = amateur_logits_per_step[n][0][token].item()
# Speculative decoding acceptance criterion
if ((q_x > p_x) and (r > (1 - p_x / q_x))):
break # Reject token and restart the loop
else:
accepted += 1
# Check length
if (input_ids.shape[1] + accepted) >= max_length:
return tokenizer.batch_decode(input_ids)
input_ids = torch.cat([input_ids, generated_ids[:, :accepted]], dim=-1)
if accepted < n_tokens:
diff = expert_logits[accepted] - amateur_logits_per_step[accepted][0]
clipped_diff = torch.clamp(diff, min=0)
# Sample a token from the adjusted expert distribution
normalized_result = clipped_diff / torch.sum(clipped_diff, dim=0, keepdim=True)
next_token = sample_from_distribution(normalized_result)
input_ids = torch.cat([input_ids, next_token.unsqueeze(1)], dim=-1)
else:
# Sample directly from the expert logits for the last accepted token
next_token = sample_from_distribution(expert_logits[-1])
input_ids = torch.cat([input_ids, next_token.unsqueeze(1)], dim=-1)
return tokenizer.batch_decode(input_ids)
# Example usage
prompt = "Large Language models are"
generated_text = speculative_decoding(prompt, n_tokens=3, max_length=25)
print(generated_text)
Evaluación
Podemos evaluar ambos enfoques de decodificación comparándolos con un método de decodificación ingenuo, donde elegimos aleatoriamente el siguiente token de la distribución de probabilidad.
def sequential_sampling(prompt, max_length=50):
"""
Perform sequential sampling with the given model.
"""
# Tokenize the input prompt
input_ids = tokenizer(prompt, return_tensors="pt").input_idswith torch.no_grad():
while input_ids.shape[1] < max_length:
# Sample from the model output logits for the last token
outputs = expert_lm(input_ids, return_dict=True)
logits = outputs.logits[:, -1, :]
probabilities = torch.softmax(logits, dim=-1)
next_token = torch.multinomial(probabilities, num_samples=1)
input_ids = torch.cat([input_ids, next_token], dim=-1)
return tokenizer.batch_decode(input_ids)
Para evaluar la decodificación contrastiva, podemos utilizar las siguientes métricas de riqueza léxica.
- Entropía de n-gramas: Mide la imprevisibilidad o diversidad de n-gramas en el texto generado. Una entropía alta indica un texto más diverso, mientras que una entropía baja sugiere repetición o previsibilidad.
- distinto-n: Mide la proporción de n-gramas únicos en el texto generado. Los valores de n distintos más altos indican una mayor diversidad léxica.
from collections import Counter
import mathdef ngram_entropy(text, n):
"""
Compute n-gram entropy for a given text.
"""
# Tokenize the text
tokens = text.split()
if len(tokens) < n:
return 0.0 # Not enough tokens to form n-grams
# Create n-grams
ngrams = [tuple(tokens[i:i + n]) for i in range(len(tokens) - n + 1)]
# Count frequencies of n-grams
ngram_counts = Counter(ngrams)
total_ngrams = sum(ngram_counts.values())
# Compute entropy
entropy = -sum((count / total_ngrams) * math.log2(count / total_ngrams)
for count in ngram_counts.values())
return entropy
def distinct_n(text, n):
"""
Compute distinct-n metric for a given text.
"""
# Tokenize the text
tokens = text.split()
if len(tokens) < n:
return 0.0 # Not enough tokens to form n-grams
# Create n-grams
ngrams = [tuple(tokens[i:i + n]) for i in range(len(tokens) - n + 1)]
# Count unique and total n-grams
unique_ngrams = set(ngrams)
total_ngrams = len(ngrams)
return len(unique_ngrams) / total_ngrams if total_ngrams > 0 else 0.0
prompts = [
"Large Language models are",
"Barack Obama was",
"Decoding strategy is important because",
"A good recipe for Halloween is",
"Stanford is known for"
]
# Initialize accumulators for metrics
naive_entropy_totals = [0, 0, 0] # For n=1, 2, 3
naive_distinct_totals = [0, 0] # For n=1, 2
contrastive_entropy_totals = [0, 0, 0]
contrastive_distinct_totals = [0, 0]
for prompt in prompts:
naive_generated_text = sequential_sampling(prompt, max_length=50)[0]
for n in range(1, 4):
naive_entropy_totals[n - 1] += ngram_entropy(naive_generated_text, n)
for n in range(1, 3):
naive_distinct_totals[n - 1] += distinct_n(naive_generated_text, n)
contrastive_generated_text = contrastive_decoding(prompt, max_length=50)[0]
for n in range(1, 4):
contrastive_entropy_totals[n - 1] += ngram_entropy(contrastive_generated_text, n)
for n in range(1, 3):
contrastive_distinct_totals[n - 1] += distinct_n(contrastive_generated_text, n)
# Compute averages
naive_entropy_averages = [total / len(prompts) for total in naive_entropy_totals]
naive_distinct_averages = [total / len(prompts) for total in naive_distinct_totals]
contrastive_entropy_averages = [total / len(prompts) for total in contrastive_entropy_totals]
contrastive_distinct_averages = [total / len(prompts) for total in contrastive_distinct_totals]
# Display results
print("Naive Sampling:")
for n in range(1, 4):
print(f"Average Entropy (n={n}): {naive_entropy_averages[n - 1]}")
for n in range(1, 3):
print(f"Average Distinct-{n}: {naive_distinct_averages[n - 1]}")
print("\nContrastive Decoding:")
for n in range(1, 4):
print(f"Average Entropy (n={n}): {contrastive_entropy_averages[n - 1]}")
for n in range(1, 3):
print(f"Average Distinct-{n}: {contrastive_distinct_averages[n - 1]}")
Los siguientes resultados nos muestran que la decodificación contrastiva supera al muestreo ingenuo para estas métricas.
Muestreo ingenuo:
Entropía promedio (n=1): 4.990499826537679
Entropía promedio (n=2): 5,174765791328267
Entropía promedio (n=3): 5.14373124004409
Promedio Distinto-1: 0.8949694135740648
Promedio Distinto-2: 0,9951219512195122Decodificación contrastiva:
Entropía promedio (n=1): 5,182773920916605
Entropía promedio (n=2): 5.3495681172235665
Entropía promedio (n=3): 5.313720275712986
Promedio Distinto-1: 0.9028425204970866
Promedio Distinto-2: 1.0
Para evaluar la decodificación especulativa, podemos observar el tiempo de ejecución promedio de un conjunto de indicaciones para diferentes norte valores.
import time
import matplotlib.pyplot as plt# Parameters
n_tokens = range(1, 11)
speculative_decoding_times = []
naive_decoding_times = []
prompts = [
"Large Language models are",
"Barack Obama was",
"Decoding strategy is important because",
"A good recipe for Halloween is",
"Stanford is known for"
]
# Loop through n_tokens values
for n in n_tokens:
avg_time_naive, avg_time_speculative = 0, 0
for prompt in prompts:
start_time = time.time()
_ = sequential_sampling(prompt, max_length=25)
avg_time_naive += (time.time() - start_time)
start_time = time.time()
_ = speculative_decoding(prompt, n_tokens=n, max_length=25)
avg_time_speculative += (time.time() - start_time)
naive_decoding_times.append(avg_time_naive / len(prompts))
speculative_decoding_times.append(avg_time_speculative / len(prompts))
avg_time_naive = sum(naive_decoding_times) / len(naive_decoding_times)
# Plotting the results
plt.figure(figsize=(8, 6))
plt.bar(n_tokens, speculative_decoding_times, width=0.6, label='Speculative Decoding Time', alpha=0.7)
plt.axhline(y=avg_time_naive, color='red', linestyle='--', label='Naive Decoding Time')
# Labels and title
plt.xlabel('n_tokens', fontsize=12)
plt.ylabel('Average Time (s)', fontsize=12)
plt.title('Speculative Decoding Runtime vs n_tokens', fontsize=14)
plt.legend()
plt.grid(axis='y', linestyle='--', alpha=0.7)
# Show the plot
plt.show()
plt.savefig("plot.png")
Podemos ver que el tiempo de ejecución promedio para la decodificación ingenua es mucho mayor que para la decodificación especulativa. norte valores.
La combinación de modelos de lenguaje grandes y pequeños para la decodificación logra un equilibrio entre calidad y eficiencia. Si bien estos enfoques introducen una complejidad adicional en el diseño del sistema y la gestión de recursos, sus beneficios se aplican a la IA conversacional, la traducción en tiempo real y la creación de contenido.
Estos enfoques requieren una cuidadosa consideración de las limitaciones de implementación. Por ejemplo, las demandas adicionales de memoria y computación al ejecutar modelos duales pueden limitar la viabilidad en los dispositivos de borde, aunque esto se puede mitigar mediante técnicas como la cuantificación de modelos.
A menos que se indique lo contrario, todas las imágenes son del autor.