Hands-on Building a Large Model
5.1 Hands-on Implementation of a LLaMA2 Large Model
Meta (formerly Facebook) released the first large language model based on the Transformer architecture, LLaMA, in February 2023, and then released the same series model LLaMA2 in July of the same year. In Chapter 4, we have learned and understood LLMs, as well as how to train LLMs. In this section, we will learn how to implement a LLaMA2 model hands-on.
The structure of the LLaMA2 model is shown in Figure 5.1:

Figure 5.1 LLaMA2 Structure
5.1.1 Define Hyperparameters
First, we need to define some hyperparameters, which include the size of the model, number of layers, number of heads, embedding dimension, hidden layer dimension, etc. These hyperparameters can be adjusted according to actual conditions.
Here, we define a ModelConfig class to store and record our hyperparameters. Here, we inherit from the PretrainedConfig class, which is a parameter class in the transformers library. By inheriting this class, we can conveniently use some functions in the transformers library and also facilitate exporting Hugging Face models later.
from transformers import PretrainedConfig
class ModelConfig(PretrainedConfig):
model_type = "Tiny-K"
def __init__(
self,
dim: int = 768, # Model dimension
n_layers: int = 12, # Number of Transformer layers
n_heads: int = 16, # Number of attention heads
n_kv_heads: int = 8, # Number of key-value heads
vocab_size: int = 6144, # Vocabulary size
hidden_dim: int = None, # Hidden layer dimension
multiple_of: int = 64,
norm_eps: float = 1e-5, # Normalization layer's epsilon
max_seq_len: int = 512, # Maximum sequence length
dropout: float = 0.0, # Dropout probability
flash_attn: bool = True, # Whether to use 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)
In the following code, when
argsappears, it is assumed to be the aboveModelConfigparameter configuration.
Let's look at the meaning of some of these hyperparameters, such as dim is the model dimension, n_layers is the number of Transformer layers, n_heads is the number of attention heads, vocab_size is the vocabulary size, max_seq_len is the maximum sequence length, etc. The above code also provides detailed comments for each parameter, and in the subsequent code, we will build our model based on these hyperparameters.
5.1.2 Build RMSNorm
RMSNorm can be represented by the following mathematical formula:
Where:
- is the th element of the input vector
- is a learnable scaling parameter (corresponding to
self.weightin the code) - is the number of dimensions of the input vector
- is a small constant used for numerical stability (to avoid division by zero)
This normalization helps stabilize the learning process by ensuring that the scale of weights does not become too large or too small, which is especially useful in deep learning models with many layers.
We can implement RMSNorm as follows:
class RMSNorm(nn.Module):
def __init__(self, dim: int, eps: float):
super().__init__()
# eps is to prevent division by zero
self.eps = eps
# weight is a learnable parameter, initialized to 1
self.weight = nn.Parameter(torch.ones(dim))
def _norm(self, x):
# Calculate the core part of RMSNorm
# x.pow(2).mean(-1, keepdim=True) calculates the mean of the square of the input x
# torch.rsqrt is the reciprocal of the square root, so it gets the denominator part of RMSNorm, and adds eps to prevent the denominator from being zero
# Finally, multiply by x to get the result of RMSNorm
return x * torch.rsqrt(x.pow(2).mean(-1, keepdim=True) + self.eps)
def forward(self, x):
# The forward function is the forward propagation of the model
# First, convert the input x to float type, then perform RMSNorm, and then convert back to the original data type
# Finally, multiply by weight, which is a learnable scaling factor of RMSNorm
output = self._norm(x.float()).type_as(x)
return output * self.weight
And we can test the RMSNorm module with the following code, which shows that the output shape is torch.Size([1, 50, 288]), consistent with the input shape, indicating that the module implementation is correct, and normalization does not change the input shape.
norm = RMSNorm(args.dim, args.norm_eps)
x = torch.randn(1, 50, args.dim)
output = norm(x)
print(output.shape)
out:
torch.Size([1, 50, 768])
5.1.3 Build LLaMA2 Attention
In the LLaMA2 model, although only the LLaMA2-70B model uses the grouped query attention mechanism (Grouped-Query Attention, GQA), we still choose to use GQA to build our LLaMA Attention module, which can improve the efficiency of the model and save some GPU memory usage.

Figure 5.2 LLaMA2 Attention Structure
5.1.3.1 repeat_kv
In the LLaMA2 model, we need to expand the dimensions of the keys and values to match the dimensions of the queries to perform attention calculations. We can implement repeat_kv as follows:
def repeat_kv(x: torch.Tensor, n_rep: int) -> torch.Tensor:
# Get the shape of the input tensor: batch size, sequence length, number of key/value heads, dimension size of each head
bs, slen, n_kv_heads, head_dim = x.shape
# If the repetition count is 1, no repetition is needed, return the original tensor directly
if n_rep == 1:
return x
# Expand and reshape the tensor to repeat the key-value pairs
return (
x[:, :, :, None, :] # Add a new dimension after the third dimension (the head dimension)
.expand(bs, slen, n_kv_heads, n_rep, head_dim) # Expand the newly added dimension to n_rep size, achieving the effect of repetition
.reshape(bs, slen, n_kv_heads * n_rep, head_dim) # Reshape to merge the number of key/value heads and the repetition count dimension
)
In the above code:
-
First, we get the shape of the input tensor: first, the code uses
x.shapeto get the shape of the input tensor, including the batch size (bs), sequence length (slen), the number of key/value heads (n_kv_heads), and the dimension size of each head (head_dim). -
Then, check the repetition count: Next, the code checks whether the repetition count
n_repis 1. If it is 1, it means that there is no need to repeat the keys and values, and the original tensorxis returned directly. -
Finally, expand and reshape the tensor:
- Add a new dimension after the third dimension (the head dimension) to form
x[:, :, :, None, :]. - Use the
expandmethod to expand the newly added dimension ton_repsize, achieving the effect of repeating the key-value pairs. - Finally, use the
reshapemethod to reshape, merging the dimensions of the number of key/value heads and the repetition count, thus achieving the desired shape for the queries.
- Add a new dimension after the third dimension (the head dimension) to form
5.1.3.2 Rotary Embedding
Next, we proceed to implement the rotary embedding, which is an important component of the LLaMA2 model. It provides stronger contextual information for the attention mechanism, thereby improving the performance of the model.
First, we need to construct functions to obtain the real and imaginary parts of the rotary embeddings:
# Note: The dim here should be dim//n_head, because we are applying rotary embeddings to each head
def precompute_freqs_cis(dim: int, end: int, theta: float = 10000.0):
# torch.arange(0, dim, 2)[: (dim // 2)].float() generates a sequence starting from 0, step of 2, with a length of half of dim
# Then each element is divided by dim, and then take the reciprocal of theta to get the frequency
freqs = 1.0 / (theta ** (torch.arange(0, dim, 2)[: (dim // 2)].float() / dim))
# Generate a sequence from 0 to end with a length of end
t = torch.arange(end, device=freqs.device)
# Compute the outer product, resulting in a two-dimensional matrix, where each row is the element of t multiplied by the element of freqs
freqs = torch.outer(t, freqs).float()
# Compute the cosine of the frequencies, obtaining the real part
freqs_cos = torch.cos(freqs)
# Compute the sine of the frequencies, obtaining the imaginary part
freqs_sin = torch.sin(freqs)
return freqs_cos, freqs_sin
- Calculate the frequency sequence:
torch.arange(0, dim, 2)[: (dim // 2)].float()generates a sequence starting from 0, with a step of 2, with a length of half ofdim.- Each element is divided by
dimand then the reciprocal ofthetais taken to get a frequency sequencefreqs. This step is to generate suitable frequencies for the rotary embeddings.
- Generate a time sequence:
t = torch.arange(end, device=freqs.device)generates a sequence from0toendwith a length ofend.endis usually the maximum length of the sequence.
- Compute the outer product
freqs = torch.outer(t, freqs).float()computes the outer product of the time sequencetand the frequency sequencefreqs, resulting in a two-dimensional matrixfreqs. Each row is the element of the time sequencetmultiplied by the element of the frequency sequencefreqs.
- Compute the real and imaginary parts
freqs_cos = torch.cos(freqs)computes the cosine of the frequency matrixfreqs, obtaining the real part of the rotary embedding.freqs_sin = torch.sin(freqs)computes the sine of the frequency matrixfreqs, obtaining the imaginary part of the rotary embedding.
Finally, the function returns two matrices freqs_cos and freqs_sin, which represent the real and imaginary parts of the rotary embeddings, respectively, for use in subsequent calculations.
Next, we construct the reshape_for_broadcast function to adjust the shape of freqs_cis, which mainly aims to align the shape of freqs_cis with x during broadcasting operations to enable correct tensor operations.
def reshape_for_broadcast(freqs_cis: torch.Tensor, x: torch.Tensor):
# Get the number of dimensions of x
ndim = x.ndim
# Assert that 0 <= 1 < ndim
assert 0 <= 1 < ndim
# Assert that the shape of freqs_cis matches the second dimension and last dimension of x
assert freqs_cis.shape == (x.shape[1], x.shape[-1])
# Construct a new shape, except for the second dimension and the last dimension, all other dimensions are set to 1, so that freqs_cis can be broadcast with x
shape = [d if i == 1 or i == ndim - 1 else 1 for i, d in enumerate(x.shape)]
# Adjust the shape of freqs_cis and return
return freqs_cis.view(shape)
Finally, we can implement the rotary embedding as follows:
def apply_rotary_emb(
xq: torch.Tensor,
xk: torch.Tensor,
freqs_cos: torch.Tensor,
freqs_sin: torch.Tensor
) -> Tuple[torch.Tensor, torch.Tensor]:
# Convert the query and key tensors to float and reshape to separate real and imaginary parts
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)
# Reshape the frequency tensor for broadcasting
freqs_cos = reshape_for_broadcast(freqs_cos, xq_r)
freqs_sin = reshape_for_broadcast(freqs_sin, xq_r)
# Apply rotation, calculating the rotated real and imaginary parts separately
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
# Merge the last two dimensions and restore the original tensor shape
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)
Here, we provide code to test the apply_rotary_emb function. You can also add breakpoints in the code to view the calculation results at each step.
xq = torch.randn(1, 50, 6, 48) # bs, seq_len, dim//n_head, n_head_dim
xk = torch.randn(1, 50, 6, 48) # bs, seq_len, dim//n_head, n_head_dim
# Use precompute_freqs_cis function to get sin and cos
cos, sin = precompute_freqs_cis(288//6, 50)
print(cos.shape, sin.shape)
xq_out, xk_out = apply_rotary_emb(xq, xk, cos, sin)
xq_out.shape, xk_out.shape
OUT:
torch.Size([50, 24]) torch.Size([50, 24])
(torch.Size([1, 50, 6, 48]), torch.Size([1, 50, 6, 48]))
5.1.3.3 Assemble LLaMA2 Attention
We have completed the implementation of the rotary embedding above, and now we can build the LLaMA2 Attention module.
class Attention(nn.Module):
def __init__(self, args: ModelConfig):
super().__init__()
# Determine the number of heads for keys (key) and values (value) based on whether n_kv_heads is specified.
self.n_kv_heads = args.n_heads if args.n_kv_heads is None else args.n_kv_heads
# Ensure the total number of heads can be divided by the number of key-value heads.
assert args.n_heads % self.n_kv_heads == 0
# Model parallel processing size, default is 1.
model_parallel_size = 1
# Local head count, equal to the total number of heads divided by the model parallel processing size.
self.n_local_heads = args.n_heads // model_parallel_size
# Local key-value head count, equal to the key-value head count divided by the model parallel processing size.
self.n_local_kv_heads = self.n_kv_heads // model_parallel_size
# Repetition count, used to expand the dimensions of keys and values.
self.n_rep = self.n_local_heads // self.n_local_kv_heads
# Dimension per head, equal to the model dimension divided by the total number of heads.
self.head_dim = args.dim // args.n_heads
# Define weight matrices.
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)
# Output weight matrix.
self.wo = nn.Linear(args.n_heads * self.head_dim, args.dim, bias=False)
# Define dropout.
self.attn_dropout = nn.Dropout(args.dropout)
self.resid_dropout = nn.Dropout(args.dropout)
# Save dropout probability.
self.dropout = args.dropout
# Check if Flash Attention is used (requires PyTorch >= 2.0).
self.flash = hasattr(torch.nn.functional, 'scaled_dot_product_attention')
if not self.flash:
# If Flash Attention is not supported, use a manually implemented attention mechanism and set the mask.
print("WARNING: using slow attention. Flash Attention requires PyTorch >= 2.0")
# Create an upper triangular matrix to mask future information.
mask = torch.full((1, 1, args.max_seq_len, args.max_seq_len), float("-inf"))
mask = torch.triu(mask, diagonal=1)
# Register as a buffer of the model
self.register_buffer("mask", mask)
def forward(self, x: torch.Tensor, freqs_cos: torch.Tensor, freqs_sin: torch.Tensor):
# Get batch size and sequence length, [batch_size, seq_len, dim]
bsz, seqlen, _ = x.shape
# Calculate queries (Q), keys (K), and values (V).
xq, xk, xv = self.wq(x), self.wk(x), self.wv(x)
# Reshape to adapt to the head dimension.
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)
# Apply rotary position embeddings (RoPE).
xq, xk = apply_rotary_emb(xq, xk, freqs_cos, freqs_sin)
# Expand keys and values to match the repetition count.
xk = repeat_kv(xk, self.n_rep)
xv = repeat_kv(xv, self.n_rep)
# Treat the heads as a batch dimension.
xq = xq.transpose(1, 2)
xk = xk.transpose(1, 2)
xv = xv.transpose(1, 2)
# Depending on whether Flash Attention is supported, choose the implementation.
if self.flash:
# Use Flash Attention.
output = torch.nn.functional.scaled_dot_product_attention(xq, xk, xv, attn_mask=None, dropout_p=self.dropout if self.training else 0.0, is_causal=True)
else:
# Use manually implemented attention mechanism.
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)
# Restore the time dimension and merge heads.
output = output.transpose(1, 2).contiguous().view(bsz, seqlen, -1)
# Project back to the residual stream.
output = self.wo(output)
output = self.resid_dropout(output)
return output
Similarly, you can use the following code to test the attention module, and you can see that the final output shape is torch.Size([1, 50, 768]), consistent with the input shape, indicating that the module implementation is correct.
# Create Attention instance
attention_model = Attention(args)
# Simulate input data
batch_size = 1
seq_len = 50 # Assuming the actual sequence length used is 50
dim = args.dim
x = torch.rand(batch_size, seq_len, dim) # Randomly generate input tensor
# freqs_cos = torch.rand(seq_len, dim // 2) # Simulate cos frequencies for RoPE
# freqs_sin = torch.rand(seq_len, dim // 2) # Simulate sin frequencies for RoPE
freqs_cos, freqs_sin = precompute_freqs_cis(dim//args.n_heads, seq_len)
# Run the Attention model
output = attention_model(x, freqs_cos, freqs_sin)
# The shape after attention is still [batch_size, seq_len, dim]
print("Output shape:", output.shape)
OUT:
Output shape: torch.Size([1, 50, 768])
5.1.4 Build LLaMA2 MLP Module
Compared to the LLaMA2 Attention module we implemented earlier, the implementation of the LLaMA2 MLP module is simpler. We can implement MLP as follows:
class MLP(nn.Module):
def __init__(self, dim: int, hidden_dim: int, multiple_of: int, dropout: float):
super().__init__()
# If the hidden dimension is not specified, we set it to 4 times the input dimension
# Then reduce it to 2/3, and finally ensure it is a multiple of multiple_of
if hidden_dim is None:
hidden_dim = 4 * dim
hidden_dim = int(2 * hidden_dim / 3)
hidden_dim = multiple_of * ((hidden_dim + multiple_of - 1) // multiple_of)
# Define the first linear transformation from input dimension to hidden dimension
self.w1 = nn.Linear(dim, hidden_dim, bias=False)
# Define the second linear transformation from hidden dimension to input dimension
self.w2 = nn.Linear(hidden_dim, dim, bias=False)
# Define the third linear transformation from input dimension to hidden dimension
self.w3 = nn.Linear(dim, hidden_dim, bias=False)
# Define the dropout layer to prevent overfitting
self.dropout = nn.Dropout(dropout)
def forward(self, x):
# Forward propagation function
# First, the input x passes through the first linear transformation and SILU activation function
# Then, the result is multiplied by the result of the input x passing through the third linear transformation
# Finally, it passes through the second linear transformation and dropout layer
return self.dropout(self.w2(F.silu(self.w1(x)) * self.w3(x)))
We focus on observing the implementation of the forward function. First, the input x passes through the first linear transformation self.w1 and the SILU activation function, then the result is multiplied by the result of the input x passing through the third linear transformation self.w3. Finally, it passes through the second linear transformation self.w2 and the dropout layer to get the final output.
Similarly, you can use the following code to test the LLaMAMLP module, and you can see that the final output shape is torch.Size([1, 50, 768]), consistent with the input shape, indicating that the module implementation is correct.
# Create MLP instance
mlp = MLP(args.dim, args.hidden_dim, args.multiple_of, args.dropout)
# Randomly generate data
x = torch.randn(1, 50, args.dim)
# Run MLP model
output = mlp(x)
print(output.shape)
OUT:
torch.Size([1, 50, 768])
5.1.5 LLaMA2 Decoder Layer
At this point, we have implemented the LLaMA2 model's Attention module and MLP module. Next, we can build the LLaMA2 Decoder Layer.
class DecoderLayer(nn.Module):
def __init__(self, layer_id: int, args: ModelConfig):
super().__init__()
# Define the number of heads for multi-head attention
self.n_heads = args.n_heads
# Define the input dimension
self.dim = args.dim
# Define the dimension per head, equal to the input dimension divided by the number of heads
self.head_dim = args.dim // args.n_heads
# Define the LLaMA2Attention object for multi-head attention calculation
self.attention = Attention(args)
# Define the LLaMAMLP object for feed-forward neural network calculation
self.feed_forward = MLP(
dim=args.dim,
hidden_dim=args.hidden_dim,
multiple_of=args.multiple_of,
dropout=args.dropout,
)
# Define the layer ID
self.layer_id = layer_id
# Define the normalization layer for attention calculation
self.attention_norm = RMSNorm(args.dim, eps=args.norm_eps)
# Define the normalization layer for feed-forward neural network calculation
self.ffn_norm = RMSNorm(args.dim, eps=args.norm_eps)
def forward(self, x, freqs_cos, freqs_sin):
# Forward propagation function
# First, the input x passes through the attention normalization layer, then performs attention calculation, the result is added to the input x to get h
# Then, h passes through the feed-forward neural network normalization layer, then performs feed-forward neural network calculation, the result is added to h to get the output
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
The DecoderLayer combines the Attention module and MLP module we completed above, implementing a complete Transformer module.
Similarly, you can use the following code to test the DecoderLayer module, and you can see that the final output shape is torch.Size([1, 50, 768]), consistent with the input shape, indicating that the module implementation is correct.
# Create LLaMADecoderLayer instance
decoderlayer = DecoderLayer(0, args)
# Simulate input data
dim = args.dim
seq_len = 50
x = torch.randn(1, seq_len, dim) # [bs, seq_len, dim]
freqs_cos, freqs_sin = precompute_freqs_cis(dim//args.n_heads, seq_len)
out = decoderlayer(x, freqs_cos, freqs_sin)
print(out.shape) # Shape is the same as the input x [batch_size, seq_len, dim]
OUT:
torch.Size([1, 50, 768])
5.1.6 Build LLaMA2 Model
Okay, we have completed the implementation of all the modules mentioned above, and now it's the exciting moment, we can build the LLaMA2 model. The LLaMA2 model is simply stacking the DecoderLayer modules to form a complete Transformer model.
class Transformer(PreTrainedModel):
config_class = ModelConfig # Configuration class
last_loss: Optional[torch.Tensor] # Record the loss of the last calculation
def __init__(self, args: ModelConfig = None):
super().__init__(args)
# Initialize model parameters
self.args = args
# Vocabulary size
self.vocab_size = args.vocab_size
# Number of layers
self.n_layers = args.n_layers
# Token embedding layer
self.tok_embeddings = nn.Embedding(args.vocab_size, args.dim)
# Dropout layer
self.dropout = nn.Dropout(args.dropout)
# Decoder layers
self.layers = torch.nn.ModuleList()
for layer_id in range(args.n_layers):
self.layers.append(DecoderLayer(layer_id, args))
# Normalization layer
self.norm = RMSNorm(args.dim, eps=args.norm_eps)
# Output layer
self.output = nn.Linear(args.dim, args.vocab_size, bias=False)
# Share the weights of the token embedding layer with the output layer
self.tok_embeddings.weight = self.output.weight
# Precompute relative position embeddings' frequencies
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)
# Initialize all weights
self.apply(self._init_weights)
# Special initialization for residual projection
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))
# Initialize the attribute for the last forward pass loss
self.last_loss = None
self.OUT = CausalLMOutputWithPast() # Output container
self._no_split_modules = [name for name, _ in self.named_modules()] # List of modules not to split
def _init_weights(self, module):
# Weight initialization function
if isinstance(module, nn.Linear):
torch.nn.init.normal_(module.weight, mean=0.0, std=0.02)
if module.bias is not None:
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], input token tensor.
- targets: Optional[torch.Tensor], target token tensor.
- kv_cache: bool, whether to use key-value cache.
- kwargs: other keyword arguments.
- self.OUT: CausalLMOutputWithPast, contains logits and loss.
"""
if 'input_ids' in kwargs:
tokens = kwargs['input_ids']
if 'attention_mask' in kwargs:
targets = kwargs['attention_mask']
# Forward propagation function
_bsz, seqlen = tokens.shape
# Through the token embedding layer and Dropout layer
h = self.tok_embeddings(tokens)
h = self.dropout(h)
# Get relative position embedding frequencies
freqs_cos = self.freqs_cos[:seqlen]
freqs_sin = self.freqs_sin[:seqlen]
# Through Decoder layers
for layer in self.layers:
h = layer(h, freqs_cos, freqs_sin)
# Through normalization layer
h = self.norm(h)
if targets is not None:
# If targets are provided, calculate loss
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:
# Inference optimization: only propagate the last position's output
logits = self.output(h[:, [-1], :])
self.last_loss = None
# Set output
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):
"""
Given an input sequence idx (a long integer tensor of shape (bz,seq_len)), generate new tokens multiple times to complete the sequence.
Runs in model.eval() mode. A less efficient sampling version without using key-value cache.
"""
index = idx.shape[1]
for _ in range(max_new_tokens):
# If the sequence context is too long, truncate it to the maximum length
idx_cond = idx if idx.size(1) <= self.args.max_seq_len else idx[:, -self.args.max_seq_len:]
# Forward propagate to get the logits of the last position in the sequence
logits = self(idx_cond).logits
logits = logits[:, -1, :] # Keep only the last time step's output
if temperature == 0.0:
# Select the most likely index
_, idx_next = torch.topk(logits, k=1, dim=-1)
else:
# Scale logits and apply softmax
logits = logits / temperature
if top_k is not None:
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
# Add the sampled index to the sequence and continue
idx = torch.cat((idx, idx_next), dim=1)
return idx[:, index:] # Return only the generated tokens
Similarly, you can use the following code to test the Transformer module, and you can see that the final output shape is torch.Size([1, 1, 6144]), consistent with the input shape, indicating that the module implementation is correct.
# LLaMA2Model.forward accepts two parameters, tokens and targets, where tokens is the input tensor, should be int type
x = torch.randint(0, 6144, (1, 50))
...