Una suave introducción al ajuste del modelo lingüístico

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 )