在前面的文章中,已经详细介绍了现代LLM(大语言模型)几乎都是基于Transformer架构实现的,并对Transformer的核心原理、结构组成以及工作流程进行了深入讲解。从多头注意力机制到前馈网络,从残差连接到层归一化,这些基础组件共同构成了当今大语言模型的基石。

这篇文章则更进一步从理论走向实践,基于Transformer架构亲自动手搭建一个小型LLM,并完成预训练(Pre-training) 和有监督微调(Supervised Fine-Tuning,SFT) 两个核心阶段。通过这个过程,不仅能够深入理解LLM的内部运作机制,还能掌握从零开始训练一个可用模型的全流程。

这篇文章的目标不是训练一个完美的大模型,而是在有限的计算资源下,跑通从架构实现到模型训练再到推理使用的完整链路,从而对LLM的训练全流程有一个直观而扎实的理解。

二、实现一个LLM架构

在前面的文档中我们提到过,LLaMA系列模型是开源社区的基石,其生态极其繁荣,围绕它社区衍生出了丰富的微调方案、部署工具和各种改进模型(如Alpaca、Vicuna、Chinese-LLaMA等)。因此,这篇文章也将以LLaMA架构为基础进行手动实现。

2.1 架构图

首先,一览LLaMA的整体架构设计图,将以此为准进行代码实现:

从架构图中可以看出,LLaMA采用了仅解码器(Decoder-only) 的Transformer结构,这与GPT系列一致。整个模型由嵌入层、多个Decoder层堆叠、以及最终的输出层组成。每个Decoder层内部包含自注意力机制和前馈网络,并使用了前置层归一化(Pre-normalization) 和残差连接,这种设计有助于提升训练的稳定性和收敛速度。

2.2 模型参数定义

在开始编写模型代码之前,首先需要定义一组超参数。这些超参数决定了模型的大小、容量和计算复杂度,是模型设计的第一步。关键的超参数包括:模型维度(即隐藏层大小)、Transformer层数、注意力头数、词汇表大小、最大序列长度等。这些参数需要根据实际的任务需求、数据集规模和可用的计算资源来综合确定。

自定义一个ModelConfig类来统一存储和管理这些超参数。该类继承自transformers库中的PretrainedConfig,这样做的好处是:一方面可以复用transformers库提供的便捷功能,另一方面也便于后续将模型导出为Hugging Face标准格式,从而与整个社区生态无缝对接。

from transformers import PretrainedConfigclass ModelConfig(PretrainedConfig):    model_type = "My-LLM"# 模型ID    def __init__(            self,            dim: int = 768, # 模型维度            n_layers: int = 12, # Transformer的层数            n_heads: int = 16, # 注意力机制的头数            n_kv_heads: int = 8, # 键值头的数量            vocab_size: int = 6144, # 词汇表大小            hidden_dim: int = None, # 隐藏层维度            multiple_of: int = 64,             norm_eps: float = 1e-5, # 归一化层的eps            max_seq_len: int = 512, # 模型最大输入序列长度            dropout: float = 0.0, # dropout概率            flash_attn: bool = True, # 是否使用Flash Attention            **kwargs,    ):        self.dim = dim        self.n_layers = n_layers        self.n_heads = n_heads        self.n_kv_heads = n_kv_heads        self.vocab_size = vocab_size        self.hidden_dim = hidden_dim        self.multiple_of = multiple_of        self.norm_eps = norm_eps        self.max_seq_len = max_seq_len        self.dropout = dropout        self.flash_attn = flash_attn        super().__init__(**kwargs)

下面解释其中几个核心超参数的含义:

  • dim(模型维度):即Transformer内部各层的隐藏维度大小,决定了模型的"宽度"。维度越大,模型的表示能力越强,但计算开销和显存占用也越大。
  • n_layers(层数):Transformer解码器的堆叠层数,决定了模型的"深度"。层数越多,模型能够捕捉的特征层次越丰富,但也越容易出现过拟合和训练困难的问题。
  • n_heads(注意力头数):多头注意力机制中的头数。每个头会从不同的子空间学习不同的注意力模式,多头并行可以增强模型的表达能力。通常要求dim能被n_heads整除。
  • vocab_size(词汇表大小):词表的大小,决定了模型能够表示的不同Token的总数量。这个值需要与使用的Tokenizer保持一致。
  • max_seq_len(最大序列长度):模型能够处理的最大输入序列长度(以Token为单位)。超过这个长度的序列会被截断或丢弃,这是Transformer架构中位置编码的固有局限。

2.3 归一化层定义

归一化层是深度神经网络中稳定训练过程的关键组件。这里使用的是RMSNorm(Root Mean Square Normalization)机制,它是LayerNorm的一种变体。与标准的LayerNorm相比,RMSNorm去掉了均值中心化的步骤,只对输入的均方根进行缩放,计算更加高效,同时在一些大规模模型中表现出了不逊色甚至更好的效果。

RMSNorm的数学公式如下:

这种归一化方式有助于防止各层输出的数值规模过大或过小,从而稳定梯度传播,特别是在具有数十甚至上百层的深度模型中,RMSNorm对于缓解梯度消失和梯度爆炸问题有着显著作用。

通过如下代码实现RMSNorm

class RMSNorm(nn.Module):    def __init__(self, dim: int, eps: float):        super().__init__()        # eps是为了防止除以0的情况        self.eps = eps        # weight是一个可学习的参数,全部初始化为1        self.weight = nn.Parameter(torch.ones(dim))    def _norm(self, x):        # 计算RMSNorm的核心部分        # x.pow(2).mean(-1, keepdim=True)计算了输入x的平方的均值        # torch.rsqrt是平方根的倒数,这样就得到了RMSNorm的分母部分,再加上eps防止分母为0        # 最后乘以x,得到RMSNorm的结果        return x * torch.rsqrt(x.pow(2).mean(-1, keepdim=True) + self.eps)    def forward(self, x):        # forward函数是模型的前向传播        # 首先将输入x转为float类型,然后进行RMSNorm,最后再转回原来的数据类型        # 最后乘以weight,这是RMSNorm的一个可学习的缩放因子        output = self._norm(x.float()).type_as(x)        return output * self.weight

2.4 注意力层定义

注意力机制是Transformer架构的核心所在。标准的Transformer使用多头注意力(Multi-Head Attention,MHA),而在LLaMA系列中,除了LLaMA2-70B等超大模型使用了分组查询注意力(Grouped-Query Attention,GQA)外,LLaMA2-70B以下的模型仍然使用MHA。

这里选择GQA来构建注意力模块。GQAMHA多查询注意力(Multi-Query Attention,MQA)之间的一种折中方案:它将查询(Query)头分成若干组,每组共享一对键(Key)和值(Value)投影。这样做既能比MHA减少KV Cache的显存占用,提高推理效率,又能比MQA保持更好的模型质量,这对于资源受限的场景尤其有价值,是当前大模型推理优化的热门方向。

即使训练的模型规模不大,采用 GQA 也有助于提前熟悉这一重要的工程优化技术,同时为后续迁移到更大规模的模型打下基础。

GQA 注意力层的实现代码如下:

class Attention(nn.Module):    def __init__(self, args: ModelConfig):        super().__init__()        # 根据是否指定n_kv_heads,确定用于键(key)和值(value)的头的数量。        self.n_kv_heads = args.n_heads if args.n_kv_heads isNoneelse args.n_kv_heads        # 确保总头数可以被键值头数整除。        assert args.n_heads % self.n_kv_heads == 0        # 模型并行处理大小,默认为1。        model_parallel_size = 1        # 本地计算头数,等于总头数除以模型并行处理大小。        self.n_local_heads = args.n_heads // model_parallel_size        # 本地键值头数,等于键值头数除以模型并行处理大小。        self.n_local_kv_heads = self.n_kv_heads // model_parallel_size        # 重复次数,用于扩展键和值的尺寸。        self.n_rep = self.n_local_heads // self.n_local_kv_heads        # 每个头的维度,等于模型维度除以头的总数。        self.head_dim = args.dim // args.n_heads        # 定义权重矩阵。        self.wq = nn.Linear(args.dim, args.n_heads * self.head_dim, bias=False)        self.wk = nn.Linear(args.dim, self.n_kv_heads * self.head_dim, bias=False)        self.wv = nn.Linear(args.dim, self.n_kv_heads * self.head_dim, bias=False)        # 输出权重矩阵。        self.wo = nn.Linear(args.n_heads * self.head_dim, args.dim, bias=False)        # 定义dropout。        self.attn_dropout = nn.Dropout(args.dropout)        self.resid_dropout = nn.Dropout(args.dropout)        # 保存dropout概率。        self.dropout = args.dropout        # 检查是否使用Flash Attention(需要PyTorch >= 2.0)。        self.flash = hasattr(torch.nn.functional, 'scaled_dot_product_attention')        ifnot self.flash:            # 若不支持Flash Attention,则使用手动实现的注意力机制,并设置mask。            print("WARNING: using slow attention. Flash Attention requires PyTorch >= 2.0")            # 创建一个上三角矩阵,用于遮蔽未来信息。            mask = torch.full((1, 1, args.max_seq_len, args.max_seq_len), float("-inf"))            mask = torch.triu(mask, diagonal=1)            # 注册为模型的缓冲区            self.register_buffer("mask", mask)    def forward(self, x: torch.Tensor, freqs_cos: torch.Tensor, freqs_sin: torch.Tensor):        # 获取批次大小和序列长度,[batch_size, seq_len, dim]        bsz, seqlen, _ = x.shape        # 计算查询(Q)、键(K)、值(V)。        xq, xk, xv = self.wq(x), self.wk(x), self.wv(x)        # 调整形状以适应头的维度。        xq = xq.view(bsz, seqlen, self.n_local_heads, self.head_dim)        xk = xk.view(bsz, seqlen, self.n_local_kv_heads, self.head_dim)        xv = xv.view(bsz, seqlen, self.n_local_kv_heads, self.head_dim)        # 应用旋转位置嵌入(RoPE)。        xq, xk = apply_rotary_emb(xq, xk, freqs_cos, freqs_sin)        # 对键和值进行扩展以适应重复次数。        xk = repeat_kv(xk, self.n_rep)        xv = repeat_kv(xv, self.n_rep)        # 将头作为批次维度处理。        xq = xq.transpose(1, 2)        xk = xk.transpose(1, 2)        xv = xv.transpose(1, 2)        # 根据是否支持Flash Attention,选择实现方式。        if self.flash:            # 使用Flash Attention。            output = torch.nn.functional.scaled_dot_product_attention(xq, xk, xv, attn_mask=None, dropout_p=self.dropout if self.training else0.0, is_causal=True)        else:            # 使用手动实现的注意力机制。            scores = torch.matmul(xq, xk.transpose(2, 3)) / math.sqrt(self.head_dim)            assert hasattr(self, 'mask')            scores = scores + self.mask[:, :, :seqlen, :seqlen]            scores = F.softmax(scores.float(), dim=-1).type_as(xq)            scores = self.attn_dropout(scores)            output = torch.matmul(scores, xv)        # 恢复时间维度并合并头。        output = output.transpose(1, 2).contiguous().view(bsz, seqlen, -1)        # 最终投影回残差流。        output = self.wo(output)        output = self.resid_dropout(output)        return output

2.5 前馈层定义

前馈层(Feed-Forward Network,FFN)是Transformer中除了注意力层之外的另一个核心组成部分,负责对注意力层的输出进行非线性变换。在LLaMA架构中,前馈层采用了SwiGLU(Swish-Gated Linear Unit) 结构,这是一种门控线性单元(GLU)的变体,使用Swish激活函数(即SiLU)作为门控机制的激活函数。

与标准Transformer中简单的两层线性变换加ReLU激活不同,SwiGLU通过引入第三个线性变换作为门控信号,使得网络能够更精细地控制信息的流动,从而在相同的参数量下获得更强的表达能力。

前馈层的实现代码如下:

class MLP(nn.Module):    def __init__(self, dim: int, hidden_dim: int, multiple_of: int, dropout: float):        super().__init__()        # 如果没有指定隐藏层的维度,我们将其设置为输入维度的4倍        # 然后将其减少到2/3,最后确保它是multiple_of的倍数        if hidden_dim isNone:            hidden_dim = 4 * dim            hidden_dim = int(2 * hidden_dim / 3)            hidden_dim = multiple_of * ((hidden_dim + multiple_of - 1) // multiple_of)        # 定义第一层线性变换,从输入维度到隐藏维度        self.w1 = nn.Linear(dim, hidden_dim, bias=False)        # 定义第二层线性变换,从隐藏维度到输入维度        self.w2 = nn.Linear(hidden_dim, dim, bias=False)        # 定义第三层线性变换,从输入维度到隐藏维度        self.w3 = nn.Linear(dim, hidden_dim, bias=False)        # 定义dropout层,用于防止过拟合        self.dropout = nn.Dropout(dropout)    def forward(self, x):        # 前向传播函数        # 首先,输入x通过第一层线性变换和SILU激活函数        # 然后,结果乘以输入x通过第三层线性变换的结果        # 最后,通过第二层线性变换和dropout层        return self.dropout(self.w2(F.silu(self.w1(x)) * self.w3(x)))

仔细分析一下forward函数的实现逻辑:

  1. 首先,输入x通过第一层线性变换self.w1和SILU激活函数进行非线性映射
  2. 同时,输入x通过第三层线性变换self.w3得到门控信号
  3. 将上述两个结果进行逐元素相乘,实现门控机制
  4. 最后,通过第二层线性变换self.w2和dropout层进行输出映射和正则化

这种门控线性单元(GLU)的设计使得网络可以自适应地决定哪些特征被保留、哪些被抑制,从而让模型能够更灵活地控制信息流动。

2.6 Decoder定义

Decoder 层是构成LLM的基本单元,它将前面定义的自注意力层和前馈层有机地组合在一起,并辅以残差连接和层归一化,形成一个完整的Transformer解码器模块。每个Decoder层都包含自注意力、前馈网络、残差连接和层归一化等组件。

这里,Decoder层采用了Pre-Norm(前置归一化)的结构设计,即在每个子层(注意力或前馈)之前先进行归一化,然后再进入子层计算,最后将计算结果与输入进行残差连接。这种设计与原始Transformer中的Post-Norm(后置归一化)相比,在大规模训练中表现出了更好的稳定性和收敛速度。

Decoder层的实现代码如下:

