Skip to content
Open
Show file tree
Hide file tree
Changes from all commits
Commits
File filter

Filter by extension

Filter by extension


Conversations
Failed to load comments.
Loading
Jump to
Jump to file
Failed to load files.
Loading
Diff view
Diff view
147 changes: 96 additions & 51 deletions llama/generation.py
Original file line number Diff line number Diff line change
Expand Up @@ -48,6 +48,84 @@ class ChatPrediction(TypedDict, total=False):
UNSAFE_ERROR = "Error: special tags are not allowed as part of the prompt."




class ChatFormatter:
def format_dialog(self, dialog: Dialog) -> List[int]:
raise NotImplementedError


class Llama2Formatter(ChatFormatter):
def __init__(self, tokenizer: Tokenizer):
self.tokenizer = tokenizer

def format_dialog(self, dialog: Dialog) -> List[int]:
if not dialog:
raise ValueError("Dialog cannot be empty")

if dialog[0]["role"] == "system":
if len(dialog) < 2:
raise ValueError("System message must be followed by a user message")
dialog = [
{
"role": dialog[1]["role"],
"content": B_SYS + dialog[0]["content"] + E_SYS + dialog[1]["content"],
}
] + dialog[2:]

dialog_tokens = []
for prompt, answer in zip(dialog[::2], dialog[1::2]):
dialog_tokens.extend(
self.tokenizer.encode(
f"{B_INST} {prompt['content'].strip()} {E_INST} {answer['content'].strip()} ",
bos=True,
eos=True,
)
)

if dialog[-1]["role"] != "user":
raise ValueError(f"Last message must be from user, got {dialog[-1]['role']}")

dialog_tokens.extend(
self.tokenizer.encode(
f"{B_INST} {dialog[-1]['content'].strip()} {E_INST}",
bos=True,
eos=False,
)
)
return dialog_tokens


class Llama3Formatter(ChatFormatter):
def __init__(self, tokenizer: Tokenizer):
self.tokenizer = tokenizer

def format_dialog(self, dialog: Dialog) -> List[int]:
if not dialog:
raise ValueError("Dialog cannot be empty")

tokens = []
for i, msg in enumerate(dialog):
if "role" not in msg or "content" not in msg:
raise ValueError(f"Message {i} missing 'role' or 'content' field")

is_first = i == 0
formatted_msg = (
f"<|start_header_id|>{msg['role']}<|end_header_id|>\n\n"
f"{msg['content']}<|eot_id|>"
)
tokens.extend(self.tokenizer.encode(formatted_msg, bos=is_first, eos=False))

tokens.extend(
self.tokenizer.encode(
"<|start_header_id|>assistant<|end_header_id|>\n\n",
bos=False,
eos=False,
)
)
return tokens


