luna-1.5-flash / modeling_lightning.py
Aobangaming's picture
Update modeling_lightning.py
c55a851 verified
Raw History Blame Contribute Delete
16.4 kB
import math
import re
import torch
import torch.nn as nn
import torch.nn.functional as F
from transformers import PreTrainedModel, PretrainedConfig
from transformers.modeling_outputs import CausalLMOutput
# ============================================================
# Sampling
# ============================================================
def top_k_top_p_sample(
logits,
top_k=40,
top_p=0.9
):
"""
Sample one token from logits using top-k and/or top-p sampling.
"""
logits = logits.float()
# --------------------------------------------------------
# Top-k
# --------------------------------------------------------
if top_k is not None and top_k > 0:
top_k = min(
top_k,
logits.size(-1)
)
values, indices = torch.topk(
logits,
top_k
)
filtered_logits = torch.full_like(
logits,
-float("inf")
)
filtered_logits.scatter_(
0,
indices,
values
)
logits = filtered_logits
# --------------------------------------------------------
# Top-p
# --------------------------------------------------------
if top_p is not None and 0.0 < top_p < 1.0:
sorted_logits, sorted_indices = torch.sort(
logits,
descending=True
)
probabilities = torch.softmax(
sorted_logits,
dim=-1
)
cumulative_probabilities = torch.cumsum(
probabilities,
dim=-1
)
remove_mask = (
cumulative_probabilities > top_p
)
# Always keep the first token above the threshold.
remove_mask[1:] = remove_mask[:-1].clone()
remove_mask[0] = False
sorted_logits[remove_mask] = -float("inf")
logits = torch.full_like(
logits,
-float("inf")
)
logits.scatter_(
0,
sorted_indices,
sorted_logits
)
probabilities = torch.softmax(
logits,
dim=-1
)
next_token = torch.multinomial(
probabilities,
num_samples=1
)
return next_token.item()
# ============================================================
# Positional Encoding
# ============================================================
class PositionalEncoding(nn.Module):
def __init__(self, d_model, max_len=5000):
super().__init__()
pe = torch.zeros(
max_len,
d_model
)
position = torch.arange(
0,
max_len,
dtype=torch.float
).unsqueeze(1)
div_term = torch.exp(
torch.arange(
0,
d_model,
2
).float()
* (
-math.log(10000.0)
/ d_model
)
)
pe[:, 0::2] = torch.sin(
position * div_term
)
pe[:, 1::2] = torch.cos(
position * div_term
)
pe = pe.unsqueeze(0).transpose(0, 1)
self.register_buffer(
"pe",
pe
)
def forward(self, x):
return x + self.pe[
:x.size(1), :
].transpose(0, 1)
# ============================================================
# Causal Self Attention
# ============================================================
class CausalSelfAttention(nn.Module):
def __init__(
self,
d_model,
nhead,
dropout=0.1
):
super().__init__()
assert d_model % nhead == 0
self.nhead = nhead
self.head_dim = (
d_model // nhead
)
self.dropout = dropout
self.qkv = nn.Linear(
d_model,
d_model * 3
)
self.out_proj = nn.Linear(
d_model,
d_model
)
def forward(self, x):
B, T, C = x.shape
q, k, v = self.qkv(x).chunk(
3,
dim=-1
)
q = q.view(
B,
T,
self.nhead,
self.head_dim
).transpose(1, 2)
k = k.view(
B,
T,
self.nhead,
self.head_dim
).transpose(1, 2)
v = v.view(
B,
T,
self.nhead,
self.head_dim
).transpose(1, 2)
y = F.scaled_dot_product_attention(
q,
k,
v,
attn_mask=None,
dropout_p=(
self.dropout
if self.training
else 0.0
),
is_causal=True
)
y = (
y.transpose(1, 2)
.contiguous()
.view(B, T, C)
)
return self.out_proj(y)
# ============================================================
# Transformer Block
# ============================================================
class TransformerBlock(nn.Module):
def __init__(
self,
d_model,
nhead,
dropout=0.1
):
super().__init__()
self.norm1 = nn.LayerNorm(
d_model
)
self.attention = CausalSelfAttention(
d_model,
nhead,
dropout
)
self.norm2 = nn.LayerNorm(
d_model
)
self.ffn = nn.Sequential(
nn.Linear(
d_model,
d_model * 4
),
nn.GELU(),
nn.Linear(
d_model * 4,
d_model
),
nn.Dropout(
dropout
)
)
def forward(self, x):
x = x + self.attention(
self.norm1(x)
)
x = x + self.ffn(
self.norm2(x)
)
return x
# ============================================================
# Transformer Language Model
# ============================================================
class TransformerLanguageModel(nn.Module):
def __init__(
self,
vocab_size,
d_model=256,
nhead=4,
num_layers=6,
dropout=0.1,
max_seq_len=200
):
super().__init__()
self.d_model = d_model
self.max_seq_len = max_seq_len
self.token_embedding = nn.Embedding(
vocab_size,
d_model
)
self.positional_encoding = PositionalEncoding(
d_model,
max_seq_len
)
self.transformer = nn.ModuleList([
TransformerBlock(
d_model,
nhead,
dropout
)
for _ in range(num_layers)
])
self.final_norm = nn.LayerNorm(
d_model
)
self.output_layer = nn.Linear(
d_model,
vocab_size,
bias=False
)
def forward(self, src):
x = self.token_embedding(src)
x = self.positional_encoding(x)
for layer in self.transformer:
x = layer(x)
x = self.final_norm(x)
return self.output_layer(x)
# ============================================================
# Lightning Config
# ============================================================
class LightningConfig(
PretrainedConfig
):
model_type = "lightning"
def __init__(
self,
vocab_size=75000,
d_model=256,
nhead=4,
num_layers=6,
dropout=0.1,
max_seq_len=200,
**kwargs
):
kwargs.setdefault(
"tie_word_embeddings",
False
)
super().__init__(
**kwargs
)
self.vocab_size = vocab_size
self.d_model = d_model
self.nhead = nhead
self.num_layers = num_layers
self.dropout = dropout
self.max_seq_len = max_seq_len
self.num_hidden_layers = (
num_layers
)
self.num_attention_heads = (
nhead
)
self.hidden_size = (
d_model
)
# ============================================================
# Lightning Causal LM
# ============================================================
class LightningForCausalLM(
PreTrainedModel
):
config_class = LightningConfig
base_model_prefix = "lightning"
def __init__(
self,
config
):
super().__init__(
config
)
self.lightning = (
TransformerLanguageModel(
vocab_size=config.vocab_size,
d_model=config.d_model,
nhead=config.nhead,
num_layers=config.num_layers,
dropout=config.dropout,
max_seq_len=config.max_seq_len
)
)
self.post_init()
# --------------------------------------------------------
# Forward
# --------------------------------------------------------
def forward(
self,
input_ids=None,
labels=None,
**kwargs
):
logits = self.lightning(
input_ids
)
loss = None
if labels is not None:
shift_logits = (
logits[..., :-1, :]
.contiguous()
)
shift_labels = (
labels[..., 1:]
.contiguous()
)
loss_fn = nn.CrossEntropyLoss()
loss = loss_fn(
shift_logits.view(
-1,
shift_logits.size(-1)
),
shift_labels.view(-1)
)
return CausalLMOutput(
loss=loss,
logits=logits
)
# --------------------------------------------------------
# Embeddings
# --------------------------------------------------------
def get_input_embeddings(
self
):
return self.lightning.token_embedding
def set_input_embeddings(
self,
value
):
self.lightning.token_embedding = value
def get_output_embeddings(
self
):
return self.lightning.output_layer
def set_output_embeddings(
self,
new_embeddings
):
self.lightning.output_layer = (
new_embeddings
)
# ============================================================
# Text Generation API
# ============================================================
def generate_text(
model,
tokenizer,
prompt,
max_len,
device,
top_k=40,
top_p=0.9,
penalty=1.2,
temperature=0.8,
chat_history=None
):
"""
Generate text from Lightning.
chat_history format:
[
{
"role": "user",
"content": "Hello"
},
{
"role": "assistant",
"content": "Hi!"
}
]
"""
model.eval()
# --------------------------------------------------------
# Maximum sequence length
# --------------------------------------------------------
msl = getattr(
model.config,
"max_seq_len",
200
)
# --------------------------------------------------------
# Build conversation
# --------------------------------------------------------
messages = []
if chat_history:
for message in chat_history:
role = message.get(
"role",
""
).lower()
content = message.get(
"content",
""
).strip()
if not content:
continue
if role == "user":
messages.append(
f"User: {content}"
)
elif role == "assistant":
messages.append(
f"Assistant: {content}"
)
messages.append(
f"User: {prompt}"
)
messages.append(
"Assistant:"
)
generation_prompt = "\n".join(
messages
)
# --------------------------------------------------------
# Tokenize
# --------------------------------------------------------
encoding = tokenizer.encode(
generation_prompt,
add_special_tokens=False
)
input_ids = torch.tensor(
[encoding.ids],
dtype=torch.long,
device=device
)
# --------------------------------------------------------
# Context window
# --------------------------------------------------------
if input_ids.size(1) > msl:
input_ids = input_ids[
:,
-msl:
]
prompt_len = input_ids.size(1)
generated = (
input_ids[0].tolist()
)
# --------------------------------------------------------
# Special tokens
# --------------------------------------------------------
eos_id = tokenizer.token_to_id(
"<|endoftext|>"
)
eor_id = tokenizer.token_to_id(
"<|eor|>"
)
pad_id = getattr(
tokenizer,
"pad_id",
None
)
# --------------------------------------------------------
# Generation
# --------------------------------------------------------
for _ in range(max_len):
src = input_ids[
:,
-msl:
]
with torch.no_grad():
output = model(
src
)
logits = (
output.logits[:, -1, :]
.squeeze(0)
)
# ----------------------------------------------------
# Prevent PAD generation
# ----------------------------------------------------
if pad_id is not None:
logits[
pad_id
] = -float("inf")
# ----------------------------------------------------
# Repetition penalty
# ----------------------------------------------------
response_tokens = (
generated[prompt_len:]
)
for idx in set(
response_tokens[-32:]
):
if logits[idx] > 0:
logits[idx] /= penalty
else:
logits[idx] *= penalty
# ----------------------------------------------------
# Temperature
# ----------------------------------------------------
if temperature <= 0:
raise ValueError(
"temperature must be > 0"
)
logits /= temperature
# ----------------------------------------------------
# Sample
# ----------------------------------------------------
next_token = top_k_top_p_sample(
logits,
top_k=top_k,
top_p=top_p
)
# ----------------------------------------------------
# Stop tokens
# ----------------------------------------------------
if (
next_token == eor_id
or next_token == eos_id
):
break
generated.append(
next_token
)
input_ids = torch.cat(
[
input_ids,
torch.tensor(
[[next_token]],
device=device
)
],
dim=1
)
# --------------------------------------------------------
# Decode
# --------------------------------------------------------
new_tokens = generated[
prompt_len:
]
response = tokenizer.decode(
new_tokens,
skip_special_tokens=True
)
response = (
response
.replace("<pad>", "")
.strip()
)
# --------------------------------------------------------
# Cleanup
# --------------------------------------------------------
response = re.sub(
r"[{}\\/]",
"",
response
)
if response.startswith(
"Assistant:"
):
response = (
response[
len("Assistant:"):
]
.strip()
)
return response