class DecoderLayer(nn.Module):    def __init__(self, layer_id: int, args: ModelConfig):        super().__init__()        # 定义多头注意力的头数        self.n_heads = args.n_heads        # 定义输入维度        self.dim = args.dim        # 定义每个头的维度,等于输入维度除以头数        self.head_dim = args.dim // args.n_heads        # 定义LLaMA2Attention对象,用于进行多头注意力计算        self.attention = Attention(args)        # 定义LLaMAMLP对象,用于进行前馈神经网络计算        self.feed_forward = MLP(            dim=args.dim,            hidden_dim=args.hidden_dim,            multiple_of=args.multiple_of,            dropout=args.dropout,        )        # 定义层的ID        self.layer_id = layer_id        # 定义注意力计算的归一化层        self.attention_norm = RMSNorm(args.dim, eps=args.norm_eps)        # 定义前馈神经网络计算的归一化层        self.ffn_norm = RMSNorm(args.dim, eps=args.norm_eps)    def forward(self, x, freqs_cos, freqs_sin):        # 前向传播函数        # 首先,输入x经过注意力归一化层,然后进行注意力计算,结果与输入x相加得到h        # 然后,h经过前馈神经网络归一化层,然后进行前馈神经网络计算,结果与h相加得到输出        h = x + self.attention.forward(self.attention_norm(x), freqs_cos, freqs_sin)        out = h + self.feed_forward.forward(self.ffn_norm(h))        return out

2.7 架构组装

在完成了所有子模块(RMSNorm、注意力层、前馈层、Decoder层)的实现之后,最后一步就是将所有这些模块组装起来,形成一个完整的大模型架构。

组装过程包括:构建Token嵌入层(将输入的Token ID映射为稠密向量)、位置编码(使用旋转位置编码RoPE或可学习位置编码)、堆叠多个Decoder层、以及最终的输出投影层(将隐藏状态映射为词汇表大小的概率分布)。

模型架构组装代码如下:

class Transformer(PreTrainedModel):    config_class = ModelConfig  # 配置类    last_loss: Optional[torch.Tensor] # 记录最后一次计算的损失    def __init__(self, args: ModelConfig = None):        super().__init__(args)        # 初始化模型参数        self.args = args        # 词汇表大小        self.vocab_size = args.vocab_size        # 层数        self.n_layers = args.n_layers        # 词嵌入层        self.tok_embeddings = nn.Embedding(args.vocab_size, args.dim)        # Dropout层        self.dropout = nn.Dropout(args.dropout)        # Decoder层        self.layers = torch.nn.ModuleList()        for layer_id in range(args.n_layers):            self.layers.append(DecoderLayer(layer_id, args))        # 归一化层        self.norm = RMSNorm(args.dim, eps=args.norm_eps)        # 输出层        self.output = nn.Linear(args.dim, args.vocab_size, bias=False)        # 将词嵌入层的权重与输出层的权重共享        self.tok_embeddings.weight = self.output.weight         # 预计算相对位置嵌入的频率        freqs_cos, freqs_sin = precompute_freqs_cis(self.args.dim // self.args.n_heads, self.args.max_seq_len)        self.register_buffer("freqs_cos", freqs_cos, persistent=False)        self.register_buffer("freqs_sin", freqs_sin, persistent=False)        # 初始化所有权重        self.apply(self._init_weights)        # 对残差投影进行特殊的缩放初始化        for pn, p in self.named_parameters():            if pn.endswith('w3.weight') or pn.endswith('wo.weight'):                torch.nn.init.normal_(p, mean=0.0, std=0.02/math.sqrt(2 * args.n_layers))        # 初始化最后一次前向传播的损失属性        self.last_loss = None        self.OUT = CausalLMOutputWithPast()  # 输出容器        self._no_split_modules = [name for name, _ in self.named_modules()]  # 不分割的模块列表    def _init_weights(self, module):        # 初始化权重的函数        if isinstance(module, nn.Linear):            torch.nn.init.normal_(module.weight, mean=0.0, std=0.02)            if module.bias isnotNone:                torch.nn.init.zeros_(module.bias)        elif isinstance(module, nn.Embedding):            torch.nn.init.normal_(module.weight, mean=0.0, std=0.02)        def forward(self, tokens: torch.Tensor, targets: Optional[torch.Tensor] = None, **kwargs) -> torch.Tensor:        """        - tokens: Optional[torch.Tensor], 输入 token 张量。        - targets: Optional[torch.Tensor], 目标 token 张量。        - kv_cache: bool, 是否使用键值缓存。        - kwargs: 其他关键字参数。        - self.OUT: CausalLMOutputWithPast, 包含 logits 和损失。        """        if'input_ids'in kwargs:            tokens = kwargs['input_ids']        if'labels'in kwargs:            targets = kwargs['labels']        # 前向传播函数        _bsz, seqlen = tokens.shape        # 通过词嵌入层和Dropout层        h = self.tok_embeddings(tokens)        h = self.dropout(h)        # 获取相对位置嵌入的频率        freqs_cos = self.freqs_cos[:seqlen]        freqs_sin = self.freqs_sin[:seqlen]        # 通过Decoder层        for layer in self.layers:            h = layer(h, freqs_cos, freqs_sin)        # 通过归一化层        h = self.norm(h)        if targets isnotNone:            # 如果给定了目标,计算损失            logits = self.output(h)            self.last_loss = F.cross_entropy(logits.view(-1, logits.size(-1)), targets.view(-1), ignore_index=0, reduction='none')        else:            # 推理时的优化:只对最后一个位置的输出进行前向传播            logits = self.output(h[:, [-1], :])             self.last_loss = None        # 设置输出        self.OUT.__setitem__('logits', logits)        self.OUT.__setitem__('last_loss', self.last_loss)        return self.OUT        @torch.inference_mode()    def generate(self, idx, stop_id=None, max_new_tokens=256, temperature=1.0, top_k=None):        """        给定输入序列 idx(形状为 (bz,seq_len) 的长整型张量),通过多次生成新 token 来完成序列。        在 model.eval() 模式下运行。        """        index = idx.shape[1]        for _ in range(max_new_tokens):            # 如果序列上下文过长,截断它到最大长度            idx_cond = idx if idx.size(1) <= self.args.max_seq_len else idx[:, -self.args.max_seq_len:]                        # 前向传播获取序列中最后一个位置的 logits            logits = self(idx_cond).logits            logits = logits[:, -1, :] # 只保留最后一个时间步的输出                        if temperature == 0.0:                # 选择最有可能的索引                _, idx_next = torch.topk(logits, k=1, dim=-1)            else:                # 缩放 logits 并应用 softmax                logits = logits / temperature                if top_k isnotNone:                    v, _ = torch.topk(logits, min(top_k, logits.size(-1)))                    logits[logits < v[:, [-1]]] = -float('Inf')                probs = F.softmax(logits, dim=-1)                idx_next = torch.multinomial(probs, num_samples=1)                        if idx_next == stop_id:                break            # 将采样的索引添加到序列中并继续            idx = torch.cat((idx, idx_next), dim=1)        return idx[:, index:] # 只返回生成的token

2.8 完整代码

将上述所有模块整合在一起,形成完整的LLM实现:

