class Lm():
    def __init__(self, arg):
        # 创建一个TensorFlow计算图
        self.graph = tf.Graph()
        with self.graph.as_default():
            # 初始化参数
            self.is_training = arg.is_training
            self.hidden_units = arg.hidden_units
            self.input_vocab_size = arg.input_vocab_size
            self.label_vocab_size = arg.label_vocab_size
            self.num_heads = arg.num_heads
            self.num_blocks = arg.num_blocks
            self.max_length = arg.max_length
            self.lr = arg.lr
            self.dropout_rate = arg.dropout_rate
            
            # 创建两个占位符,用于输入和标签
            self.x = tf.placeholder(tf.int32, shape=(None, None))
            self.y = tf.placeholder(tf.int32, shape=(None, None))
            
            # 对输入进行embedding
            self.emb = embedding(self.x, vocab_size=self.input_vocab_size, num_units=self.hidden_units, scale=True, scope='enc_embed')
            
            # 对输入进行位置编码
            self.enc = self.emb + embedding(tf.tile(tf.expand_dims(tf.range(tf.shape(self.x)[1]), 0), [tf.shape(self.x)[0], 1]),
              vocab_size=self.max_length,num_units=self.hidden_units, zero_pad=False, scale=False,scope='enc_pe')
            
            # 对位置编码后的输入进行dropout
            self.enc = tf.layers.dropout(self.enc, 
                                        rate=self.dropout_rate, 
                                        training=tf.convert_to_tensor(self.is_training))
            
            # 进行多个块的Transformer编码
            for i in range(self.num_blocks):
                with tf.variable_scope('num_blocks_{}'.format(i)):
                    ### Multihead Attention
                    # 对输入进行多头注意力
                    self.enc = multihead_attention(emb = self.emb,
                                                    queries=self.enc, 
                                                    keys=self.enc, 
                                                    num_units=self.hidden_units, 
                                                    num_heads=self.num_heads, 
                                                    dropout_rate=self.dropout_rate,
                                                    is_training=self.is_training,
                                                    causality=False)
                    
            # 对编码后的输出进行前馈神经网络处理
            self.outputs = feedforward(self.enc, num_units=[4*self.hidden_units, self.hidden_units])
            
            # 对前馈神经网络的输出进行全连接层处理,得到预测结果
            self.logits = tf.layers.dense(self.outputs, self.label_vocab_size)
            
            # 对预测结果进行softmax,得到最终预测
            self.preds = tf.to_int32(tf.argmax(self.logits, axis=-1))
            
            # 进行准确率计算
            self.istarget = tf.to_float(tf.not_equal(self.y, 0))
            self.acc = tf.reduce_sum(tf.to_float(tf.equal(self.preds, self.y))*self.istarget)/ (tf.reduce_sum(self.istarget))
            tf.summary.scalar('acc', self.acc)
            
            if self.is_training:  
                # 进行标签平滑处理,并计算损失
                self.y_smoothed = label_smoothing(tf.one_hot(self.y, depth=self.label_vocab_size))
                self.loss = tf.nn.softmax_cross_entropy_with_logits_v2(logits=self.logits, labels=self.y_smoothed)
                self.mean_loss = tf.reduce_sum(self.loss*self.istarget) / (tf.reduce_sum(self.istarget))
                
                # 定义全局步数变量
                self.global_step = tf.Variable(0, name='global_step', trainable=False)
                
                # 定义优化器
                self.optimizer = tf.train.AdamOptimizer(learning_rate=self.lr, beta1=0.9, beta2=0.98, epsilon=1e-8)
                
                # 进行训练操作
                self.train_op = self.optimizer.minimize(self.mean_loss, global_step=self.global_step)
                
                # 记录损失和准确率
                tf.summary.scalar('mean_loss', self.mean_loss)
                self.merged = tf.summary.merge_all()
Transformer Language Model Implementation in TensorFlow

原文地址: https://www.cveoy.top/t/topic/nBlz 著作权归作者所有。请勿转载和采集!

免费AI点我,无需注册和登录