class Llama:
@staticmethod
def build(
Expand Down Expand Up @@ -288,6 +366,7 @@ def chat_completion(
top_p: float = 0.9,
max_gen_len: Optional[int] = None,
logprobs: bool = False,
formatter_override: Optional[ChatFormatter] = None,
) -> List[ChatPrediction]:
"""
Generate assistant responses for a list of conversational dialogs using the language generation model.
Expand All @@ -299,66 +378,32 @@ def chat_completion(
max_gen_len (Optional[int], optional): Maximum length of the generated response sequence.
If not provided, it's set to the model's maximum sequence length minus 1.
logprobs (bool, optional): Flag indicating whether to compute token log probabilities. Defaults to False.
formatter_override (Optional[ChatFormatter]): Override the default formatter.

Returns:
List[ChatPrediction]: List of chat predictions, each containing the assistant's generated response.

Raises:
AssertionError: If the last message in a dialog is not from the user.
AssertionError: If the dialog roles are not in the required 'user', 'assistant', and optional 'system' order.

Note:
This method generates assistant responses for the provided conversational dialogs.
It employs nucleus sampling to introduce controlled randomness in text generation.
If logprobs is True, token log probabilities are computed for each generated token.

"""
if max_gen_len is None:
max_gen_len = self.model.params.max_seq_len - 1

if formatter_override:
formatter = formatter_override
elif hasattr(self.tokenizer, "tiktoken_model") and self.tokenizer.tiktoken_model is not None:
formatter = Llama3Formatter(self.tokenizer)
else:
formatter = Llama2Formatter(self.tokenizer)

prompt_tokens = []
unsafe_requests = []
for dialog in dialogs:
unsafe_requests.append(
any([tag in msg["content"] for tag in SPECIAL_TAGS for msg in dialog])
)
if dialog[0]["role"] == "system":
dialog = [
{
"role": dialog[1]["role"],
"content": B_SYS
+ dialog[0]["content"]
+ E_SYS
+ dialog[1]["content"],
}
] + dialog[2:]
assert all([msg["role"] == "user" for msg in dialog[::2]]) and all(
[msg["role"] == "assistant" for msg in dialog[1::2]]
), (
"model only supports 'system', 'user' and 'assistant' roles, "
"starting with 'system', then 'user' and alternating (u/a/u/a/u...)"
)
dialog_tokens: List[int] = sum(
[
self.tokenizer.encode(
f"{B_INST} {(prompt['content']).strip()} {E_INST} {(answer['content']).strip()} ",
bos=True,
eos=True,
)
for prompt, answer in zip(
dialog[::2],
dialog[1::2],
)
],
[],
)
assert (
dialog[-1]["role"] == "user"
), f"Last message must be from user, got {dialog[-1]['role']}"
dialog_tokens += self.tokenizer.encode(
f"{B_INST} {(dialog[-1]['content']).strip()} {E_INST}",
bos=True,
eos=False,
)
if isinstance(formatter, Llama2Formatter):
unsafe_requests.append(
any([tag in msg["content"] for tag in SPECIAL_TAGS for msg in dialog])
)
else:
unsafe_requests.append(False)

dialog_tokens = formatter.format_dialog(dialog)
prompt_tokens.append(dialog_tokens)

generation_tokens, generation_logprobs = self.generate(
Expand Down
75 changes: 73 additions & 2 deletions llama/model.py
Original file line number Diff line number Diff line change
Expand Up @@ -26,11 +26,84 @@ class ModelArgs:
multiple_of: int = 256 # make SwiGLU hidden layer size multiple of large power of 2
ffn_dim_multiplier: Optional[float] = None
norm_eps: float = 1e-5
rope_theta: float = 10000.0

max_batch_size: int = 32
max_seq_len: int = 2048


def precompute_freqs_cis(dim: int, end: int, theta: float = 10000.0):
"""
Precompute the frequency tensor for complex exponentials (cis) with given dimensions.

This function calculates a frequency tensor with complex exponentials using the given dimension 'dim'
and the end index 'end'. The 'theta' parameter scales the frequencies.
The returned tensor contains complex values in complex64 data type.

Args:
dim (int): Dimension of the frequency tensor.
end (int): End index for precomputing frequencies.
theta (float, optional): Scaling factor for frequency computation. Defaults to 10000.0.

Returns:
torch.Tensor: Precomputed frequency tensor with complex exponentials.




"""
freqs = 1.0 / (theta ** (torch.arange(0, dim, 2)[: (dim // 2)].float() / dim))
t = torch.arange(end, device=freqs.device) # type: ignore
freqs = torch.outer(t, freqs).float() # type: ignore
freqs_cis = torch.polar(torch.ones_like(freqs), freqs) # complex64
return freqs_cis



class Transformer(nn.Module):
def __init__(self, params: ModelArgs):
"""
Initialize a Transformer model.

Args:
params (ModelArgs): Model configuration parameters.

Attributes:
params (ModelArgs): Model configuration parameters.
vocab_size (int): Vocabulary size.
n_layers (int): Number of layers in the model.
tok_embeddings (ParallelEmbedding): Token embeddings.
layers (torch.nn.ModuleList): List of Transformer blocks.
norm (RMSNorm): Layer normalization for the model output.
output (ColumnParallelLinear): Linear layer for final output.
freqs_cis (torch.Tensor): Precomputed cosine and sine frequencies.

"""
super().__init__()
self.params = params
self.vocab_size = params.vocab_size
self.n_layers = params.n_layers

self.tok_embeddings = ParallelEmbedding(
params.vocab_size, params.dim, init_method=lambda x: x
)

self.layers = torch.nn.ModuleList()
for layer_id in range(params.n_layers):
self.layers.append(TransformerBlock(layer_id, params))

self.norm = RMSNorm(params.dim, eps=params.norm_eps)
self.output = ColumnParallelLinear(
params.dim, params.vocab_size, bias=False, init_method=lambda x: x
)

self.freqs_cis = precompute_freqs_cis(
self.params.dim // self.params.n_heads,
self.params.max_seq_len * 2,
theta=params.rope_theta,
)


class RMSNorm(torch.nn.Module):
def __init__(self, dim: int, eps: float = 1e-6):
"""
Expand Down Expand Up @@ -448,8 +521,6 @@ def __init__(self, params: ModelArgs):
)

self.freqs_cis = precompute_freqs_cis(
# Note that self.params.max_seq_len is multiplied by 2 because the token limit for the Llama 2 generation of models is 4096.
# Adding this multiplier instead of using 4096 directly allows for dynamism of token lengths while training or fine-tuning.
self.params.dim // self.params.n_heads, self.params.max_seq_len * 2
)

Expand Down
87 changes: 71 additions & 16 deletions llama/tokenizer.py
Original file line number Diff line number Diff line change
Expand Up @@ -3,6 +3,7 @@

import os
from logging import getLogger
from pathlib import Path
from typing import List

from sentencepiece import SentencePieceProcessor
Expand All @@ -12,28 +13,76 @@


class Tokenizer:
"""tokenizing and encoding/decoding text using SentencePiece."""
"""tokenizing and encoding/decoding text using SentencePiece or Tiktoken."""
def __init__(self, model_path: str):
"""
Initializes the Tokenizer with a SentencePiece model.
Initializes the Tokenizer with a SentencePiece model or Tiktoken model file.

Args:
model_path (str): The path to the SentencePiece model file.
model_path (str): The path to the SentencePiece model file or Tiktoken model file.
"""
# reload tokenizer
assert os.path.isfile(model_path), model_path
self.sp_model = SentencePieceProcessor(model_file=model_path)
logger.info(f"Reloaded SentencePiece model from {model_path}")
assert os.path.isfile(model_path), f"Tokenizer model not found: {model_path}"

self.sp_model = None
self.tiktoken_model = None

if model_path.endswith(".model"):
self.sp_model = SentencePieceProcessor(model_file=model_path)
self.n_words: int = self.sp_model.vocab_size()
self.bos_id: int = self.sp_model.bos_id()
self.eos_id: int = self.sp_model.eos_id()
self.pad_id: int = self.sp_model.pad_id()
logger.info(f"Loaded SentencePiece model from {model_path}")
else:
try:
import tiktoken
from tiktoken.load import load_tiktoken_bpe
except ImportError as e:
logger.error("Tiktoken not installed. Install with: pip install tiktoken")
raise ImportError(
"Tiktoken is required for Llama 3 models. Install with: pip install tiktoken"
) from e

mergeable_ranks = load_tiktoken_bpe(model_path)

num_base_tokens = len(mergeable_ranks)
special_tokens = [
"<|begin_of_text|>",
"<|end_of_text|>",
"<|reserved_special_token_0|>",
"<|reserved_special_token_1|>",
"<|finetune_right_pad_id|>",
"<|step_id|>",
"<|start_header_id|>",
"<|end_header_id|>",
"<|eom_id|>",
"<|eot_id|>",
"<|python_tag|>",
]
reserved_tokens = [
f"<|reserved_special_token_{2+i}|>"
for i in range(256 - len(special_tokens))
]
special_tokens.extend(reserved_tokens)

self.tiktoken_model = tiktoken.Encoding(
name=Path(model_path).name,
pat_str=r"(?i:'s|'t|'re|'ve|'m|'ll|'d)|[^\r\n\p{L}\p{N}]?\p{L}+|\p{N}{1,3}| ?[^\s\p{L}\p{N}]+[\r\n]*|\s*[\r\n]+|\s+(?!\S)|\s+",
mergeable_ranks=mergeable_ranks,
special_tokens={
token: num_base_tokens + i for i, token in enumerate(special_tokens)
},
)
self.n_words = self.tiktoken_model.n_vocab
self.bos_id = self.tiktoken_model.encode_single_token("<|begin_of_text|>")
self.eos_id = self.tiktoken_model.encode_single_token("<|end_of_text|>")
self.pad_id = self.tiktoken_model.encode_single_token("<|finetune_right_pad_id|>")

logger.info(f"Loaded Tiktoken model from {model_path}")

# BOS / EOS token IDs
self.n_words: int = self.sp_model.vocab_size()
self.bos_id: int = self.sp_model.bos_id()
self.eos_id: int = self.sp_model.eos_id()
self.pad_id: int = self.sp_model.pad_id()
logger.info(
f"#words: {self.n_words} - BOS ID: {self.bos_id} - EOS ID: {self.eos_id}"
)
assert self.sp_model.vocab_size() == self.sp_model.get_piece_size()

def encode(self, s: str, bos: bool, eos: bool) -> List[int]:
"""
Expand All @@ -47,8 +96,11 @@ def encode(self, s: str, bos: bool, eos: bool) -> List[int]:
Returns:
List[int]: A list of token IDs.
"""
assert type(s) is str
t = self.sp_model.encode(s)
if not isinstance(s, str):
raise TypeError(f"Expected str, got {type(s)}")

t = self.sp_model.encode(s) if self.sp_model else self.tiktoken_model.encode(s)

if bos:
t = [self.bos_id] + t
if eos:
Expand All @@ -65,4 +117,7 @@ def decode(self, t: List[int]) -> str:
Returns:
str: The decoded string.
"""
return self.sp_model.decode(t)
if self.sp_model:
return self.sp_model.decode(t)
else:
return self.tiktoken_model.decode(t)
1 change: 1 addition & 0 deletions requirements.txt
Original file line number Diff line number Diff line change
Expand Up @@ -2,3 +2,4 @@ torch
fairscale
fire
sentencepiece
tiktoken
Loading