# model_struct.pyimport mathimport inspectfrom dataclasses import dataclassfrom typing import Any, Optional, Tupleimport torchimport torch.nn.functional as Ffrom torch import nnfrom transformers import PreTrainedModel, AutoTokenizerfrom transformers.modeling_outputs import CausalLMOutputWithPastfrom transformers import PretrainedConfigclass ModelConfig(PretrainedConfig):    model_type = "My-LLM"# 模型ID    def __init__(            self,            dim: int = 768, # 模型维度            n_layers: int = 12, # Transformer的层数            n_heads: int = 16, # 注意力机制的头数            n_kv_heads: int = 8, # 键值头的数量            vocab_size: int = 6144, # 词汇表大小            hidden_dim: int = None, # 隐藏层维度            multiple_of: int = 64,             norm_eps: float = 1e-5, # 归一化层的eps            max_seq_len: int = 512, # 模型最大输入序列长度            dropout: float = 0.0, # dropout概率            flash_attn: bool = True, # 是否使用Flash Attention            **kwargs,    ):        self.dim = dim        self.n_layers = n_layers        self.n_heads = n_heads        self.n_kv_heads = n_kv_heads        self.vocab_size = vocab_size        self.hidden_dim = hidden_dim        self.multiple_of = multiple_of        self.norm_eps = norm_eps        self.max_seq_len = max_seq_len        self.dropout = dropout        self.flash_attn = flash_attn        super().__init__(**kwargs)class RMSNorm(nn.Module):    def __init__(self, dim: int, eps: float):        super().__init__()        # eps是为了防止除以0的情况        self.eps = eps        # weight是一个可学习的参数,全部初始化为1        self.weight = nn.Parameter(torch.ones(dim))    def _norm(self, x):        # 计算RMSNorm的核心部分        # x.pow(2).mean(-1, keepdim=True)计算了输入x的平方的均值        # torch.rsqrt是平方根的倒数,这样就得到了RMSNorm的分母部分,再加上eps防止分母为0        # 最后乘以x,得到RMSNorm的结果        return x * torch.rsqrt(x.pow(2).mean(-1, keepdim=True) + self.eps)    def forward(self, x):        # forward函数是模型的前向传播        # 首先将输入x转为float类型,然后进行RMSNorm,最后再转回原来的数据类型        # 最后乘以weight,这是RMSNorm的一个可学习的缩放因子        output = self._norm(x.float()).type_as(x)        return output * self.weight# 获得旋转嵌入的实部和虚部# 注意:此处的dim应为 dim//n_head,因为我们是对每个head进行旋转嵌入def precompute_freqs_cis(dim: int, end: int, theta: float = 10000.0):    # torch.arange(0, dim, 2)[: (dim // 2)].float()生成了一个从0开始,步长为2的序列,长度为dim的一半    # 然后每个元素除以dim,再取theta的倒数,得到频率    freqs = 1.0 / (theta ** (torch.arange(0, dim, 2)[: (dim // 2)].float() / dim))    # 生成一个从0到end的序列,长度为end    t = torch.arange(end, device=freqs.device)    # 计算外积,得到一个二维矩阵,每一行是t的元素乘以freqs的元素    freqs = torch.outer(t, freqs).float()    # 计算频率的余弦值,得到实部    freqs_cos = torch.cos(freqs)    # 计算频率的正弦值,得到虚部    freqs_sin = torch.sin(freqs)    return freqs_cos, freqs_sin# 此函数的作用是将freqs_cis调整为与x的形状相同,以便能够与x进行广播操作def reshape_for_broadcast(freqs_cis: torch.Tensor, x: torch.Tensor):    # 获取x的维度数    ndim = x.ndim    # 断言,确保1在x的维度范围内    assert0 <= 1 < ndim    # 断言,确保freqs_cis的形状与x的第二维和最后一维相同    assert freqs_cis.shape == (x.shape[1], x.shape[-1])    # 构造一个新的形状,除了第二维和最后一维,其他维度都为1,这样做是为了能够将freqs_cis与x进行广播操作    shape = [d if i == 1or i == ndim - 1else1for i, d in enumerate(x.shape)]    # 将freqs_cis调整为新的形状,并返回    return freqs_cis.view(shape)def apply_rotary_emb(    xq: torch.Tensor,    xk: torch.Tensor,    freqs_cos: torch.Tensor,    freqs_sin: torch.Tensor) -> Tuple[torch.Tensor, torch.Tensor]:    # 将查询和键张量转换为浮点数,并重塑形状以分离实部和虚部    xq_r, xq_i = xq.float().reshape(xq.shape[:-1] + (-1, 2)).unbind(-1)    xk_r, xk_i = xk.float().reshape(xk.shape[:-1] + (-1, 2)).unbind(-1)    # 重新塑形频率张量以进行广播    freqs_cos = reshape_for_broadcast(freqs_cos, xq_r)    freqs_sin = reshape_for_broadcast(freqs_sin, xq_r)    # 应用旋转,分别计算旋转后的实部和虚部    xq_out_r = xq_r * freqs_cos - xq_i * freqs_sin    xq_out_i = xq_r * freqs_sin + xq_i * freqs_cos    xk_out_r = xk_r * freqs_cos - xk_i * freqs_sin    xk_out_i = xk_r * freqs_sin + xk_i * freqs_cos    # 将最后两个维度合并,并还原为原始张量的形状    xq_out = torch.stack([xq_out_r, xq_out_i], dim=-1).flatten(3)    xk_out = torch.stack([xk_out_r, xk_out_i], dim=-1).flatten(3)    return xq_out.type_as(xq), xk_out.type_as(xk)def repeat_kv(x: torch.Tensor, n_rep: int) -> torch.Tensor:    # 获取输入张量的形状:批量大小、序列长度、键/值对头的数量、每个头的维度大小    bs, slen, n_kv_heads, head_dim = x.shape        # 如果重复次数为1,则不需要重复,直接返回原始张量    if n_rep == 1:        return x        # 对张量进行扩展和重塑操作以重复键值对    return (        x[:, :, :, None, :]  # 在第四个维度(头的维度前)添加一个新的维度        .expand(bs, slen, n_kv_heads, n_rep, head_dim)  # 将新添加的维度扩展到n_rep大小,实现重复的效果        .reshape(bs, slen, n_kv_heads * n_rep, head_dim)  # 重新塑形,合并键/值对头的数量和重复次数的维度    )class Attention(nn.Module):    def __init__(self, args: ModelConfig):        super().__init__()        # 根据是否指定n_kv_heads,确定用于键(key)和值(value)的头的数量。        self.n_kv_heads = args.n_heads if args.n_kv_heads isNoneelse args.n_kv_heads        # 确保总头数可以被键值头数整除。        assert args.n_heads % self.n_kv_heads == 0        # 模型并行处理大小,默认为1。        model_parallel_size = 1        # 本地计算头数,等于总头数除以模型并行处理大小。        self.n_local_heads = args.n_heads // model_parallel_size        # 本地键值头数,等于键值头数除以模型并行处理大小。        self.n_local_kv_heads = self.n_kv_heads // model_parallel_size        # 重复次数,用于扩展键和值的尺寸。        self.n_rep = self.n_local_heads // self.n_local_kv_heads        # 每个头的维度,等于模型维度除以头的总数。        self.head_dim = args.dim // args.n_heads        # 定义权重矩阵。        self.wq = nn.Linear(args.dim, args.n_heads * self.head_dim, bias=False)        self.wk = nn.Linear(args.dim, self.n_kv_heads * self.head_dim, bias=False)        self.wv = nn.Linear(args.dim, self.n_kv_heads * self.head_dim, bias=False)        # 输出权重矩阵。        self.wo = nn.Linear(args.n_heads * self.head_dim, args.dim, bias=False)        # 定义dropout。        self.attn_dropout = nn.Dropout(args.dropout)        self.resid_dropout = nn.Dropout(args.dropout)        # 保存dropout概率。        self.dropout = args.dropout        # 检查是否使用Flash Attention(需要PyTorch >= 2.0)。        self.flash = hasattr(torch.nn.functional, 'scaled_dot_product_attention')        ifnot self.flash:            # 若不支持Flash Attention,则使用手动实现的注意力机制,并设置mask。            print("WARNING: using slow attention. Flash Attention requires PyTorch >= 2.0")            # 创建一个上三角矩阵,用于遮蔽未来信息。            mask = torch.full((1, 1, args.max_seq_len, args.max_seq_len), float("-inf"))            mask = torch.triu(mask, diagonal=1)            # 注册为模型的缓冲区            self.register_buffer("mask", mask)    def forward(self, x: torch.Tensor, freqs_cos: torch.Tensor, freqs_sin: torch.Tensor, attention_mask: Optional[torch.Tensor] = None):        # 获取批次大小和序列长度,[batch_size, seq_len, dim]        bsz, seqlen, _ = x.shape        # 计算查询(Q)、键(K)、值(V)。        xq, xk, xv = self.wq(x), self.wk(x), self.wv(x)        # 调整形状以适应头的维度。        xq = xq.view(bsz, seqlen, self.n_local_heads, self.head_dim)        xk = xk.view(bsz, seqlen, self.n_local_kv_heads, self.head_dim)        xv = xv.view(bsz, seqlen, self.n_local_kv_heads, self.head_dim)        # 应用旋转位置嵌入(RoPE)。        xq, xk = apply_rotary_emb(xq, xk, freqs_cos, freqs_sin)        # 对键和值进行扩展以适应重复次数。        xk = repeat_kv(xk, self.n_rep)        xv = repeat_kv(xv, self.n_rep)        # 将头作为批次维度处理。        xq = xq.transpose(1, 2)        xk = xk.transpose(1, 2)        xv = xv.transpose(1, 2)        key_padding_mask = None        if attention_mask isnotNone:            key_padding_mask = attention_mask[:, None, None, :].to(dtype=torch.bool)        # 根据是否支持Flash Attention,选择实现方式。        if self.flash:            # 使用Flash Attention。            if key_padding_mask isnotNone:                causal_mask = torch.ones((seqlen, seqlen), dtype=torch.bool, device=x.device).tril()                full_attn_mask = causal_mask[None, None, :, :] & key_padding_mask                output = torch.nn.functional.scaled_dot_product_attention(                    xq,                    xk,                    xv,                    attn_mask=full_attn_mask,                    dropout_p=self.dropout if self.training else0.0,                    is_causal=False,                )            else:                output = torch.nn.functional.scaled_dot_product_attention(                    xq,                    xk,                    xv,                    attn_mask=None,                    dropout_p=self.dropout if self.training else0.0,                    is_causal=True,                )        else:            # 使用手动实现的注意力机制。            scores = torch.matmul(xq, xk.transpose(2, 3)) / math.sqrt(self.head_dim)            assert hasattr(self, 'mask')            scores = scores + self.mask[:, :, :seqlen, :seqlen]            if key_padding_mask isnotNone:                scores = scores.masked_fill(~key_padding_mask, float("-inf"))            scores = F.softmax(scores.float(), dim=-1).type_as(xq)            scores = self.attn_dropout(scores)            output = torch.matmul(scores, xv)        # 恢复时间维度并合并头。        output = output.transpose(1, 2).contiguous().view(bsz, seqlen, -1)        # 最终投影回残差流。        output = self.wo(output)        output = self.resid_dropout(output)        return outputclass MLP(nn.Module):    def __init__(self, dim: int, hidden_dim: int, multiple_of: int, dropout: float):        super().__init__()        # 如果没有指定隐藏层的维度,我们将其设置为输入维度的4倍        # 然后将其减少到2/3,最后确保它是multiple_of的倍数        if hidden_dim isNone:            hidden_dim = 4 * dim            hidden_dim = int(2 * hidden_dim / 3)            hidden_dim = multiple_of * ((hidden_dim + multiple_of - 1) // multiple_of)        # 定义第一层线性变换,从输入维度到隐藏维度        self.w1 = nn.Linear(dim, hidden_dim, bias=False)        # 定义第二层线性变换,从隐藏维度到输入维度        self.w2 = nn.Linear(hidden_dim, dim, bias=False)        # 定义第三层线性变换,从输入维度到隐藏维度        self.w3 = nn.Linear(dim, hidden_dim, bias=False)        # 定义dropout层,用于防止过拟合        self.dropout = nn.Dropout(dropout)    def forward(self, x):        # 前向传播函数        # 首先,输入x通过第一层线性变换和SILU激活函数        # 然后,结果乘以输入x通过第三层线性变换的结果        # 最后,通过第二层线性变换和dropout层        return self.dropout(self.w2(F.silu(self.w1(x)) * self.w3(x)))    class DecoderLayer(nn.Module):    def __init__(self, layer_id: int, args: ModelConfig):        super().__init__()        # 定义多头注意力的头数        self.n_heads = args.n_heads        # 定义输入维度        self.dim = args.dim        # 定义每个头的维度,等于输入维度除以头数        self.head_dim = args.dim // args.n_heads        # 定义LLaMA2Attention对象,用于进行多头注意力计算        self.attention = Attention(args)        # 定义LLaMAMLP对象,用于进行前馈神经网络计算        self.feed_forward = MLP(            dim=args.dim,            hidden_dim=args.hidden_dim,            multiple_of=args.multiple_of,            dropout=args.dropout,        )        # 定义层的ID        self.layer_id = layer_id        # 定义注意力计算的归一化层        self.attention_norm = RMSNorm(args.dim, eps=args.norm_eps)        # 定义前馈神经网络计算的归一化层        self.ffn_norm = RMSNorm(args.dim, eps=args.norm_eps)    def forward(self, x, freqs_cos, freqs_sin, attention_mask: Optional[torch.Tensor] = None):        # 前向传播函数        # 首先,输入x经过注意力归一化层,然后进行注意力计算,结果与输入x相加得到h        # 然后,h经过前馈神经网络归一化层,然后进行前馈神经网络计算,结果与h相加得到输出        h = x + self.attention.forward(self.attention_norm(x), freqs_cos, freqs_sin, attention_mask=attention_mask)        out = h + self.feed_forward.forward(self.ffn_norm(h))        return outclass Transformer(PreTrainedModel):    config_class = ModelConfig  # 配置类    last_loss: Optional[torch.Tensor] # 记录最后一次计算的损失    def __init__(self, args: ModelConfig = None):        super().__init__(args)        # 初始化模型参数        self.args = args        # 词汇表大小        self.vocab_size = args.vocab_size        # 层数        self.n_layers = args.n_layers        # 词嵌入层        self.tok_embeddings = nn.Embedding(args.vocab_size, args.dim)        # Dropout层        self.dropout = nn.Dropout(args.dropout)        # Decoder层        self.layers = torch.nn.ModuleList()        for layer_id in range(args.n_layers):            self.layers.append(DecoderLayer(layer_id, args))        # 归一化层        self.norm = RMSNorm(args.dim, eps=args.norm_eps)        # 输出层        self.output = nn.Linear(args.dim, args.vocab_size, bias=False)        # 将词嵌入层的权重与输出层的权重共享        self.tok_embeddings.weight = self.output.weight         # 预计算相对位置嵌入的频率        freqs_cos, freqs_sin = precompute_freqs_cis(self.args.dim // self.args.n_heads, self.args.max_seq_len)        self.register_buffer("freqs_cos", freqs_cos, persistent=False)        self.register_buffer("freqs_sin", freqs_sin, persistent=False)        # 初始化所有权重        self.apply(self._init_weights)        # 对残差投影进行特殊的缩放初始化        for pn, p in self.named_parameters():            if pn.endswith('w3.weight') or pn.endswith('wo.weight'):                torch.nn.init.normal_(p, mean=0.0, std=0.02/math.sqrt(2 * args.n_layers))        # 初始化最后一次前向传播的损失属性        self.last_loss = None        self.OUT = CausalLMOutputWithPast()  # 输出容器        self._no_split_modules = [name for name, _ in self.named_modules()]  # 不分割的模块列表    def _init_weights(self, module):        # 初始化权重的函数        if isinstance(module, nn.Linear):            torch.nn.init.normal_(module.weight, mean=0.0, std=0.02)            if module.bias isnotNone:                torch.nn.init.zeros_(module.bias)        elif isinstance(module, nn.Embedding):            torch.nn.init.normal_(module.weight, mean=0.0, std=0.02)    def _prepare_attention_mask(self, attention_mask: Optional[torch.Tensor], tokens: torch.Tensor) -> Optional[torch.Tensor]:        if attention_mask isNone:            returnNone        if attention_mask.dim() == 4:            attention_mask = attention_mask[:, 0, 0, :]        elif attention_mask.dim() == 3:            attention_mask = attention_mask[:, 0, :]        attention_mask = attention_mask.to(tokens.device)        if attention_mask.dtype != torch.bool:            attention_mask = attention_mask > 0        if attention_mask.shape != tokens.shape:            raise ValueError(f"attention_mask shape {attention_mask.shape} must match input_ids shape {tokens.shape}")        return attention_mask    def _left_pad_by_attention_mask(            self,            idx: torch.Tensor,            attention_mask: Optional[torch.Tensor],            pad_token_id: int    ) -> Tuple[torch.Tensor, Optional[torch.Tensor]]:        if attention_mask isNoneor attention_mask.all():            return idx, attention_mask        bsz = idx.size(0)        lengths = attention_mask.long().sum(dim=1)        max_len = max(int(lengths.max().item()), 1)        packed_idx = idx.new_full((bsz, max_len), pad_token_id)        packed_mask = attention_mask.new_zeros((bsz, max_len), dtype=torch.bool)        for row in range(bsz):            valid_len = int(lengths[row].item())            if valid_len <= 0:                continue            valid_tokens = idx[row][attention_mask[row]]            packed_idx[row, max_len - valid_len:] = valid_tokens            packed_mask[row, max_len - valid_len:] = True        return packed_idx, packed_mask        def forward(self, tokens: torch.Tensor, targets: Optional[torch.Tensor] = None, **kwargs) -> torch.Tensor:        """        - tokens: Optional[torch.Tensor], 输入 token 张量。        - targets: Optional[torch.Tensor], 目标 token 张量。        - kv_cache: bool, 是否使用键值缓存。        - kwargs: 其他关键字参数。        - self.OUT: CausalLMOutputWithPast, 包含 logits 和损失。        """        if'input_ids'in kwargs:            tokens = kwargs['input_ids']        if'labels'in kwargs:            targets = kwargs['labels']        attention_mask = self._prepare_attention_mask(kwargs.get('attention_mask'), tokens)        # 前向传播函数        _bsz, seqlen = tokens.shape        # 通过词嵌入层和Dropout层        h = self.tok_embeddings(tokens)        h = self.dropout(h)        # 获取相对位置嵌入的频率        freqs_cos = self.freqs_cos[:seqlen]        freqs_sin = self.freqs_sin[:seqlen]        # 通过Decoder层        for layer in self.layers:            h = layer(h, freqs_cos, freqs_sin, attention_mask=attention_mask)        # 通过归一化层        h = self.norm(h)        if targets isnotNone:            # 如果给定了目标,计算损失            logits = self.output(h)            ignore_index = self.args.pad_token_id if self.args.pad_token_id isnotNoneelse0            if torch.any(targets == -100):                ignore_index = -100            self.last_loss = F.cross_entropy(                logits.view(-1, logits.size(-1)),                targets.view(-1),                ignore_index=ignore_index,                reduction='none'            )        else:            # 推理时的小优化:只对最后一个位置的输出进行前向传播            if attention_mask isNone:                logits = self.output(h[:, [-1], :])            else:                full_logits = self.output(h)                last_token_pos = attention_mask.long().sum(dim=1).clamp(min=1) - 1                logits = full_logits[torch.arange(_bsz, device=tokens.device), last_token_pos].unsqueeze(1)            self.last_loss = None        # 设置输出        self.OUT.__setitem__('logits', logits)        self.OUT.__setitem__('last_loss', self.last_loss)        return self.OUT        @torch.inference_mode()    def generate(            self,            idx,            stop_id=None,            max_new_tokens=256,            temperature=1.0,            top_k=None,            attention_mask: Optional[torch.Tensor] = None,            pad_token_id: Optional[int] = None    ):        """        给定输入序列 idx(形状为 (bz,seq_len) 的长整型张量),通过多次生成新 token 来完成序列。        在 model.eval() 模式下运行。        """        if pad_token_id isNone:            pad_token_id = self.args.pad_token_id if self.args.pad_token_id isnotNoneelse0        attention_mask = self._prepare_attention_mask(attention_mask, idx)        idx, attention_mask = self._left_pad_by_attention_mask(idx, attention_mask, pad_token_id)        finished = torch.zeros(idx.size(0), dtype=torch.bool, device=idx.device)        index = idx.shape[1]        for _ in range(max_new_tokens):            # 如果序列上下文过长,截断它到最大长度            idx_cond = idx if idx.size(1) <= self.args.max_seq_len else idx[:, -self.args.max_seq_len:]            mask_cond = None            if attention_mask isnotNone:                mask_cond = attention_mask if attention_mask.size(1) <= self.args.max_seq_len else attention_mask[:, -self.args.max_seq_len:]                        # 前向传播获取序列中最后一个位置的 logits            logits = self(idx_cond, attention_mask=mask_cond).logits            logits = logits[:, -1, :] # 只保留最后一个时间步的输出                        if temperature == 0.0:                # 选择最有可能的索引                _, idx_next = torch.topk(logits, k=1, dim=-1)            else:                # 缩放 logits 并应用 softmax                logits = logits / temperature                if top_k isnotNone:                    v, _ = torch.topk(logits, min(top_k, logits.size(-1)))                    logits[logits < v[:, [-1]]] = -float('Inf')                probs = F.softmax(logits, dim=-1)                idx_next = torch.multinomial(probs, num_samples=1)            prev_finished = finished.clone()            if stop_id isnotNone:                if prev_finished.any():                    fill_token = pad_token_id if pad_token_id isnotNoneelse stop_id                    idx_next = torch.where(prev_finished[:, None], torch.full_like(idx_next, fill_token), idx_next)                finished = prev_finished | idx_next[:, 0].eq(stop_id)            # 将采样的索引添加到序列中并继续            idx = torch.cat((idx, idx_next), dim=1)            if attention_mask isnotNone:                next_mask = torch.ones((attention_mask.size(0), 1), dtype=attention_mask.dtype, device=attention_mask.device)                if prev_finished.any():                    next_mask[prev_finished] = False                attention_mask = torch.cat((attention_mask, next_mask), dim=1)            if stop_id isnotNoneand finished.all():                break        return idx[:, index:] # 只返回生成的token        def _greedy_decode(self, logits: torch.Tensor) -> torch.Tensor:        """        贪婪解码:选择概率最大的token        Args:            logits: 模型输出的logits,形状为 (batch_size, vocab_size)        Returns:            选择的token索引,形状为 (batch_size, 1)        """        _, idx_next = torch.topk(logits, k=1, dim=-1)        return idx_next    def _random_sample(self, logits: torch.Tensor, temperature: float = 1.0, top_k: int = None) -> torch.Tensor:        """        随机采样:基于概率分布随机选择token        Args:            logits: 模型输出的logits,形状为 (batch_size, vocab_size)            temperature: 温度参数,控制随机性            top_k: 只考虑概率最高的k个token        Returns:            选择的token索引,形状为 (batch_size, 1)        """        # 缩放 logits        logits = logits / temperature        # 应用top-k过滤        if top_k isnotNone:            v, _ = torch.topk(logits, min(top_k, logits.size(-1)))            # 将不在 top-k 内的 logits 设为负无穷            logits[logits < v[:, [-1]]] = -float('Inf')        # 计算概率并采样        probs = F.softmax(logits, dim=-1)        idx_next = torch.multinomial(probs, num_samples=1)        return idx_next    def _beam_search(self, idx: torch.Tensor, max_new_tokens: int, num_beams: int,                     temperature: float = 1.0, top_k: int = None, stop_id: int = None) -> torch.Tensor:        """        束搜索:维护多个候选序列,选择最优路径        束搜索的核心思想:在每一步生成时,不是只选择一个最佳token,        而是保留多个候选路径,最终选择累积概率最高的完整序列。        Args:            idx: 输入序列,形状为 (batch_size, seq_len)            max_new_tokens: 最大生成token数量            num_beams: 束宽度,表示保留的候选路径数量            temperature: 温度参数,控制分布的平滑程度            top_k: top-k过滤参数,限制候选token范围            stop_id: 停止生成的token ID,遇到则停止        Returns:            生成的token序列,形状为 (batch_size, generated_length)            只返回新生成的部分,不包含原始输入序列        """        # 获取输入序列的基本信息        batch_size = idx.shape[0]  # 批次大小,通常为1        seq_len = idx.shape[1]     # 输入序列长度        # 初始化束:创建 num_beams 个候选序列        beams = [idx.clone() for _ in range(num_beams)]        # 初始化每个候选序列的累积对数概率分数        beam_scores = torch.zeros(num_beams, device=idx.device)        # 第一个候选是原始输入序列,分数为0        beam_scores[0] = 0.0        # 其他候选初始分数设为负无穷,表示尚未生成        beam_scores[1:] = float('-inf')        # 主循环:逐步生成新的token,最多生成 max_new_tokens 个        for step in range(max_new_tokens):            # 每轮迭代收集新的候选序列和分数            new_beams = []   # 新的候选序列列表            new_scores = []  # 对应的分数列表            # 遍历当前的所有候选序列            for beam_idx, beam in enumerate(beams):                # 跳过无效候选(分数为负无穷的序列)                if beam_scores[beam_idx] == float('-inf'):                    continue                # 序列长度检查:如果超过最大长度,截取最后的部分                beam_cond = beam if beam.size(1) <= self.args.max_seq_len else beam[:, -self.args.max_seq_len:]                # 前向传播:获取模型对当前序列的预测                output = self(beam_cond)                # 提取最后一个位置的logits,用于预测下一个token                logits = output.logits[:, -1, :]  # 形状: (1, vocab_size)                # 温度缩放:调整logits的分布                if temperature != 1.0:                    logits = logits / temperature                    # 温度 > 1:分布更平滑,增加随机性                    # 温度 < 1:分布更尖锐,更确定                # Top-k过滤:限制候选token的范围,提高质量                if top_k isnotNone:                    # 找到logits中前top_k个最大的值                    v, _ = torch.topk(logits, min(top_k, logits.size(-1)))                    # 将不在前top_k内的logits设为负无穷                    logits[logits < v[:, [-1]]] = -float('Inf')                    # 这样采样时只会考虑前top_k个token                # 计算对数概率:使用log_softmax避免数值不稳定                log_probs = F.log_softmax(logits, dim=-1)                # 获取前 num_beams 个最可能的候选token                # 注意:这里的top-k与上面的top-k不同                # 上面的top-k是全局过滤,这里是束搜索的分支选择                top_log_probs, top_indices = torch.topk(log_probs, k=num_beams, dim=-1)                # 为当前候选序列生成 num_beams 个扩展序列                for k in range(num_beams):                    # 选择第k个候选token                    token = top_indices[:, k:k+1]      # token ID                    log_prob = top_log_probs[:, k]     # 对应的对数概率                    # 扩展序列:将新token添加到当前序列末尾                    new_beam = torch.cat([beam, token], dim=1)                    # 更新累积分数:原序列分数 + 新token的对数概率                    new_score = beam_scores[beam_idx] + log_prob.item()                    # 保存新的候选序列和分数                    new_beams.append(new_beam)                    new_scores.append(new_score)            # 安全检查:如果没有生成任何有效候选,提前结束            ifnot new_beams:                break            # 筛选最佳候选:从所有新生成的候选中选择分数最高的 num_beams 个            # 按分数降序排序,获取索引            sorted_indices = sorted(range(len(new_scores)), key=lambda i: new_scores[i], reverse=True)            # 选择前 num_beams 个最佳候选            beams = [new_beams[i] for i in sorted_indices[:num_beams]]            beam_scores = [new_scores[i] for i in sorted_indices[:num_beams]]            # 停止条件检查:检查最佳序列是否以停止token结尾            if stop_id isnotNoneand beams[0][0, -1] == stop_id:                break        # 返回得分最高的序列,只返回新生成的部分(去掉原始输入)        # beams[0] 是最终得分最高的完整序列        # [:, seq_len:] 切片只保留生成部分        return beams[0][:, seq_len:]    @torch.inference_mode()    def generate_super(self,                       idx,                       stop_id=None,                       max_new_tokens=256,                       temperature=1.0,                       top_k=None,                       do_sample=False,                       num_beams=1,                       attention_mask: Optional[torch.Tensor] = None,                       pad_token_id: Optional[int] = None                       ):        """        高级文本生成函数,支持三种解码策略:        1. 贪婪解码(Greedy Search):           - 参数:do_sample=False, num_beams=1           - 特点:每步选择概率最大的token,速度快、结果确定        2. 随机采样(Random Sampling):           - 参数:do_sample=True, num_beams=1           - 特点:基于概率分布随机采样,可配合temperature和top-k控制多样性        3. 束搜索(Beam Search):           - 参数:do_sample=False, num_beams>1           - 特点:维护多条候选路径,选择总概率最高的序列,质量更高但速度较慢        Args:            idx: 输入序列张量,形状为 (batch_size, seq_len)            stop_id: 停止生成的token ID            max_new_tokens: 最大生成token数量            temperature: 温度参数,控制随机性,越高越随机            top_k: 只考虑概率最高的k个token,None表示不考虑            do_sample: 是否使用随机采样,False时使用确定性解码            num_beams: 束搜索的束宽度,1表示不使用束搜索        Returns:            生成的token序列,形状为 (batch_size, generated_length)        """        # 参数验证        if temperature <= 0:            temperature = 0.001# 避免除零错误        if num_beams < 1:            num_beams = 1        if top_k isnotNoneand top_k < 1:            top_k = None        if pad_token_id isNone:            pad_token_id = self.args.pad_token_id if self.args.pad_token_id isnotNoneelse0        attention_mask = self._prepare_attention_mask(attention_mask, idx)        idx, attention_mask = self._left_pad_by_attention_mask(idx, attention_mask, pad_token_id)        # 束搜索逻辑        ifnot do_sample and num_beams > 1:            return self._beam_search(idx, max_new_tokens, num_beams, temperature, top_k, stop_id)        # 贪婪解码和随机采样逻辑        finished = torch.zeros(idx.size(0), dtype=torch.bool, device=idx.device)        index = idx.shape[1]        for _ in range(max_new_tokens):            # 如果序列上下文过长,截断它到最大长度            idx_cond = idx if idx.size(1) <= self.args.max_seq_len else idx[:, -self.args.max_seq_len:]            mask_cond = None            if attention_mask isnotNone:                mask_cond = attention_mask if attention_mask.size(1) <= self.args.max_seq_len else attention_mask[:, -self.args.max_seq_len:]            # 前向传播获取序列中最后一个位置的 logits            logits = self(idx_cond, attention_mask=mask_cond).logits            logits = logits[:, -1, :] # 只保留最后一个时间步的输出            # 根据参数选择解码策略            if do_sample:                idx_next = self._random_sample(logits, temperature, top_k)            else:                # 当temperature=0时使用贪婪解码                if temperature < 0.1:                    idx_next = self._greedy_decode(logits)                else:                    # 低温度下的随机采样(接近贪婪)                    idx_next = self._random_sample(logits, temperature, top_k)            prev_finished = finished.clone()            if stop_id isnotNone:                if prev_finished.any():                    fill_token = pad_token_id if pad_token_id isnotNoneelse stop_id                    idx_next = torch.where(prev_finished[:, None], torch.full_like(idx_next, fill_token), idx_next)                finished = prev_finished | idx_next[:, 0].eq(stop_id)            # 将选择的token添加到序列中            idx = torch.cat((idx, idx_next), dim=1)            if attention_mask isnotNone:                next_mask = torch.ones((attention_mask.size(0), 1), dtype=attention_mask.dtype, device=attention_mask.device)                if prev_finished.any():                    next_mask[prev_finished] = False                attention_mask = torch.cat((attention_mask, next_mask), dim=1)            if stop_id isnotNoneand finished.all():                break        return idx[:, index:]

