importar clases de datos
importar tokenizadores
importar antorcha
antorcha de importación . nn como nn
antorcha de importación . nn . funcional como F
desde la importación de antorcha Tensor
# Arquitectura del modelo igual que el script de entrenamiento
@ clases de datos . clase de datos
clase Configuración de llama :
"" "Definir los hiperparámetros del modelo Llama." ""
tamaño_vocab : entero = 50000
max_position_embeddings : entero = 2048
tamaño_oculto : entero = 768
tamaño_intermedio : entero = 4 * 768
num_capas_ocultas : entero = 12
num_attention_heads : entero = 12
num_key_value_heads : entero = 3
clase RotaryPositionEncoding ( nn . Módulo ) :
"" "Codificación de posición giratoria." ""
def __init__ ( yo , tenue : entero , max_position_embeddings : int ) -> Ninguno :
súper ( ) . __inicio__ ( )
ser . oscuro = oscuro
ser . max_position_embeddings = max_position _ incrustaciones
norte = 10_000.0
frecuencia_inv = 1.0 / ( N * * ( antorcha . arange ( 0 , tenue , 2 ) / tenue ) )
frecuencia_inv = antorcha . gato ( ( inv_freq , inv_freq ) , tenue = – 1 )
posición = antorcha . arange ( max_position_embeddings )
sinusoide_inp = antorcha . exterior ( posición , frecuencia_inv )
ser . registro_buffer ( "cos" , sinusoid_inp . porque ( ) )
ser . registrarse_buffer ( "pecado" , sinusoid_inp . pecado ( ) )
def adelante ( yo , x : tensor ) -> Tensores :
tamaño_lote , secuencia_len , numero_cabezas , cabeza_dim = x . forma
dispositivo = x . dispositivo
tipo d = x . tipo d
porque = ser . porque . a ( dispositivo , dtype ) [ : seq_len ] . ver ( 1 , secuencia_len , 1 , – 1 )
pecado = ser . pecado . a ( dispositivo , dtype ) [ : seq_len ] . ver ( 1 , secuencia_len , 1 , – 1 )
x1 , x2 = x . trozo ( 2 , tenue = – 1 )
girado = antorcha . gato ( ( – x2 , x1 ) , tenue = – 1 )
devolver ( x * porque ) + ( girado * pecado )
clase LlamaAttention ( nn . Módulo ) :
"" "Atención de consultas agrupadas con incrustaciones rotativas." ""
def __init__ ( yo , configuración : LlamaConfig ) -> Ninguno :
súper ( ) . __inicio__ ( )
ser . tamaño_oculto = configuración . tamaño_oculto
ser . numero_cabezas = configuración . num_atención_cabezas
ser . cabeza_dim = ser . tamaño_oculto // self.num_heads
ser . num_kv_cabezas = configuración . num_key_value_heads
afirmar ( yo . head_dim * ser . núm_cabezas ) == ser . tamaño_oculto
ser . q_proyecto = nn . Lineal ( self . tamaño_oculto , ser . numero_cabezas * ser . cabeza_tenue , sesgo = falso )
ser . k_proj = nn . Lineal ( self . tamaño_oculto , ser . num_kv_heads * ser . cabeza_tenue , sesgo = falso )
ser . v_proyecto = nn . Lineal ( self . tamaño_oculto , ser . num_kv_heads * ser . cabeza_tenue , sesgo = falso )
ser . o_proyecto = nn . Lineal ( self . num_heads * ser . cabeza_tenue , ser . tamaño_oculto , sesgo = falso )
def adelante ( yo , estados_ocultos : tensores , soga : Codificación de posición rotativa ) -> Tensores :
bs , secuencia_len , oscuro = estados_ocultos . tamaño ( )
estados_consulta = ser . q_proj ( estados_ocultos ) . ver ( bs , secuencia_len , ser . numero_cabezas , ser . cabeza_dim )
estados_clave = ser . k_proj ( estados_ocultos ) . ver ( bs , secuencia_len , ser . num_kv_heads , ser . cabeza_dim )
estados_valor = ser . v_proj ( estados_ocultos ) . ver ( bs , secuencia_len , ser . num_kv_heads , ser . cabeza_dim )
atención_salida = F. atención_producto_punto_escalado (
cuerda ( query_states ) . transponer ( 1 , 2 ) ,
cuerda ( key_states ) . transponer ( 1 , 2 ) ,
estados_valor . transponer ( 1 , 2 ) ,
is_causal = Verdadero ,
abandono_p = 0.0 ,
enable_gqa = Verdadero ,
)
atención_salida = atención_salida . transponer ( 1 , 2 ) . remodelar ( bs , secuencia_len , ser . tamaño_oculto )
devolver ser . o_proj ( attn_output )
clase LlamaMLP ( nn . Módulo ) :
"" "Red feed-forward con activación SwiGLU." ""
def __init__ ( yo , configuración : LlamaConfig ) -> Ninguno :
súper ( ) . __inicio__ ( )
ser . puerta_proj = nn . Lineal ( config . tamaño_oculto , configuración . tamaño_intermedio , sesgo = falso )
ser . up_proyecto = nn . Lineal ( config . tamaño_oculto , configuración . tamaño_intermedio , sesgo = falso )
ser . actuar_fn = F. silú
ser . down_proj = nn . Lineal ( config . tamaño_intermedio , configuración . tamaño_oculto , sesgo = falso )
def adelante ( yo , x : tensor ) -> Tensores :
puerta = ser . act_fn ( self . gate_proj ( x ) )
arriba = ser . up_proj ( x )
devolver ser . down_proj ( puerta * arriba )
clase LlamaDecoderLayer ( nn . Módulo ) :
"" "Capa de transformador única para un modelo Llama." ""
def __init__ ( yo , configuración : LlamaConfig ) -> Ninguno :
súper ( ) . __inicio__ ( )
ser . norma_capa_entrada = nn . RMSNorm ( config . tamaño_oculto , pasos = 1e – 5 )
ser . atención propia = LlamaAtención ( config )
ser . post_attention_layernorm = nn . RMSNorm ( config . tamaño_oculto , pasos = 1e – 5 )
ser . mlp = LlamaMLP ( configuración )
def adelante ( yo , estados_ocultos : tensores , soga : Codificación de posición rotativa ) -> Tensores :
residual = estados_ocultos
estados_ocultos = ser . input_layernorm ( estados_ocultos )
atención_salidas = ser . self_attn ( estados_ocultos , cuerda = cuerda )
estados_ocultos = atención_salidas + residual
residual = estados_ocultos
estados_ocultos = ser . post_attention_layernorm ( estados_ocultos )
devolver ser . mlp ( estados_ocultos ) + residual
clase LlamaModel ( nn . Módulo ) :
"" "El modelo Llama completo sin cabezas de entrenamiento previo." ""
def __init__ ( yo , configuración : LlamaConfig ) -> Ninguno :
súper ( ) . __inicio__ ( )
ser . rotativo_emb = Codificación de posición rotativa (
configuración . tamaño_oculto // config.num_attention_heads,
configuración . max_position_embeddings ,
)
ser . tokens_incrustados = nn . Incrustar ( config.vocab_size , configuración . tamaño_oculto )
ser . capas = nn . Lista de módulos ( [
LlamaDecoderLayer ( configuración ) para _ en rango ( config . num_hidden_layers )
] )
ser . norma = nn . RMSNorm ( config . tamaño_oculto , pasos = 1e – 5 )
def adelante ( yo , identificadores de entrada : tensor ) -> Tensores :
estados_ocultos = ser . incrustar_tokens ( id_entrada )
para capa en ser . capas :
estados_ocultos = capa ( estados_ocultos , cuerda = yo . rotativo_emb )
devolver ser . norma ( estados_ocultos )
clase LlamaForPretraining ( nn . Módulo ) :
def __init__ ( yo , configuración : LlamaConfig ) -> Ninguno :
súper ( ) . __inicio__ ( )
ser . modelo_base = LlamaModel ( configuración )
ser . lm_cabeza = nn . Lineal ( config . tamaño_oculto , configuración . tamaño_vocab , sesgo = falso )
def adelante ( yo , identificadores de entrada : tensor ) -> Tensores :
estados_ocultos = ser . modelo_base ( id_entrada )
devolver ser . lm_head ( estados_ocultos )
def aplicar_repetición_penalidad ( logits : tensores , fichas : lista [ int ] , penalización : flotar ) -> Tensores :
"" "Aplicar penalización por repetición a los logits." ""
para tomar en fichas :
si logits [ tok ] > 0 :
logits [ tok ] /= pena
demás :
logits [ tok ] *= pena
devolver logits
@ antorcha . no_graduado ( )
def generar ( modelo , tokenizador , inmediato , tokens_max = 100 , temperatura = 1,0 , penalización_repetición = 1.0 ,
rango_penalización_repetición = 10 , top_k = 50 , dispositivo = Ninguno ) -> cadena :
"" "Generar texto autorregresivamente desde un mensaje.
Argumentos:
modelo: El modelo entrenado LlamaForPretraining
tokenizador: El tokenizador
mensaje: mensaje de entrada de texto
max_tokens: Número máximo de tokens a generar
temperatura: Temperatura de muestreo (más alta = más aleatoria)
repetition_penalty: Penalización por repetir tokens
repetition_penalty_range: Número de tokens anteriores a considerar para la penalización por repetición
top_k: Solo muestra de los k tokens más probables
dispositivo: Dispositivo en el que está cargado el modelo
Devoluciones:
Texto generado
" ""
# Cambiar el modelo al modo de evaluación: la capa de norma funcionará de manera diferente
modelo . evaluar ( )
# Obtener ID de token especiales
bot_id = tokenizador . token_to_id ( "[BOT]" )
eot_id = tokenizador . token_to_id ( "[EOT]" )
# Tokenizar el mensaje en tensor entero
tokens_indicadores = [ bot_id ] + tokenizador . codificar ( " " + inmediato ) . identificaciones
ids_entrada = antorcha . tensor ( [ prompt_tokens ] , tipo d = antorcha . int64 , dispositivo = dispositivo )
# Generar tokens recursivamente
tokens_generados = [ ]
para _entrar rango ( max_tokens ) :
# Modelo de paso directo
logits = modelo ( input_ids )
# Obtener logits para el último token
next_token_logits = logits [ 0 , – 1 , : ] / temperatura
# Aplicar penalización por repetición
si penalización_repetición != 1.0 y len ( tokens_generados ) > 0 :
next_token_logits = aplicar_repetición_penalización (
next_token_logits ,
tokens_generados [ -rango_penalización_repetición : ] ,
penalización_repetición ,
)
# Aplicar filtrado top-k
si top_k > 0 :
top_k_logits = antorcha . topk ( siguiente_token_logits , top_k ) [ 0 ]
índices_para_eliminar = next_token_logits < top_k_logits [ -1 ]
next_token_logits [ índices_a_eliminar ] = flotador ( "-inf" )
# Muestra de la distribución filtrada
problemas = F. softmax ( siguiente_token_logits , tenue = – 1 )
siguiente_token = antorcha . multinomial ( problemas , núm_muestras = 1 )
# Parada anticipada si se genera el token EOT
si token_siguiente . artículo ( ) == eot_id :
romper
# Agregar el nuevo token a input_ids para la próxima iteración
ids_entrada = antorcha . gato ( [ input_ids , token_siguiente . descomprimir ( 0 ) ] , tenue = 1 )
tokens_generados . agregar ( siguiente_token . elemento ( ) )
# Decodificar todos los tokens generados
devolver tokenizador . decodificar ( generado_tokens )
control = "llama_model_final.pth" # punto de control del modelo guardado
tokenizador = "bpe_50K.json" # tokenizador guardado
tokens_max = 100
temperatura = 0,9
top_k = 50
pena = 1.1
rango_pena = 10
# Cargar tokenizador y modelo
dispositivo = antorcha . dispositivo ( "cuda" si antorcha . cuda . está_disponible ( ) demás "UPC" )
tokenizador = tokenizadores . Tokenizador . from_file ( tokenizador )
configuración = LlamaConfig ( )
modelo = LlamaForPretraining ( config ) . a ( dispositivo )
modelo . load_state_dict ( antorcha . cargar ( punto de control , map_location = dispositivo ) )
inmediato = "Había una vez"
respuesta = generar (
modelo = modelo ,
tokenizador = tokenizador ,
aviso = aviso ,
max_tokens = max_tokens ,
temperatura = temperatura ,
top_k = top_k ,
repetición_pena = penalización ,
rango_penalización_repetición = rango_penalización ,
dispositivo = dispositivo ,
)
imprimir ( indicador )
imprimir ( "-" * 20 )
imprimir ( respuesta )