三、预训练LLM

预训练(Pre-training)是LLM训练流程中的第一个阶段,也是最为关键的一步。在这个阶段,模型在大规模的通用文本语料上进行自监督学习,通过预测下一个Token的任务来学习语言的通用表示和世界知识。经过预训练后的模型已经具备了基本的语言理解和生成能力,为后续的微调阶段提供了良好的初始参数。

3.1 训练数据准备

对于预训练阶段,依然使用上一篇文章中的出门问问序列猴子数据集。该数据集中的通用文本子集由来自网页、百科、博客、问答、开源代码、书籍、报刊、专利、教材、考题等多种公开可获取的数据经过汇总和清洗后形成,覆盖面广、质量较高,非常适合用于大语言模型的预训练任务。整个数据集的总量大约在 10B Token 左右,对于一个小型模型的预训练来说已经足够。

原始数据采用键值对格式存储,样例如下:

几乎所有的LLM都有最大输入序列长度max_seq_len)的限制。在前面的配置中,这个值设置为512max_seq_len: int = 512),意味着模型在单次前向传播中最多只能处理512个Token。因此,在训练前需要将原始文本按照 512 的长度进行切分处理,确保每一条训练样本的长度不超过这一限制。

数据预处理代码如下:

import osimport jsonfrom tqdm import tqdmpretrain_data = '~/Desktop/train-llm/code/train_data/mobvoi_seq_monkey_general_open_corpus.jsonl'output_pretrain_data = '~/Desktop/train-llm/code/train_data/seq_monkey_pretrain.jsonl'def split_text(text, chunk_size=512):    """将文本按指定长度切分成块"""    return [text[i:i+chunk_size] for i in range(0, len(text), chunk_size)]with open(output_pretrain_data, 'a', encoding='utf-8') as pretrain:    with open(pretrain_data, 'r', encoding='utf-8') as f:        data = f.readlines()        for line in tqdm(data, desc=f"Processing lines in {pretrain_data}", leave=False):             line = json.loads(line)            text = line['text']            chunks = split_text(text)            for chunk in chunks:                pretrain.write(json.dumps({'text': chunk}, ensure_ascii=False) + '\n')

运行脚本如下,对文本进行处理:

切分后的数据样例如下,每个样本都包含了固定长度的Token序列,可以直接输入模型进行训练:

3.2 预训练实现

预训练的第一个必要组件是Tokenizer(分词器)。这里复用上一篇文章中训练好的Tokenizer,它采用了BPE(Byte Pair Encoding)算法,已经在所使用的语料上进行了适配。使用方式如下:

tokenizer = AutoTokenizer.from_pretrained('./tokenizer/')

完整的预训练代码如下,其中包含了模型初始化、数据集加载、训练循环、损失计算、梯度更新以及模型保存等全部流程:

# pretrain_llm.pyimport osimport reimport globimport argparseimport timeimport warningsimport mathimport jsonimport randomimport pickleimport numpy as npimport pandas as pdimport torchfrom torch import optimfrom contextlib import nullcontextfrom torch.utils.data import Dataset, DataLoaderfrom transformers import AutoTokenizerfrom model_struct import ModelConfig, Transformerimport swanlab# 忽略警告信息warnings.filterwarnings('ignore')def get_device():    """自动检测最优设备:CUDA > MPS > CPU"""    if torch.cuda.is_available():        return"cuda"    elif torch.backends.mps.is_available():        return"mps"    else:        return"cpu"def Logger(content):    print(content)def atomic_save(obj, path):    """原子保存:先写临时文件,再原子替换,防止中断导致文件损坏"""    tmp = path + '.tmp'    torch.save(obj, tmp)    os.replace(tmp, path)def build_checkpoint(model, optimizer, scaler, epoch, step, accumulation_step, run=None):    """    构建检查点字典,保存训练所需的全部状态    Args:        model: 模型(支持 DataParallel)        optimizer: 优化器        scaler: 混合精度梯度缩放器        epoch: 当前 epoch        step: 当前 step(数据加载步数)        accumulation_step: 当前梯度累积周期内已完成的步数        run: SwanLab run 对象(可选)    Returns:        dict: 完整的检查点字典    """    state_dict = model.module.state_dict() if isinstance(model, torch.nn.DataParallel) else model.state_dict()    return {        'model': state_dict,        'optimizer': optimizer.state_dict(),        'scaler': scaler.state_dict(),        'epoch': epoch,        'step': step,        'accumulation_step': accumulation_step,        'iter_per_epoch': iter_per_epoch,        'swanlab_run_id': run.id if run isnotNoneelseNone,        # 保存随机状态,确保恢复后训练可复现        'torch_rng_state': torch.random.get_rng_state(),        'numpy_rng_state': np.random.get_state(),        'python_rng_state': random.getstate(),        'cuda_rng_state': torch.cuda.get_rng_state() if torch.cuda.is_available() elseNone,        # 保存 DataLoader shuffle 生成器状态,确保恢复后数据顺序一致        'shuffle_generator_state': shuffle_generator.get_state() if shuffle_generator isnotNoneelseNone,    }class PretrainDataset(Dataset):    def __init__(self, data_path, tokenizer, max_length=512, max_samples=None):        super().__init__()        self.data_path = data_path        self.tokenizer = tokenizer        self.max_length = max_length        self.padding = tokenizer.pad_token_id if tokenizer.pad_token_id isnotNoneelse0        # 预计算每行的起始字节偏移量        self._offsets = []        with open(data_path, 'rb') as f:            self._offsets.append(0)            while f.readline():                self._offsets.append(f.tell())        self._total_lines = len(self._offsets) - 1# 最后一个 tell() 是 EOF        # 限制样本数量        if max_samples isnotNone:            self._total_lines = min(self._total_lines, max_samples)    def __len__(self):        return self._total_lines    def __getitem__(self, index: int):        # 尝试读取当前行,如果失败则跳到下一行        max_offset = len(self._offsets) - 1        for attempt in range(10):            try:                idx = min(index + attempt, max_offset - 1)                with open(self.data_path, 'rb') as f:                    f.seek(self._offsets[idx])                    line = f.readline().decode('utf-8').strip()                ifnot line:                    continue                sample = json.loads(line)                break            except (json.JSONDecodeError, IndexError):                continue        else:            raise RuntimeError(f"无法读取有效数据,index={index}")        text = f"{self.tokenizer.bos_token}{sample['text']}"        input_id = self.tokenizer(text).data['input_ids'][:self.max_length]        text_len = len(input_id)        # 没满最大长度的剩余部分        padding_len = self.max_length - text_len        input_id = input_id + [self.padding] * padding_len        # 0表示不计算损失        loss_mask = [1] * text_len + [0] * padding_len        input_id = np.array(input_id)        X = np.array(input_id[:-1]).astype(np.int64)        Y = np.array(input_id[1:]).astype(np.int64)        loss_mask = np.array(loss_mask[1:]).astype(np.int64)        return torch.from_numpy(X), torch.from_numpy(Y), torch.from_numpy(loss_mask)def get_lr(it, all):    """    计算当前迭代的学习率,使用余弦退火调度策略        学习率调度策略:    1. Warmup阶段:学习率从0线性增长到目标学习率    2. 余弦退火阶段:学习率按余弦函数衰减到最小学习率    3. 超出训练步数后:保持最小学习率        Args:        it (int): 当前迭代步数        all (int): 总迭代步数            Returns:        float: 当前步数对应的学习率    """    warmup_iters = args.warmup_iters  # 预热迭代次数    lr_decay_iters = all  # 学习率衰减的总迭代次数    min_lr = args.learning_rate / 10# 最小学习率,为初始学习率的1/10    # Warmup阶段:线性增长    if it < warmup_iters:        return args.learning_rate * it / warmup_iters        # 超出训练步数:保持最小学习率    if it > lr_decay_iters:        return min_lr        # 余弦退火阶段    decay_ratio = (it - warmup_iters) / (lr_decay_iters - warmup_iters)    assert0 <= decay_ratio <= 1    coeff = 0.5 * (1.0 + math.cos(math.pi * decay_ratio))  # 余弦系数    return min_lr + coeff * (args.learning_rate - min_lr)def train_epoch(epoch, resume_step=-1, resume_accumulation_step=0):    """    训练一个epoch的函数    实现了完整的训练循环,包括:    1. 数据加载和设备转移    2. 动态学习率调整    3. 前向传播和损失计算    4. 梯度累积和反向传播    5. 梯度裁剪和优化器更新    6. 日志记录和模型保存    Args:        epoch (int): 当前epoch编号        resume_step (int): 恢复训练时跳过的步数,默认-1表示不跳过        resume_accumulation_step (int): 恢复时梯度累积周期内已完成的步数,默认0    """    start_time = time.time()  # 记录开始时间    accumulation_step = resume_accumulation_step  # 恢复梯度累积进度    # 遍历数据加载器中的每个batch    for step, (X, Y, loss_mask) in enumerate(train_loader):        # 跳过已训练的步数,用于epoch内恢复        if step <= resume_step:            continue        # 将数据转移到指定设备(GPU/CPU)        X = X.to(args.device)  # 输入序列        Y = Y.to(args.device)  # 目标序列        loss_mask = loss_mask.to(args.device)  # 损失掩码,用于忽略padding token        # 计算当前步骤的学习率        lr = get_lr(epoch * iter_per_epoch + step, args.epochs * iter_per_epoch)        # 更新优化器中所有参数组的学习率        for param_group in optimizer.param_groups:            param_group['lr'] = lr        # 使用混合精度训练上下文        with ctx:            # 前向传播            out = model(X, Y)            # 计算损失并除以累积步数(用于梯度累积)            loss = out.last_loss / args.accumulation_steps            # 将loss_mask展平为一维            loss_mask = loss_mask.view(-1)            # 应用掩码计算有效损失(忽略padding位置)            loss = torch.sum(loss * loss_mask) / loss_mask.sum()        # 使用scaler进行混合精度的反向传播        scaler.scale(loss).backward()        accumulation_step += 1        # 每accumulation_steps步执行一次优化器更新        if accumulation_step == args.accumulation_steps:            # 取消梯度缩放,准备梯度裁剪            scaler.unscale_(optimizer)            # 梯度裁剪,防止梯度爆炸            torch.nn.utils.clip_grad_norm_(model.parameters(), args.grad_clip)            # 执行优化器步骤            scaler.step(optimizer)            # 更新scaler的缩放因子            scaler.update()            # 清零梯度,set_to_none=True可以节省内存            optimizer.zero_grad(set_to_none=True)            accumulation_step = 0# 重置梯度累积计数器        # 每log_interval步记录一次日志        if step % args.log_interval == 0:            spend_time = time.time() - start_time            # 打印训练进度信息            Logger(                'Epoch:[{}/{}]({}/{}) loss:{:.3f} lr:{:.7f} epoch_Time:{}min;'.format(                    epoch + 1,                    args.epochs,                    step,                    iter_per_epoch,                    loss.item() * args.accumulation_steps,  # 恢复真实的loss值                    optimizer.param_groups[-1]['lr'],                    spend_time / (step + 1) * iter_per_epoch // 60 - spend_time // 60))                        # 如果启用SwanLab,记录训练指标            if args.use_swanlab:                swanlab.log({                    "loss": loss.item() * args.accumulation_steps,                    "lr": optimizer.param_groups[-1]['lr']                })        # 每save_interval步保存一次模型        if (step + 1) % args.save_interval == 0:            checkpoint = build_checkpoint(model, optimizer, scaler, epoch, step,                                          accumulation_step, run if args.use_swanlab elseNone)            # 最新模型(覆盖)            ckp_latest = f'{args.save_dir}/pretrain_{lm_config.dim}_{lm_config.n_layers}_{lm_config.vocab_size}.pth'            # 带步数的检查点(不覆盖)            ckp_step = f'{args.save_dir}/pretrain_{lm_config.dim}_{lm_config.n_layers}_{lm_config.vocab_size}_step{step+1}.pth'            atomic_save(checkpoint, ckp_latest)            atomic_save(checkpoint, ckp_step)            # Logger(f'检查点已保存: {ckp_step}')            # 清理旧的step检查点,只保留最近5个(按步数数字排序)            step_ckps = glob.glob(f'{args.save_dir}/pretrain_{lm_config.dim}_{lm_config.n_layers}_{lm_config.vocab_size}_step*.pth')            step_ckps.sort(key=lambda x: int(re.search(r'_step(\d+)\.pth$', x).group(1)))            for old_ckp in step_ckps[:-5]:                os.remove(old_ckp)    # 训练结束,保存最终模型和epoch检查点    checkpoint = build_checkpoint(model, optimizer, scaler, epoch, step,                                  accumulation_step, run if args.use_swanlab elseNone)    ckp_latest = f'{args.save_dir}/pretrain_{lm_config.dim}_{lm_config.n_layers}_{lm_config.vocab_size}.pth'    ckp_epoch = f'{args.save_dir}/pretrain_{lm_config.dim}_{lm_config.n_layers}_{lm_config.vocab_size}_epoch{epoch+1}.pth'    atomic_save(checkpoint, ckp_latest)    atomic_save(checkpoint, ckp_epoch)    Logger(f'训练完成,模型已保存: {ckp_epoch}')def init_model():    """    初始化模型和分词器        功能包括:    1. 加载预训练的分词器    2. 创建Transformer模型    3. 设置多GPU并行训练(如果可用)    4. 将模型移动到指定设备    5. 统计并打印模型参数量        Returns:        tuple: (model, tokenizer) 初始化后的模型和分词器    """    def count_parameters(model):        """        统计模型中可训练参数的数量                Args:            model: PyTorch模型                    Returns:            int: 可训练参数总数        """        return sum(p.numel() for p in model.parameters() if p.requires_grad)    # 从本地路径加载预训练的分词器    tokenizer = AutoTokenizer.from_pretrained('./tokenizer/')    if tokenizer.pad_token_id isnotNone:        lm_config.pad_token_id = tokenizer.pad_token_id    # 根据配置创建Transformer模型    model = Transformer(lm_config)        # 多卡初始化(使用 DataParallel 简化演示)# ⚠️ 重要说明:DataParallel 仅适用于小模型演示场景。对于真实大模型训练(7B+),请务必使用:#   - DistributedDataParallel (DDP):更适合多机多卡,性能优于DP#   - FSDP (Fully Sharded Data Parallel):支持模型分片,可训练超大模型# 参考实现:#   - DDP: https://pytorch.org/tutorials/intermediate/ddp_tutorial.html#   - FSDP: https://pytorch.org/docs/stable/fsdp.html    if"cuda"in args.device:        num_gpus = torch.cuda.device_count()        if num_gpus > 1:            Logger(f"Using {num_gpus} GPUs with DataParallel!")            Logger(f"⚠️ For large models (7B+), consider using DDP or FSDP instead.")            model = torch.nn.DataParallel(model)    elif"mps"in args.device:        Logger("Using Apple Silicon MPS")        # 将模型移动到指定设备(GPU或CPU)    model = model.to(args.device)        # 计算并打印模型参数量(以百万为单位)    Logger(f'LLM总参数量:{count_parameters(model) / 1e6:.3f} 百万')    return model, tokenizerif __name__ == "__main__":    # ==================== 命令行参数解析 ====================    parser = argparse.ArgumentParser(description="预训练")        # 基础训练参数    parser.add_argument("--out_dir", type=str, default="base_model", help="模型输出目录")    parser.add_argument("--epochs", type=int, default=1, help="训练轮数")    parser.add_argument("--batch_size", type=int, default=48, help="批次大小")    parser.add_argument("--learning_rate", type=float, default=2e-4, help="学习率")    parser.add_argument("--device", type=str, default=None, help="训练设备 (自动检测: cuda/mps/cpu)")    parser.add_argument("--dtype", type=str, default="float32", help="数据类型 (MPS建议用float32)")        # 实验跟踪和数据加载参数    parser.add_argument("--use_swanlab", action="store_true", help="是否使用SwanLab进行实验跟踪")    parser.add_argument("--num_workers", type=int, default=8, help="数据加载的工作进程数")    parser.add_argument("--data_path", type=str, help="训练数据路径")    parser.add_argument("--max_samples", type=int, default=None, help="最大训练样本数,用于调试和节省内存")        # 训练优化参数    parser.add_argument("--accumulation_steps", type=int, default=8, help="梯度累积步数")    parser.add_argument("--grad_clip", type=float, default=1.0, help="梯度裁剪阈值")    parser.add_argument("--warmup_iters", type=int, default=0, help="学习率预热迭代次数")        # 日志和保存参数    parser.add_argument("--log_interval", type=int, default=100, help="日志记录间隔")    parser.add_argument("--save_interval", type=int, default=100, help="模型保存检查点间隔")        # 多GPU训练参数 (仅CUDA有效)    parser.add_argument("--gpus", type=str, default=None, help="CUDA GPU ID,用逗号分隔 (例如: '0,1,2')")    args = parser.parse_args()    # ==================== 设备环境设置 ====================    # 设置可见的GPU设备    if args.gpus isnotNoneand torch.cuda.is_available():        os.environ["CUDA_VISIBLE_DEVICES"] = args.gpus        args.device = "cuda:0"    elif args.device isNone:        args.device = get_device()    # ==================== 模型配置 ====================    # 定义语言模型的配置参数    lm_config = ModelConfig(        dim=768,      # 模型维度        n_layers=12,   # Transformer层数    )    # ==================== 训练环境设置 ====================    max_seq_len = lm_config.max_seq_len  # 最大序列长度    args.save_dir = os.path.join(args.out_dir)  # 模型保存目录    # 创建必要的目录    os.makedirs(args.out_dir, exist_ok=True)    # ==================== 检查点检测 ====================    # 检测是否存在检查点,用于恢复训练和SwanLab实验    ckp_path = f'{args.save_dir}/pretrain_{lm_config.dim}_{lm_config.n_layers}_{lm_config.vocab_size}.pth'    checkpoint = None    resume_run_id = None    saved_shuffle_generator_state = None    if os.path.exists(ckp_path):        try:            checkpoint = torch.load(ckp_path, map_location='cpu')            resume_run_id = checkpoint.get('swanlab_run_id', None)            saved_shuffle_generator_state = checkpoint.get('shuffle_generator_state', None)            Logger(f'发现检查点,将从 epoch {checkpoint["epoch"]}, step {checkpoint["step"]} 恢复训练')        except (RuntimeError, EOFError, pickle.UnpicklingError):            # latest 检查点损坏,尝试从最新的 step 检查点恢复            step_ckps = sorted(glob.glob(f'{args.save_dir}/pretrain_{lm_config.dim}_{lm_config.n_layers}_{lm_config.vocab_size}_step*.pth'))            if step_ckps:                checkpoint = torch.load(step_ckps[-1], map_location='cpu')                resume_run_id = checkpoint.get('swanlab_run_id', None)                saved_shuffle_generator_state = checkpoint.get('shuffle_generator_state', None)                Logger(f'最新检查点损坏,从历史检查点 {step_ckps[-1]} 恢复')            else:                Logger(f'最新检查点损坏且无历史检查点,从头开始训练')    # ==================== 实验跟踪初始化 ====================    run = None    if args.use_swanlab:        # 注意:使用前需要先登录 swanlab.login(api_key='your key')        run = swanlab.init(            project="My-LLM",  # 项目名称            experiment_name="Pretrain",  # 实验名称            config=vars(args),  # 保存所有超参数            id=resume_run_id,  # 恢复上次的run(如有)            resume="allow",    # 有则恢复,无则新建        )    # 设置随机种子以确保结果可复现    torch.manual_seed(42)        # 确定设备类型(用于选择合适的上下文管理器)    device_type = "cuda"if"cuda"in args.device else ("mps"if"mps"in args.device else"cpu")    # 设置混合精度训练的上下文管理器    # MPS/CPU训练时使用nullcontext,CUDA训练时使用autocast    if device_type == "cuda":        ptdtype = {'float32': torch.float32, 'bfloat16': torch.bfloat16, 'float16': torch.float16}[args.dtype]        ctx = torch.cuda.amp.autocast(dtype=ptdtype)    else:        ctx = nullcontext()  # MPS/CPU 使用 float32    # ==================== 模型和数据初始化 ====================    # 初始化模型和分词器    model, tokenizer = init_model()        # 创建训练数据集    train_ds = PretrainDataset(args.data_path, tokenizer, max_length=max_seq_len, max_samples=args.max_samples)    # 创建显式 shuffle 生成器,确保恢复训练后数据顺序与中断前一致    shuffle_generator = torch.Generator()    shuffle_generator.manual_seed(42)  # 固定种子,保证不同运行间初始顺序一致    # 如果从检查点恢复,还原 shuffle 生成器状态    if saved_shuffle_generator_state isnotNone:        shuffle_generator.set_state(saved_shuffle_generator_state)        Logger('已恢复 DataLoader shuffle 生成器状态')    # 创建数据加载器    # MPS/CPU 不支持 pin_memory,num_workers 设为 0 避免多进程问题    is_cuda = "cuda"in args.device    train_loader = DataLoader(        train_ds,        batch_size=args.batch_size,  # 批次大小        pin_memory=is_cuda,          # 仅CUDA使用pin_memory加速        drop_last=False,             # 不丢弃最后一个不完整的批次        shuffle=True,                # 随机打乱数据        num_workers=args.num_workers if is_cuda else0,  # MPS/CPU用0避免多进程问题        generator=shuffle_generator,  # 使用显式生成器,确保 shuffle 顺序可恢复    )    # ==================== 优化器和训练组件初始化 ====================    # 初始化混合精度训练的梯度缩放器    # 只有CUDA + float16/bfloat16时才启用    scaler = torch.amp.GradScaler("cuda", enabled=(device_type == "cuda"and args.dtype in ['float16', 'bfloat16']))        # 初始化Adam优化器    optimizer = optim.Adam(model.parameters(), lr=args.learning_rate)    # ==================== 加载检查点(如有) ====================    start_epoch = 0    resume_step = -1# 默认不跳过任何步数    resume_accumulation_step = 0# 默认梯度累积从0开始    saved_iter_per_epoch = None    if checkpoint isnotNone:        # 处理 DataParallel 兼容性:多卡保存 → 单卡加载 或 单卡保存 → 多卡加载        model_state_dict = checkpoint['model']        is_dataparallel = isinstance(model, torch.nn.DataParallel)        state_dict_is_dataparallel = any(k.startswith('module.') for k in model_state_dict.keys())        if is_dataparallel andnot state_dict_is_dataparallel:            # 单卡保存 → 多卡加载:添加 module. 前缀            model_state_dict = {f'module.{k}': v for k, v in model_state_dict.items()}            Logger('兼容性处理:单卡检查点 → 多卡模型')        elifnot is_dataparallel and state_dict_is_dataparallel:            # 多卡保存 → 单卡加载:移除 module. 前缀            model_state_dict = {k.replace('module.', '', 1): v for k, v in model_state_dict.items()}            Logger('兼容性处理:多卡检查点 → 单卡模型')        model.load_state_dict(model_state_dict)        optimizer.load_state_dict(checkpoint['optimizer'])        scaler.load_state_dict(checkpoint['scaler'])        start_epoch = checkpoint['epoch']        resume_step = checkpoint['step']  # 恢复到中断时的步数,跳过已训练部分        resume_accumulation_step = checkpoint.get('accumulation_step', 0)  # 恢复梯度累积进度        # 恢复随机状态,确保训练可复现        if'torch_rng_state'in checkpoint:            torch.random.set_rng_state(checkpoint['torch_rng_state'])            np.random.set_state(checkpoint['numpy_rng_state'])            random.setstate(checkpoint['python_rng_state'])            if checkpoint.get('cuda_rng_state') isnotNoneand torch.cuda.is_available():                torch.cuda.set_rng_state(checkpoint['cuda_rng_state'])            Logger('已恢复随机状态')        saved_iter_per_epoch = checkpoint.get('iter_per_epoch', None)        Logger(f'已恢复模型、优化器和scaler状态,从 epoch {start_epoch}, step {resume_step + 1} 继续训练'               f'(梯度累积进度: {resume_accumulation_step}/{args.accumulation_steps})')        del checkpoint  # 释放内存    # ==================== 开始训练 ====================    # 计算每个epoch的迭代次数    iter_per_epoch = len(train_loader)    # 校验 iter_per_epoch 是否与检查点一致(学习率调度依赖此值)    if saved_iter_per_epoch isnotNoneand saved_iter_per_epoch != iter_per_epoch:        Logger(f'⚠️ 警告:检查点记录 iter_per_epoch={saved_iter_per_epoch},当前为 {iter_per_epoch},'               f'学习率调度可能不连续(数据集大小发生了变化)')    # 开始训练循环(从恢复的epoch开始)    for epoch in range(start_epoch, args.epochs):        train_epoch(epoch,                    resume_step if epoch == start_epoch else-1,                    resume_accumulation_step if epoch == start_epoch else0)

这里使用了 PyTorch 的 DataParallel(DP)进行多卡训练。需要明确的是:

  1. 适用范围:DP 适合模型参数量较小(<1B)、单机多卡的场景。这里的模型约83M参数,DP 足以胜任。
  2. 生产环境选择:在大规模训练(7B、13B、70B)中,推荐使用:
  • DistributedDataParallel (DDP):相比 DP,DDP 在多机多卡场景下性能更优,且支持梯度分桶优化,通信开销更小。
  • Fully Sharded Data Parallel (FSDP):PyTorch 原生支持,可将模型参数、梯度和优化器状态分片到多个 GPU 上,显著降低显存占用,适合训练百亿级以上模型。
  • DeepSpeed:微软开源的高效训练框架,支持 ZeRO 优化,与 DDP/FSDP 互补。
  1. 快速迁移路径:如需从 DP 迁移到 DDP,只需将 DataParallel 替换为 DistributedDataParallel,并添加相应的分布式启动脚本(如 torchruntorch.distributed.launch)。

3.3 执行训练

先来看一下训练脚本的帮助说明,了解它支持哪些命令行参数:

从帮助信息中可以看到,脚本支持配置输出目录、是否使用SwanLab可视化监控、数据路径、最大样本数等多个选项,方便进行不同的配置。

执行如下训练命令:

python pretrain_llm.py --out_dir base_model --use_swanlab --data_path ~/Desktop/train-llm/code/train_data/seq_monkey_pretrain.jsonl --max_samples 30000

训练启动后的控制台输出如下,可以看到训练的轮次(Epoch)、步数(Step)、损失值(Loss)以及学习率(LR)等关键信息在实时更新:

系统资源消耗情况:

通过SwanLab监控的训练指标:

训练完成后的结果类似如下所示,模型检查点文件被保存在了指定的输出目录中:

四、SFT监督微调

有监督微调(Supervised Fine-Tuning,SFT)是LLM训练流程中的第二个阶段。在预训练的基础上,SFT使用高质量的标注数据集(通常是对话形式或指令-回答对)对模型进行进一步训练,使模型学会遵循指令以对话方式与用户交互。SFT是将一个"基础模型"转变为"可用助手"的关键步骤,也是LLM技术栈中最核心的环节之一。

SFT训练代码和预训练基本类似,主要的区别在于使用的数据集不同以及损失计算方式的不同(SFT通常只计算回答部分的损失,而忽略用户输入部分的损失)。这里使用的是专门为对话场景设计的SFTDataset类来处理多轮对话数据。

4.1 SFT训练数据

SFT阶段使用的是 BelleGroup 数据集,该数据集包含了约350万条中文对话数据,涵盖了人机对话、人人对话、人物对话等多种对话场景,数据质量较高、覆盖面广,非常适合用于训练对话生成模型。该数据集的地址如下:

https://huggingface.co/datasets/BelleGroup/train_3.5M_CN

使用如下命令下载数据集:

huggingface-cli download --repo-type dataset --resume-download BelleGroup/train_3.5M_CN --local-dir BelleGroup

原始数据结构如下,每条数据是一个JSON对象,包含了对话的多个轮次:

原始数据格式并不能直接用于SFT训练,需要进行格式转换,将其整理成标准的多轮对话格式。

数据转换代码如下:

# sft_data_deal.pyimport osimport jsonfrom tqdm import tqdm# 原始文件路径sft_data = '~/Desktop/train-llm/code/sft_data/train_3.5M_CN.json'# 转换后文件路径output_sft_data = '~/Desktop/train-llm/code/sft_data/BelleGroup_sft.jsonl'# 处理SFT数据def convert_message(data):    """    将原始数据转换为标准格式    """    message = [        {"role": "system", "content": "你是一个AI助手"},    ]    for item in data:        if item['from'] == 'human':            message.append({'role': 'user', 'content': item['value']})        elif item['from'] == 'assistant':            message.append({'role': 'assistant', 'content': item['value']})    return messagewith open(output_sft_data, 'a', encoding='utf-8') as sft:    with open(sft_data, 'r', encoding='utf-8') as f:        data = f.readlines()        for item in tqdm(data, desc="Processing", unit="lines"):            item = json.loads(item)            message = convert_message(item['conversations'])            sft.write(json.dumps(message, ensure_ascii=False) + '\n')

执行数据转换代码:

转换后的数据格式如下,数据已被标准化为统一的对话格式,每条样本都包含了清晰的角色标记和对话内容,可以直接用于SFT训练:

4.2 SFT训练实现

SFT训练同样复用上一篇文章中训练的Tokenizer,保持词表的一致性,使用方式如下:

tokenizer = AutoTokenizer.from_pretrained('./tokenizer/')

完整的SFT训练代码如下,与预训练代码的主要区别在于使用了SFTDataset替代预训练数据集,并且在计算损失时通过ignore_index将用户输入部分的Token排除在损失计算之外,确保模型只学习生成回答:

# sft_train_llm.pyimport osimport reimport globimport jsonimport argparseimport timeimport warningsimport mathimport randomimport pickleimport numpy as npimport torchfrom torch import optimfrom contextlib import nullcontextfrom torch.utils.data import Dataset, DataLoaderfrom transformers import AutoTokenizerfrom model_struct import ModelConfig, Transformerimport swanlab# 忽略警告warnings.filterwarnings('ignore')def get_device():    """自动检测最优设备:CUDA > MPS > CPU"""    if torch.cuda.is_available():        return"cuda"    elif torch.backends.mps.is_available():        return"mps"    else:        return"cpu"def Logger(content):    """日志记录器"""    print(content)def atomic_save(obj, path):    """原子保存:先写临时文件,再原子替换,防止中断导致文件损坏"""    tmp = path + '.tmp'    torch.save(obj, tmp)    os.replace(tmp, path)def build_checkpoint(model, optimizer, scaler, epoch, step, accumulation_step, iter_per_epoch, run=None, shuffle_generator=None):    """构建检查点字典,保存训练所需的全部状态"""    state_dict = model.module.state_dict() if isinstance(model, torch.nn.DataParallel) else model.state_dict()    return {        'model': state_dict,        'optimizer': optimizer.state_dict(),        'scaler': scaler.state_dict(),        'epoch': epoch,        'step': step,        'accumulation_step': accumulation_step,        'iter_per_epoch': iter_per_epoch,        'swanlab_run_id': run.id if run isnotNoneelseNone,        'torch_rng_state': torch.random.get_rng_state(),        'numpy_rng_state': np.random.get_state(),        'python_rng_state': random.getstate(),        'cuda_rng_state': torch.cuda.get_rng_state() if torch.cuda.is_available() elseNone,        'shuffle_generator_state': shuffle_generator.get_state() if shuffle_generator isnotNoneelseNone,    }class SFTDataset(Dataset):    def __init__(self, data_path, tokenizer, max_length=512, max_samples=None):        super().__init__()        self.data_path = data_path        self.tokenizer = tokenizer        self.max_length = max_length        self.padding = tokenizer.pad_token_id if tokenizer.pad_token_id isnotNoneelse0        self._offsets = []        with open(data_path, 'rb') as f:            self._offsets.append(0)            while f.readline():                self._offsets.append(f.tell())        self._total_lines = len(self._offsets) - 1        # 限制样本数量        if max_samples isnotNone:            self._total_lines = min(self._total_lines, max_samples)    def __len__(self):        return self._total_lines    def generate_loss_mask(self, input_ids):        # 生成 loss mask, 0 表示不计算损失, 1 表示计算损失        mask = [0] * len(input_ids)        a_sequence = self.tokenizer("<|im_start|>assistant\n")['input_ids']  # <|im_start|>assistant\n        a_length = len(a_sequence)        n = len(input_ids)        i = 0                while i <= n - a_length:            # 检查当前位置是否匹配目标子序列            match = True            for k in range(a_length):                if input_ids[i + k] != a_sequence[k]:                    match = False                    break            if match:                # 从子序列结束的位置开始查找第一个 eos_token                j = None                for idx in range(i + a_length, n):                    if input_ids[idx] == self.tokenizer.eos_token_id:                        j = idx                        break                if j isnotNone:                    start = i + a_length                    end = j  # 结束位置设为j(包含4)                    # 标记区间为1(包括start到end)                    if start <= end:                        for pos in range(start, end + 1):                            if pos < len(mask):                                mask[pos] = 1                # 跳过当前子序列,避免重叠匹配                i += a_length            else:                i += 1        return mask    def __getitem__(self, index: int):        # 尝试读取当前行,如果失败则跳到下一行        max_offset = len(self._offsets) - 1        for attempt in range(10):            try:                idx = min(index + attempt, max_offset - 1)                with open(self.data_path, 'rb') as f:                    f.seek(self._offsets[idx])                    line = f.readline().decode('utf-8').strip()                ifnot line:                    continue                sample = json.loads(line)                break            except (json.JSONDecodeError, IndexError):                continue        else:            raise RuntimeError(f"无法读取有效数据,index={index}")        text = self.tokenizer.apply_chat_template(sample, tokenize=False, add_generation_prompt=False)        input_id = self.tokenizer(text).data['input_ids'][:self.max_length]        text_len = len(input_id)        # 没满最大长度的剩余部分        padding_len = self.max_length - text_len        input_id = input_id + [self.padding] * padding_len        # 0表示不计算损失        loss_mask = self.generate_loss_mask(input_id)        input_id = np.array(input_id)        X = np.array(input_id[:-1]).astype(np.int64)        Y = np.array(input_id[1:]).astype(np.int64)        loss_mask = np.array(loss_mask[1:]).astype(np.int64)        return torch.from_numpy(X), torch.from_numpy(Y), torch.from_numpy(loss_mask)def get_lr(it, all):    """获取学习率(余弦退火调度)"""    warmup_iters = args.warmup_iters    lr_decay_iters = all    min_lr = args.learning_rate / 10    if it < warmup_iters:        return args.learning_rate * it / warmup_iters    if it > lr_decay_iters:        return min_lr    decay_ratio = (it - warmup_iters) / (lr_decay_iters - warmup_iters)    assert0 <= decay_ratio <= 1    coeff = 0.5 * (1.0 + math.cos(math.pi * decay_ratio))    return min_lr + coeff * (args.learning_rate - min_lr)def train_epoch(epoch, resume_step=-1, resume_accumulation_step=0):    """训练一个epoch"""    start_time = time.time()    accumulation_step = resume_accumulation_step    for step, (X, Y, loss_mask) in enumerate(train_loader):        # 跳过已训练的步数        if step <= resume_step:            continue        X = X.to(args.device)        Y = Y.to(args.device)        loss_mask = loss_mask.to(args.device)        # 获取学习率并更新优化器        lr = get_lr(epoch * iter_per_epoch + step, args.epochs * iter_per_epoch)        for param_group in optimizer.param_groups:            param_group['lr'] = lr        # 前向传播        with ctx:            out = model(X, Y)            loss = out.last_loss / args.accumulation_steps            loss_mask = loss_mask.view(-1)            loss = torch.sum(loss * loss_mask) / loss_mask.sum()        # 反向传播        scaler.scale(loss).backward()        accumulation_step += 1        # 更新权重        if accumulation_step == args.accumulation_steps:            scaler.unscale_(optimizer)            torch.nn.utils.clip_grad_norm_(model.parameters(), args.grad_clip)            scaler.step(optimizer)            scaler.update()            optimizer.zero_grad(set_to_none=True)            accumulation_step = 0        # 打印日志        if step % args.log_interval == 0:            spend_time = time.time() - start_time            Logger(                'Epoch:[{}/{}]({}/{}) loss:{:.3f} lr:{:.7f} epoch_Time:{}min;'.format(                    epoch + 1,                    args.epochs,                    step,                    iter_per_epoch,                    loss.item() * args.accumulation_steps,                    optimizer.param_groups[-1]['lr'],                    spend_time / (step + 1) * iter_per_epoch // 60 - spend_time // 60))            if args.use_swanlab:                swanlab.log({                    "loss": loss.item() * args.accumulation_steps,                    "lr": optimizer.param_groups[-1]['lr']                })        # 每save_interval步保存一次模型        if (step + 1) % args.save_interval == 0:            checkpoint = build_checkpoint(model, optimizer, scaler, epoch, step,                                          accumulation_step, iter_per_epoch, run if args.use_swanlab elseNone, shuffle_generator)            # 最新模型(覆盖)            ckp_latest = f'{args.save_dir}/sft_{lm_config.dim}_{lm_config.n_layers}_{lm_config.vocab_size}.pth'            # 带步数的检查点(不覆盖)            ckp_step = f'{args.save_dir}/sft_{lm_config.dim}_{lm_config.n_layers}_{lm_config.vocab_size}_step{step+1}.pth'            atomic_save(checkpoint, ckp_latest)            atomic_save(checkpoint, ckp_step)            # Logger(f'检查点已保存: {ckp_step}')            # 清理旧的step检查点,只保留最近5个(按步数数字排序)            step_ckps = glob.glob(f'{args.save_dir}/sft_{lm_config.dim}_{lm_config.n_layers}_{lm_config.vocab_size}_step*.pth')            step_ckps.sort(key=lambda x: int(re.search(r'_step(\d+)\.pth$', x).group(1)))            for old_ckp in step_ckps[:-5]:                os.remove(old_ckp)    # 训练结束,保存最终模型和epoch检查点    checkpoint = build_checkpoint(model, optimizer, scaler, epoch, step,                                  accumulation_step, iter_per_epoch, run if args.use_swanlab elseNone, shuffle_generator)    ckp_latest = f'{args.save_dir}/sft_{lm_config.dim}_{lm_config.n_layers}_{lm_config.vocab_size}.pth'    ckp_epoch = f'{args.save_dir}/sft_{lm_config.dim}_{lm_config.n_layers}_{lm_config.vocab_size}_epoch{epoch+1}.pth'    atomic_save(checkpoint, ckp_latest)    atomic_save(checkpoint, ckp_epoch)    Logger(f'训练完成,模型已保存: {ckp_epoch}')def init_model():    """初始化模型"""    def count_parameters(model):        return sum(p.numel() for p in model.parameters() if p.requires_grad)    # 加载分词器    tokenizer = AutoTokenizer.from_pretrained('./tokenizer/')    if tokenizer.pad_token_id isnotNone:        lm_config.pad_token_id = tokenizer.pad_token_id    # 初始化模型    model = Transformer(lm_config)    # 加载预训练权重    ckp = args.pretrain_path    if os.path.exists(ckp):        checkpoint = torch.load(ckp, map_location='cpu')        # 兼容新旧检查点格式        if isinstance(checkpoint, dict) and'model'in checkpoint:            state_dict = checkpoint['model']        else:            state_dict = checkpoint        # 处理 _orig_mod. 前缀(torch.compile 兼容性)        unwanted_prefix = '_orig_mod.'        for k, _ in list(state_dict.items()):            if k.startswith(unwanted_prefix):                state_dict[k[len(unwanted_prefix):]] = state_dict.pop(k)        model.load_state_dict(state_dict, strict=False)        Logger(f'已加载预训练权重: {ckp}')    else:        Logger(f'⚠️ 预训练权重不存在: {ckp},将从头开始SFT训练')    # 多卡初始化    if"cuda"in args.device:        num_gpus = torch.cuda.device_count()        if num_gpus > 1:            Logger(f"Using {num_gpus} GPUs with DataParallel!")            model = torch.nn.DataParallel(model)    elif"mps"in args.device:        Logger("Using Apple Silicon MPS")    model = model.to(args.device)    Logger(f'LLM总参数量:{count_parameters(model) / 1e6:.3f} 百万')    return model, tokenizerif __name__ == "__main__":    parser = argparse.ArgumentParser(description="SFT微调训练")    parser.add_argument("--out_dir", type=str, default="sft_model", help="输出目录")    parser.add_argument("--epochs", type=int, default=1, help="训练轮数")    parser.add_argument("--batch_size", type=int, default=48, help="批处理大小")    parser.add_argument("--learning_rate", type=float, default=2e-4, help="学习率")    parser.add_argument("--device", type=str, default=None, help="使用的设备 (自动检测: cuda/mps/cpu)")    parser.add_argument("--dtype", type=str, default="float32", help="数据类型 (MPS建议用float32)")    parser.add_argument("--use_swanlab", action="store_true", help="是否使用SwanLab进行实验跟踪")    parser.add_argument("--num_workers", type=int, default=0, help="数据加载的工作进程数")    parser.add_argument("--data_path", type=str, help="SFT训练数据路径")    parser.add_argument("--pretrain_path", type=str, default="./base_model/pretrain_768_12_6144.pth", help="预训练模型路径")    parser.add_argument("--accumulation_steps", type=int, default=8, help="梯度累积步数")    parser.add_argument("--grad_clip", type=float, default=1.0, help="梯度裁剪阈值")    parser.add_argument("--warmup_iters", type=int, default=0, help="预热迭代次数")    parser.add_argument("--log_interval", type=int, default=100, help="日志记录间隔")    parser.add_argument("--save_interval", type=int, default=100, help="模型保存间隔")    parser.add_argument("--max_samples", type=int, default=None, help="最大训练样本数,用于调试和节省内存")    parser.add_argument("--gpus", type=str, default=None, help="CUDA GPU ID,逗号分隔 (例如 '0,1,2')")    args = parser.parse_args()    # 设置可见GPU    if args.gpus isnotNoneand torch.cuda.is_available():        os.environ["CUDA_VISIBLE_DEVICES"] = args.gpus        args.device = "cuda:0"    elif args.device isNone:        args.device = get_device()    # ==================== 模型配置 ====================    lm_config = ModelConfig(        dim=768,        n_layers=12,    )    max_seq_len = lm_config.max_seq_len    args.save_dir = os.path.join(args.out_dir)    os.makedirs(args.out_dir, exist_ok=True)    # ==================== 检查点检测 ====================    ckp_path = f'{args.save_dir}/sft_{lm_config.dim}_{lm_config.n_layers}_{lm_config.vocab_size}.pth'    checkpoint = None    resume_run_id = None    saved_shuffle_generator_state = None    if os.path.exists(ckp_path):        try:            checkpoint = torch.load(ckp_path, map_location='cpu')            resume_run_id = checkpoint.get('swanlab_run_id', None)            saved_shuffle_generator_state = checkpoint.get('shuffle_generator_state', None)            Logger(f'发现检查点,将从 epoch {checkpoint["epoch"]}, step {checkpoint["step"]} 恢复训练')        except (RuntimeError, EOFError, pickle.UnpicklingError):            step_ckps = glob.glob(f'{args.save_dir}/sft_{lm_config.dim}_{lm_config.n_layers}_{lm_config.vocab_size}_step*.pth')            step_ckps.sort(key=lambda x: int(re.search(r'_step(\d+)\.pth$', x).group(1)))            if step_ckps:                checkpoint = torch.load(step_ckps[-1], map_location='cpu')                resume_run_id = checkpoint.get('swanlab_run_id', None)                saved_shuffle_generator_state = checkpoint.get('shuffle_generator_state', None)                Logger(f'最新检查点损坏,从历史检查点 {step_ckps[-1]} 恢复')            else:                Logger(f'最新检查点损坏且无历史检查点,从头开始训练')    # ==================== 实验跟踪初始化 ====================    run = None    if args.use_swanlab:        run = swanlab.init(            project="My-LLM",            experiment_name="SFT",            config=vars(args),            id=resume_run_id,            resume="allow",        )    # 设置随机种子    torch.manual_seed(42)    device_type = "cuda"if"cuda"in args.device else ("mps"if"mps"in args.device else"cpu")    # 设置混合精度训练    if device_type == "cuda":        ptdtype = {'float32': torch.float32, 'bfloat16': torch.bfloat16, 'float16': torch.float16}[args.dtype]        ctx = torch.cuda.amp.autocast(dtype=ptdtype)    else:        ctx = nullcontext()    # ==================== 模型和数据初始化 ====================    model, tokenizer = init_model()    # 创建数据集    train_ds = SFTDataset(args.data_path, tokenizer, max_length=max_seq_len, max_samples=args.max_samples)    # 创建显式 shuffle 生成器    shuffle_generator = torch.Generator()    shuffle_generator.manual_seed(42)    if saved_shuffle_generator_state isnotNone:        shuffle_generator.set_state(saved_shuffle_generator_state)        Logger('已恢复 DataLoader shuffle 生成器状态')    # 创建数据加载器    is_cuda = "cuda"in args.device    train_loader = DataLoader(        train_ds,        batch_size=args.batch_size,        pin_memory=is_cuda,        drop_last=False,        shuffle=True,        num_workers=args.num_workers if is_cuda else0,        generator=shuffle_generator,    )    # ==================== 优化器和训练组件初始化 ====================    scaler = torch.amp.GradScaler(device_type, enabled=(device_type == "cuda"and args.dtype in ['float16', 'bfloat16']))    optimizer = optim.AdamW(model.parameters(), lr=args.learning_rate)    # ==================== 加载检查点(如有) ====================    start_epoch = 0    resume_step = -1    resume_accumulation_step = 0    saved_iter_per_epoch = None    if checkpoint isnotNone:        model_state_dict = checkpoint['model']        is_dataparallel = isinstance(model, torch.nn.DataParallel)        state_dict_is_dataparallel = any(k.startswith('module.') for k in model_state_dict.keys())        if is_dataparallel andnot state_dict_is_dataparallel:            model_state_dict = {f'module.{k}': v for k, v in model_state_dict.items()}            Logger('兼容性处理:单卡检查点 → 多卡模型')        elifnot is_dataparallel and state_dict_is_dataparallel:            model_state_dict = {k.replace('module.', '', 1): v for k, v in model_state_dict.items()}            Logger('兼容性处理:多卡检查点 → 单卡模型')        model.load_state_dict(model_state_dict)        optimizer.load_state_dict(checkpoint['optimizer'])        scaler.load_state_dict(checkpoint['scaler'])        start_epoch = checkpoint['epoch']        resume_step = checkpoint['step']        resume_accumulation_step = checkpoint.get('accumulation_step', 0)        # 恢复随机状态        if'torch_rng_state'in checkpoint:            torch.random.set_rng_state(checkpoint['torch_rng_state'])            np.random.set_state(checkpoint['numpy_rng_state'])            random.setstate(checkpoint['python_rng_state'])            if checkpoint.get('cuda_rng_state') isnotNoneand torch.cuda.is_available():                torch.cuda.set_rng_state(checkpoint['cuda_rng_state'])            Logger('已恢复随机状态')        saved_iter_per_epoch = checkpoint.get('iter_per_epoch', None)        Logger(f'已恢复模型、优化器和scaler状态,从 epoch {start_epoch}, step {resume_step + 1} 继续训练'               f'(梯度累积进度: {resume_accumulation_step}/{args.accumulation_steps})')        del checkpoint    # ==================== 开始训练 ====================    iter_per_epoch = len(train_loader)    if saved_iter_per_epoch isnotNoneand saved_iter_per_epoch != iter_per_epoch:        Logger(f'⚠️ 警告:检查点记录 iter_per_epoch={saved_iter_per_epoch},当前为 {iter_per_epoch},'               f'学习率调度可能不连续')    for epoch in range(start_epoch, args.epochs):        train_epoch(epoch,                    resume_step if epoch == start_epoch else-1,                    resume_accumulation_step if epoch == start_epoch else0)

4.3 执行SFT训练

先查看SFT训练脚本的帮助说明,了解可用的命令行参数配置:

使用如下命令启动SFT训练,指定输出到./sft_model目录,开启SwanLab监控,并加载之前预训练好的模型作为初始参数:

python sft_train_llm.py --out_dir ./sft_model --use_swanlab --data_path ~/Desktop/train-llm/code/sft_data/BelleGroup_sft.jsonl

系统资源占用情况:

SwanLab监控的SFT训练指标:

训练结束后,SFT模型被保存在指定目录中,如下所示:

五、使用模型

模型训练完成后,可以加载模型并测试其实际的生成效果。通过对比预训练模型和SFT模型的回答效果,可以直观地感受到SFT带来的变化——预训练模型只会基于上下文进行续写,而SFT模型则会尝试以对话的方式回应用户的提问。

使用如下代码分别加载预训练模型和SFT微调后的模型,进行推理测试:

# model_test.pyimport osimport picklefrom contextlib import nullcontextimport torchfrom model_struct import ModelConfig, Transformerfrom transformers import AutoTokenizer, AutoModelForCausalLMimport argparseimport warnings# 忽略警告warnings.filterwarnings('ignore')class TextGenerator:    def __init__(self,                  checkpoint='./base_model/pretrain_768_12_6144.pth',  # 模型检查点路径                 tokenizer_model_path='./tokenizer/',  # 分词器模型路径                 seed=42,  # 随机种子,确保可重复性                 device=None,  # 设备,优先使用 CUDA,如果没有可用的 CUDA,则使用 CPU                 dtype="bfloat16"):# 数据类型,默认为 float32,可以选择 float16 或 bfloat16        """        初始化 TextGenerator 类,加载模型、设置设备和分词器等。        """        # 模型加载配置        self.checkpoint = checkpoint  # 保存的模型检查点路径        self.tokenizer_model_path = tokenizer_model_path  # 分词器模型文件路径        self.seed = seed  # 随机数种子,用于生成的可重复性        # 自动检测设备:CUDA > MPS > CPU        if device:            self.device = device        elif torch.cuda.is_available():            self.device = 'cuda:0'        elif torch.backends.mps.is_available():            self.device = 'mps'        else:            self.device = 'cpu'        self.dtype = dtype  # 模型的浮点数类型        self.device_type = 'cuda'if'cuda'in self.device else ('mps'if'mps'in self.device else'cpu')                # 设置随机种子,确保生成的可重复性        torch.manual_seed(seed)  # 设置 CPU 随机种子        if torch.cuda.is_available():            torch.cuda.manual_seed(seed)  # 设置 CUDA 随机种子            torch.backends.cuda.matmul.allow_tf32 = True# 允许 CUDA 使用 TF32 精度进行矩阵乘法运算            torch.backends.cudnn.allow_tf32 = True# 允许 cuDNN 使用 TF32 精度加速                # 根据 dtype 选择适当的自动混合精度上下文        # MPS 暂不完整支持 autocast,使用 nullcontext        if self.device_type == 'cuda':            ptdtype = {'float32': torch.float32, 'bfloat16': torch.bfloat16, 'float16': torch.float16}[self.dtype]            self.ctx = torch.amp.autocast(device_type='cuda', dtype=ptdtype)        else:            self.ctx = nullcontext()  # MPS/CPU 使用 float32        # 初始化分词器        self.tokenizer = AutoTokenizer.from_pretrained(self.tokenizer_model_path)  # 根据指定的路径加载分词器        # 加载模型检查点文件        checkpoint_dict = torch.load(self.checkpoint, map_location=self.device)  # 加载模型参数 # 初始化模型参数        self.model = Transformer(            ModelConfig(                dim=768,                n_layers=12,                pad_token_id=self.tokenizer.pad_token_id if self.tokenizer.pad_token_id isnotNoneelse0            )        )  # 实例化 Transformer 模型        sunwanted_prefix = '_orig_mod.'        for k, v in list(checkpoint_dict.items()):            if k.startswith(sunwanted_prefix):                checkpoint_dict[k[len(sunwanted_prefix):]] = checkpoint_dict.pop(k)        self.model.load_state_dict(checkpoint_dict, strict=False)                # 计算模型参数量        num_params = sum(p.numel() for p in self.model.parameters() if p.requires_grad)        print(f"Model has {num_params / 1e6:.3f} M parameters.")        # 设置模型为评估模式(evaluation mode),防止训练模式下的 dropout 等操作影响结果        self.model.eval()        # 将模型放置到正确的设备上(GPU 或 CPU)        self.model.to(self.device)    def chat_template(self, prompt):        message = [            {"role": "system", "content": "你是一个AI助手,你的名字叫小明。"},            {"role": "user", "content": prompt}        ]        return self.tokenizer.apply_chat_template(message, tokenize=False, add_generation_prompt=True)    def sft_sample(self,                start="Hello!",  # 生成文本的起始提示词,可以是任意字符串               num_samples=3,  # 生成样本的数量,默认生成 3 个样本               max_new_tokens=256,  # 每个样本生成的最大 token 数,默认最多生成 256 个 token               temperature=0.7,  # 控制生成的随机性,1.0 为标准,值越大越随机               top_k=300):# 保留概率最高的 top_k 个 token,限制生成时的选择范围        """        根据给定的起始文本生成样本。                :param start: 生成文本的起始提示词        :param num_samples: 要生成的文本样本数        :param max_new_tokens: 每个样本生成的最大 token 数        :param temperature: 控制生成的随机性,值越小生成越确定,值越大生成越随机        :param top_k: 限制生成时选择的 token 范围        :return: 生成的文本样本列表        """        start = self.chat_template(start)        # 将起始文本编码为 token id 序列        start_ids = self.tokenizer(start).data['input_ids']        # print('start_ids:', start_ids)        x = (torch.tensor(start_ids, dtype=torch.long, device=self.device)[None, ...])  # 将编码后的 token id 转为 PyTorch 张量        generated_texts = []  # 用于保存生成的文本样本        with torch.no_grad():  # 禁用梯度计算,提升效率            with self.ctx:  # 进入自动混合精度的上下文(如果是 GPU 并使用 float16 时)                for k in range(num_samples):  # 循环生成指定数量的样本                    y = self.model.generate(x, self.tokenizer.eos_token_id, max_new_tokens, temperature=temperature, top_k=top_k)  # 生成文本                    generated_texts.append(self.tokenizer.decode(y[0].tolist()))  # 解码生成的 token 序列为可读文本        return generated_texts  # 返回生成的文本样本    def pretrain_sample(self,                start="Hello!",  # 生成文本的起始提示词,可以是任意字符串               num_samples=3,  # 生成样本的数量,默认生成 3 个样本               max_new_tokens=256,  # 每个样本生成的最大 token 数,默认最多生成 256 个 token               temperature=0.7,  # 控制生成的随机性,1.0 为标准,值越大越随机               top_k=300):# 保留概率最高的 top_k 个 token,限制生成时的选择范围        """        根据给定的起始文本生成样本。                :param start: 生成文本的起始提示词        :param num_samples: 要生成的文本样本数        :param max_new_tokens: 每个样本生成的最大 token 数        :param temperature: 控制生成的随机性,值越小生成越确定,值越大生成越随机        :param top_k: 限制生成时选择的 token 范围        :return: 生成的文本样本列表        """        # 如果 start 是以 'FILE:' 开头,表示从文件中读取起始文本        if start.startswith('FILE:'):            with open(start[5:], 'r', encoding='utf-8') as f:                start = f.read()  # 读取文件内容作为起始文本                # 将起始文本编码为 token id 序列        start_ids = self.tokenizer(start).data['input_ids']        # print('start_ids:', start_ids)        x = (torch.tensor(start_ids, dtype=torch.long, device=self.device)[None, ...])  # 将编码后的 token id 转为 PyTorch 张量        # print(x.shape)        generated_texts = []  # 用于保存生成的文本样本        with torch.no_grad():  # 禁用梯度计算,提升效率            with self.ctx:  # 进入自动混合精度的上下文(如果是 GPU 并使用 float16 时)                for k in range(num_samples):  # 循环生成指定数量的样本                    y = self.model.generate(x, max_new_tokens=max_new_tokens, temperature=temperature, top_k=top_k)  # 生成文本                    generated_texts.append(self.tokenizer.decode(y[0].tolist()))  # 解码生成的 token 序列为可读文本                return generated_texts  # 返回生成的文本样本    if __name__ == "__main__":    print("------------------- Pretrain 测试 ------------------- \n")    pretrain_prompt_datas = [        '<|im_start|>在计算税款损失时,要不要将进项留抵税额包括在内?',        '<|im_start|>你好,这是一个LLM训练',        '<|im_start|>北京天安门是'            ]    generator = TextGenerator(checkpoint='./base_model/pretrain_768_12_6144.pth')  # 初始化生成器    for i in range(len(pretrain_prompt_datas)):        samples = generator.pretrain_sample(start=pretrain_prompt_datas[i], num_samples=1, max_new_tokens=120, temperature=0.75)        print(f"\nSample {i+1}:\n{pretrain_prompt_datas[i]}{samples[0]}\n{'-'*20}")  # 打印生成的样本并用分隔线分割    print("\n ------------------- SFT 测试 ------------------- \n")    sft_prompt_datas = [        '你好呀',        "苏州是哪个省的?",        "1+2等于多少?",        "你是谁?"    ]    generator = TextGenerator(checkpoint='./sft_model/sft_768_12_6144.pth')  # 初始化生成器    for i in range(len(sft_prompt_datas)):        samples = generator.sft_sample(start=sft_prompt_datas[i], num_samples=1, max_new_tokens=128, temperature=0.6)        print(f"\nSample {i+1}:\nQuestion: {sft_prompt_datas[i]} \nAI answer: {samples[0]}\n{'-'*20}")  # 打印生成的样本并用分隔线分割

需要说明的是,真正训练出一个效果良好的模型需要大量的计算资源和时间(“大力出奇迹”),这里使用的是训练了没多久的checkpoint检查点,因此模型的生成效果不尽如人意。但这并不影响学习和理解整个训练流程。

测试结果如下,权当一乐儿:

六、导出模型

经过预训练和SFT微调后的模型,通常需要导出为标准化的格式,以便分发给其他人使用,或者上传到Hugging Face等模型平台供社区共享。导出操作将模型的状态字典(State Dict)与配置文件(Config)一起打包,生成一个完整的、可独立加载的模型目录。

6.1 导出模型

导出使用如下代码,将训练好的模型参数和配置统一保存到指定目录中:

# export_model.pyimport torchimport warningsfrom transformers import AutoTokenizerfrom model_struct import Transformer, ModelConfig# 忽略警告warnings.filterwarnings('ignore')def count_parameters(model):    return sum(p.numel() for p in model.parameters() if p.requires_grad)def export_model(tokenizer_path, model_config, model_ckpt_path, save_directory):    # 注册自定义类和配置    ModelConfig.register_for_auto_class()    Transformer.register_for_auto_class("AutoModelForCausalLM")    # 加载tokenizer    tokenizer = AutoTokenizer.from_pretrained(        tokenizer_path,        trust_remote_code=True,        use_fast=False    )    if tokenizer.pad_token_id isnotNone:        model_config.pad_token_id = tokenizer.pad_token_id    # 初始化模型    model = Transformer(model_config)    # 自动检测设备:CUDA > MPS > CPU    if torch.cuda.is_available():        device = torch.device('cuda')    elif torch.backends.mps.is_available():        device = torch.device('mps')    else:        device = torch.device('cpu')    # 加载模型权重    state_dict = torch.load(model_ckpt_path, map_location=device)    # 移除可能存在的多余前缀    unwanted_prefix = '_orig_mod.'    for k in list(state_dict.keys()):        if k.startswith(unwanted_prefix):            state_dict[k[len(unwanted_prefix):]] = state_dict.pop(k)        # 加载权重到模型    model.load_state_dict(state_dict, strict=False)    print(f'模型参数: {count_parameters(model)/1e6:.2f}M = {count_parameters(model)/1e9:.2f}B')    # 保存完整模型和tokenizer    model.save_pretrained(save_directory, safe_serialization=False)    tokenizer.save_pretrained(save_directory)    print(f'模型和tokenizer已保存至: {save_directory}')if __name__ == '__main__':    config = ModelConfig(        dim=768,        n_layers=12,    )    export_model(        tokenizer_path='./tokenizer/',        model_config=config,        model_ckpt_path='./sft_model/sft_768_12_6144.pth',        save_directory="My-LLM-83M"    )

运行导出脚本,控制台输出如下,显示模型被序列化保存到磁盘:

导出的模型目录结构如下,包含了模型配置文件(config.json)、模型权重文件(pytorch_model.bin)以及其他必要的文件:

6.2 推理

导出的模型可以很方便地分发给其他人使用,或者上传到Hugging Face平台供社区下载。那么,拿到这个模型后,该如何加载和使用?下面提供了完整的推理代码示例,展示如何从导出的目录中加载模型Tokenizer,并进行推理:

# inference.pyimport torchfrom transformers import AutoTokenizer, AutoModelForCausalLM# 选择设备:Mac 用 MPS,有 CUDA 用 CUDA,否则 CPUif torch.backends.mps.is_available():    device = "mps"elif torch.cuda.is_available():    device = "cuda"else:    device = "cpu"print(f"使用设备: {device}")# 加载模型和 tokenizermodel_path = "./My-LLM-83M"tokenizer = AutoTokenizer.from_pretrained(model_path, trust_remote_code=True)model = AutoModelForCausalLM.from_pretrained(    model_path,    trust_remote_code=True,    torch_dtype=torch.float32,    device_map=None# 手动管理设备,避免 device_map="auto" 在 Mac 上回退 CPU)model = model.to(device)model.eval()print("=" * 50)print("My-LLM-83M 对话模式")print("输入内容开始对话,输入 quit 退出")print("=" * 50)history = ""whileTrue:    user_input = input("\n你: ").strip()    ifnot user_input:        continue    if user_input.lower() in ("quit", "exit", "q"):        print("再见!")        break    history += f"{tokenizer.bos_token}{user_input}{tokenizer.eos_token}"    input_ids = tokenizer(history, return_tensors="pt").input_ids.to(device)    with torch.no_grad():        output = model.generate_super(            input_ids,            stop_id=tokenizer.eos_token_id,            max_new_tokens=256,            temperature=0.8,            top_k=50        )    response = tokenizer.decode(output[0], skip_special_tokens=True)    print(f"AI: {response}")    history += f"{response}{tokenizer.eos_token}"

推理运行结果如下(训练时间较短,输出效果有限,仅供流程参考):

七、总结

这篇文章完整演示了从零构建一个小型LLM的端到端流程:首先手动实现了包括RMSNorm、分组查询注意力(GQA)和SwiGLU前馈网络在内的核心模块,组装出一个可用的Decoder-only模型;然后使用序列猴子数据集进行预训练,让模型学习通用的语言表示和世界知识;接着在预训练基础上,利用BelleGroup对话数据集进行有监督微调(SFT),使模型具备对话交互和指令遵循的能力;最后展示了模型的推理测试和标准化导出流程。需要说明的是,受限于计算资源和训练时间,这里的模型规模和训练步数都较为有限,最终效果尚不能与商业大模型相提并论,但核心目的在于跑通从架构设计、数据准备、预训练、SFT微调到模型导出的完整技术链路,帮助对大语言模型的训练全流程建立一个直观而系统的认知,为后续更深入的研究和实践奠定基础。

学AI大模型的正确顺序,千万不要搞错了

🤔2026年AI风口已来!各行各业的AI渗透肉眼可见,超多公司要么转型做AI相关产品,要么高薪挖AI技术人才,机遇直接摆在眼前!

有往AI方向发展,或者本身有后端编程基础的朋友,直接冲AI大模型应用开发转岗超合适!

就算暂时不打算转岗,了解大模型、RAG、Prompt、Agent这些热门概念,能上手做简单项目,也绝对是求职加分王🔋

在这里插入图片描述

📝给大家整理了超全最新的AI大模型应用开发学习清单和资料,手把手帮你快速入门!👇👇

学习路线:

✅大模型基础认知—大模型核心原理、发展历程、主流模型(GPT、文心一言等)特点解析
✅核心技术模块—RAG检索增强生成、Prompt工程实战、Agent智能体开发逻辑
✅开发基础能力—Python进阶、API接口调用、大模型开发框架(LangChain等)实操
✅应用场景开发—智能问答系统、企业知识库、AIGC内容生成工具、行业定制化大模型应用
✅项目落地流程—需求拆解、技术选型、模型调优、测试上线、运维迭代
✅面试求职冲刺—岗位JD解析、简历AI项目包装、高频面试题汇总、模拟面经

以上6大模块,看似清晰好上手,实则每个部分都有扎实的核心内容需要吃透!

我把大模型的学习全流程已经整理📚好了!抓住AI时代风口,轻松解锁职业新可能,希望大家都能把握机遇,实现薪资/职业跃迁~

这份完整版的大模型 AI 学习资料已经上传CSDN,朋友们如果需要可以微信扫描下方CSDN官方认证二维码免费领取【保证100%免费

在这里插入图片描述

Logo

汇聚全球AI编程工具,助力开发者即刻编程。

更多推荐