ThoughtVectors / ThoughtVectors.py
noch inator
Fixed more bugs, solved an issue with 16FP used for the GRU.
e8bfacb verified
Raw
History Blame Contribute Delete
43.7 kB
import atexit
import random
import math
import torch
import torch.nn as nn
import torch.optim as optim
import sentencepiece as spm
from torch.amp import autocast, GradScaler
from typing import List, Union, Generator, Sequence
import os
import csv
import tempfile
import tarfile
import shutil
class PositionalEncoding(nn.Module):
"""Adds positional information to token embeddings using sine and cosine functions.
This module can either precompute positional encodings up to a specified `max_len` or compute
them dynamically based on the input sequence length. If `max_len` is 0, encodings are computed
on-the-fly in the forward pass; otherwise, they are precomputed during initialization.
Args:
d_model (int): Dimensionality of the model embeddings.
dropout (float, optional): Dropout probability applied after adding encodings. Defaults to 0.1.
max_len (int, optional): Maximum sequence length for precomputed encodings; if 0, computes
dynamically. Defaults to 5000.
Attributes:
dropout (nn.Dropout): Dropout layer for regularization.
pe (torch.Tensor, optional): Precomputed positional encodings, shape (1, max_len, d_model),
present only if max_len > 0.
"""
def __init__(self, d_model: int, dropout: float = 0.1, max_len: int = 5000):
super().__init__()
self.d_model = d_model
self.dropout = nn.Dropout(dropout)
self.max_len = max_len
if max_len > 0:
# Precompute positional encodings
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 - (d_model % 2), 2, dtype=torch.float) *
(-math.log(10000.0) / d_model))
pe[:, 0::2] = torch.sin(position * div_term)
pe[:, 1::2] = torch.cos(position * div_term[:self.d_model // 2 + (self.d_model % 2)])
pe = pe.unsqueeze(0)
self.register_buffer('pe', pe)
else:
self.pe = None
def forward(self, x: torch.Tensor) -> torch.Tensor:
"""Applies positional encodings to the input tensor.
Args:
x (torch.Tensor): Input embeddings, shape (batch_size, seq_len, d_model).
Returns:
torch.Tensor: Embeddings with positional encodings added and dropout applied.
Raises:
ValueError: If max_len > 0 and seq_len exceeds max_len.
"""
seq_len = x.size(1)
if self.max_len > 0:
if seq_len > self.max_len:
raise ValueError(f"Sequence length {seq_len} exceeds max_len {self.max_len}")
pe = self.pe[:, :seq_len]
else:
pe = torch.zeros(seq_len, self.d_model, device=x.device)
position = torch.arange(0, seq_len, dtype=torch.float, device=x.device).unsqueeze(1)
div_term = torch.exp(torch.arange(0, self.d_model - (self.d_model % 2), 2, dtype=torch.float,
device=x.device) * (-math.log(10000.0) / self.d_model))
pe[:, 0::2] = torch.sin(position * div_term)
pe[:, 1::2] = torch.cos(position * div_term[:self.d_model // 2 + 1] if self.d_model % 2 else div_term)
pe = pe.unsqueeze(0)
x = x + pe
return self.dropout(x)
class ThoughtEncoder(nn.Module):
"""Encodes text into thought vectors, representing abstract ideas or summaries.
Args:
vocab_size (int): Size of the vocabulary for token embeddings.
d_model (int, optional): Dimensionality of the model embeddings. Defaults to 256.
max_thoughts (int, optional): Maximum number of thought vectors to generate. Defaults to 16.
nhead (int, optional): Number of attention heads in the transformer encoder. Defaults to 8.
num_layers (int, optional): Number of transformer encoder layers. Defaults to 2.
dropout (float, optional): Dropout probability. Defaults to 0.1.
pad_id (int, optional): ID for padding tokens. Defaults to 0.
termination_threshold (float, optional): Threshold for stopping thought vector generation.
Defaults to 0.75.
"""
def __init__(self, vocab_size: int, d_model: int = 256, max_thoughts: int = 16,
nhead: int = 8, num_layers: int = 2, dropout: float = 0.1, pad_id: int = 0,
termination_threshold: float = 0.75, max_len: int = 5000):
super().__init__()
if nhead % 2 != 0:
raise ValueError(f"nhead must be even for thought_attention, got {nhead}")
self.d_model = d_model
self.max_thoughts = max_thoughts
self.pad_id = pad_id
self.termination_threshold = termination_threshold
self.embedding = nn.Embedding(vocab_size, d_model, padding_idx=pad_id)
self.positional_encoding = PositionalEncoding(d_model, dropout, max_len=max_len)
encoder_layer = nn.TransformerEncoderLayer(d_model=d_model, nhead=nhead, dropout=dropout, batch_first=True)
self.text_encoder = nn.TransformerEncoder(encoder_layer, num_layers=num_layers)
self.thought_attention = nn.MultiheadAttention(d_model, num_heads=nhead // 2, dropout=dropout, batch_first=True)
self.thought_rnn = nn.GRU(d_model, d_model, batch_first=True)
self.fc_thought = nn.Linear(2 * d_model, d_model)
self.fc_terminate = nn.Linear(d_model, 1)
def forward(self, input_tokens: torch.Tensor, force_single_vector: bool = False) -> torch.Tensor:
"""Encodes input tokens into a sequence of thought vectors.
Args:
input_tokens (torch.Tensor): Token IDs, shape (batch_size, seq_len).
force_single_vector (bool, optional): If True, generates only one thought vector.
Defaults to False.
Returns:
torch.Tensor: Thought vectors, shape (batch_size, num_thoughts, d_model).
"""
batch_size, seq_len = input_tokens.size()
max_thoughts = 1 if force_single_vector else self.max_thoughts
src_key_padding_mask = (input_tokens == self.pad_id)
x = self.embedding(input_tokens)
x = self.positional_encoding(x)
encoded_text = self.text_encoder(x, src_key_padding_mask=src_key_padding_mask)
valid_mask = (~src_key_padding_mask).unsqueeze(-1).float()
text_context = (encoded_text * valid_mask).sum(dim=1) / valid_mask.sum(dim=1).clamp(min=1.0)
thought_vectors_list = []
thought_hidden = text_context.unsqueeze(0)
finished = torch.zeros(batch_size, dtype=torch.bool, device=input_tokens.device)
for t in range(max_thoughts):
if t > 0:
prev_thoughts = torch.stack(thought_vectors_list, dim=1)
attn_output, _ = self.thought_attention(prev_thoughts, prev_thoughts, prev_thoughts)
prev_context = attn_output.mean(dim=1)
else:
prev_context = torch.zeros_like(text_context)
combined = torch.cat([text_context, prev_context], dim=1)
next_thought = self.fc_thought(combined)
with autocast('cuda', enabled=False): # FP32
thought_hidden = self.thought_rnn(next_thought.unsqueeze(1).float(), thought_hidden.float())[1]
termination_logit = self.fc_terminate(thought_hidden.squeeze(0))
termination_score = torch.sigmoid(termination_logit).squeeze(-1)
thought_vectors_list.append(next_thought)
finished = finished | (termination_score > self.termination_threshold)
if finished.all() and not force_single_vector:
break
return torch.stack(thought_vectors_list, dim=1)
class ThoughtDecoder(nn.Module):
"""Decodes thought vectors back into token sequences using a transformer decoder.
Args:
vocab_size (int): Size of the vocabulary for token embeddings.
d_model (int, optional): Dimensionality of the model embeddings. Defaults to 256.
num_layers (int, optional): Number of transformer decoder layers. Defaults to 2.
nhead (int, optional): Number of attention heads in the transformer decoder. Defaults to 8.
dropout (float, optional): Dropout probability. Defaults to 0.1.
"""
def __init__(self, vocab_size: int, d_model: int = 256, num_layers: int = 2,
nhead: int = 8, dropout: float = 0.1, max_len: int = 5000):
super().__init__()
self.vocab_size = vocab_size
self.d_model = d_model
self.embedding = nn.Embedding(vocab_size, d_model, padding_idx=0)
self.positional_encoding = PositionalEncoding(d_model, dropout, max_len)
decoder_layer = nn.TransformerDecoderLayer(d_model=d_model, nhead=nhead, dropout=dropout, batch_first=True)
self.transformer_decoder = nn.TransformerDecoder(decoder_layer, num_layers=num_layers)
self.fc_out = nn.Linear(d_model, vocab_size)
def generate_square_subsequent_mask(self, sz: int, device: torch.device) -> torch.Tensor:
"""Generates a causal mask for autoregressive decoding.
Args:
sz (int): Size of the mask (sequence length).
device (torch.device): Device to create the mask on.
Returns:
torch.Tensor: Mask tensor, shape (sz, sz).
"""
if self.positional_encoding.max_len > 0 and sz > self.positional_encoding.max_len:
raise ValueError(f"Sequence length {sz} exceeds max_len {self.positional_encoding.max_len}")
mask = (torch.triu(torch.ones(sz, sz, device=device)) == 1).transpose(0, 1)
mask = mask.float().masked_fill(mask == 0, float('-inf')).masked_fill(mask == 1, float(0.0))
return mask
def forward(self, thought_vectors: torch.Tensor, target_tokens: torch.Tensor) -> torch.Tensor:
"""Decodes thought vectors into token logits.
Args:
thought_vectors (torch.Tensor): Thought vectors, shape (batch_size, num_thoughts, d_model).
target_tokens (torch.Tensor): Target token IDs, shape (batch_size, tgt_seq_len).
Returns:
torch.Tensor: Logits over vocabulary, shape (batch_size, tgt_seq_len, vocab_size).
"""
tgt_emb = self.embedding(target_tokens)
tgt_emb = self.positional_encoding(tgt_emb)
tgt_mask = self.generate_square_subsequent_mask(tgt_emb.size(1), tgt_emb.device)
output = self.transformer_decoder(tgt=tgt_emb, memory=thought_vectors, tgt_mask=tgt_mask)
return self.fc_out(output)
def group_collate_fn(batch: List[List[List[int]]]) -> List[torch.Tensor]:
"""Collates a batch of tokenized groups into padded tensors.
Args:
batch (List[List[List[int]]]): Batch of tokenized groups.
Returns:
List[torch.Tensor]: Padded tensors for each group, shape (num_sentences, max_seq_len).
"""
collated_groups = []
for group in batch:
if not group:
collated_groups.append(torch.tensor([]))
continue
group_tensors = [torch.tensor(seq, dtype=torch.long) for seq in group]
padded_group = nn.utils.rnn.pad_sequence(group_tensors, batch_first=True, padding_value=0)
collated_groups.append(padded_group)
return collated_groups
class LazyList(Sequence):
def __init__(self, data: Union[str, List[List[str]]], sample_prob: float = 0.75):
"""Initialize LazyList with data source and sampling probability.
Args:
data (Union[str, List[List[str]]]): A filepath to a CSV or a list of sentence groups.
sample_prob (float, optional): Probability of returning a row in __getitem__ (0.0 to 1.0).
Defaults to 0.75 (75% chance to return, 25% to skip).
"""
if isinstance(data, str):
self.filepath = data
self.data_list = None
elif isinstance(data, list):
self.filepath = None
self.data_list = data
else:
raise ValueError("Data must be a filepath (str) or a list of lists of strings")
self.sample_prob = max(0.0, min(1.0, sample_prob)) # Clamp between 0 and 1
self._length = None
self._file_handle = None # Persistent file handle for sequential reading
self._reader = None # Persistent CSV reader
def _open_file(self):
"""Open the file and initialize the reader if not already open."""
if self.filepath and os.path.exists(self.filepath) and self._file_handle is None:
self._file_handle = open(self.filepath, 'r', encoding='utf-8', newline='')
self._reader = csv.reader(self._file_handle)
def __iter__(self) -> Generator[List[str], None, None]:
"""Iterate through all rows sequentially without sampling."""
if self.data_list is not None:
yield from self.data_list
elif self.filepath and os.path.exists(self.filepath):
with open(self.filepath, 'r', encoding='utf-8', newline='') as f:
reader = csv.reader(f)
for row in reader:
if row:
yield row
else:
dummy_data = [
["Hello world", "AI is cool"],
["This is a test", "Another sentence"],
["Python is fun", "Coding rocks"],
["Short sentence", "Quick test"],
]
yield from dummy_data
def __len__(self) -> int:
"""Compute and cache the total number of rows."""
if self._length is None:
if self.data_list is not None:
self._length = len(self.data_list)
elif self.filepath and os.path.exists(self.filepath):
with open(self.filepath, 'r', encoding='utf-8', newline='') as file:
reader = csv.reader(file)
self._length = sum(1 for row in reader if row)
else:
self._length = 4
return self._length
def __getitem__(self, index: Union[int, slice]) -> Union[List[str], List[List[str]]]:
if isinstance(index, slice):
start, stop, step = index.indices(len(self))
count = (stop - start) // (step or 1)
return [self._get_single_item() for _ in range(max(0, count))]
return self._get_single_item()
def _get_single_item(self) -> List[str]:
if self.data_list is not None:
return random.choice(self.data_list)
elif self.filepath and os.path.exists(self.filepath):
self._open_file()
try:
while True:
row = next(self._reader)
if row and random.random() <= self.sample_prob:
return row
except StopIteration:
self._file_handle.seek(0)
self._reader = csv.reader(self._file_handle)
row = next(self._reader)
return row if row else random.choice(self._dummy_data())
print("WARN: file error, file may be empty, missing, or inaccessible. Dummy data returned") # change to raise in release
return random.choice(self._dummy_data()) # remove in release
# remove in release
def _dummy_data(self) -> List[List[str]]:
return [
["Hello world", "AI is cool"],
["This is a test", "Another sentence"],
["Python is fun", "Coding rocks"],
["Short sentence", "Quick test"],
]
def __del__(self):
"""Close the file handle when the object is destroyed."""
if self._file_handle is not None:
self._file_handle.close()
self._file_handle = None
self._reader = None
class ThoughtVectors:
"""Main class for a thought vector model, encoding text into latent vectors and decoding them back.
Attributes:
device (torch.device): Computation device (CUDA if available, else CPU).
encoder (ThoughtEncoder): Encoder module.
decoder (ThoughtDecoder): Decoder module.
sp (spm.SentencePieceProcessor): SentencePiece processor.
"""
def __init__(self):
self.device = torch.device("cuda" if torch.cuda.is_available() else "cpu")
self.encoder = None
self.decoder = None
self.vocab_size = None
self.sp = None
self.d_model = None
self.max_thoughts = None
self.encoder_nhead = None
self.decoder_nhead = None
self.encoder_layers = None
self.decoder_layers = None
self.dropout = None
self.max_len = None
self.termination_threshold = None
# Track files to clean up
self._temp_files = set()
# Register cleanup function
atexit.register(self._cleanup)
def train(self,
group_data: Union[str, List[List[str]]],
test_data: Union[str, List[List[str]]] = None,
model: str = None,
num_epochs: int = 2048,
batch_size: int = 128,
batches_per_val: int = 16,
val_batches: int = 8,
accum_steps: int = 4,
learning_rate: float = 2e-4,
weight_decay: float = 1e-5,
length_penalty: float = 0.01,
single_vector_prob: float = 0.1,
save_path: str = "thought_vectors.tar",
spm_model_prefix: str = "spm",
vocab_size: int = 4096,
d_model: int = 256,
encoder_nhead: int = 8,
decoder_nhead: int = 8,
encoder_layers: int = 2,
decoder_layers: int = 2,
max_thoughts: int = 8,
dropout: float = 0.1,
max_len: int = 1024,
termination_threshold: float = 0.75,
patience: int = 5) -> None:
"""Trains the thought vector model on grouped text data.
Args:
group_data (Union[str, List[List[str]]]): Filepath to a CSV or list of sentence groups for training.
test_data (Union[str, List[List[str]]], optional): Filepath to a CSV or list of sentence groups for validation. If None, no validation is performed. Defaults to None.
model (str, optional): Path to pre-existing model to load. If None, creates a new model. Defaults to None.
num_epochs (int, optional): Number of training epochs. Defaults to 30.
batch_size (int, optional): Batch size. Defaults to 8.
batches_per_val (int, optional): Validate model every x batches. Defaults to 16
val_batches (int, optional): Number of batches to use in validation, speeds up validation. Defaults to 8
accum_steps (int, optional): Number of gradient accumulation steps. Defaults to 4.
learning_rate (float, optional): Learning rate for Adam optimizer. Defaults to 2e-4.
weight_decay (float, optional): Weight decay for regularization. Defaults to 1e-5.
length_penalty (float, optional): Penalty per additional thought vector. Defaults to 0.01.
single_vector_prob (float, optional): Probability of forcing a single thought vector. Defaults to 0.1.
save_path (str, optional): Path to save the model tar file. Defaults to "thought_vectors.tar".
spm_model_prefix (str, optional): Prefix for SentencePiece files. Defaults to "spm".
vocab_size (int, optional): Vocabulary size for tokenization. Defaults to 4096.
d_model (int, optional): Embedding dimensionality. Defaults to 256.
encoder_nhead (int, optional): Number of encoder attention heads. Defaults to 8.
decoder_nhead (int, optional): Number of decoder attention heads. Defaults to 8.
encoder_layers (int, optional): Number of encoder layers. Defaults to 2.
decoder_layers (int, optional): Number of decoder layers. Defaults to 2.
max_thoughts (int, optional): Maximum number of thought vectors. Defaults to 8.
dropout (float, optional): Dropout probability. Defaults to 0.1.
max_len (int, optional): Maximum sequence length for positional encoding. Defaults to 1024.
termination_threshold (float, optional): Threshold for stopping thought generation. Defaults to 0.75.
patience (int, optional): BATCHES to wait for improvement before early stopping. Defaults to 5.
"""
self.d_model = d_model
self.max_thoughts = max_thoughts
self.encoder_nhead = encoder_nhead
self.decoder_nhead = decoder_nhead
self.encoder_layers = encoder_layers
self.decoder_layers = decoder_layers
self.dropout = dropout
self.max_len = max_len
self.termination_threshold = termination_threshold
self.device = torch.device("cuda" if torch.cuda.is_available() else "cpu")
self.sp = spm.SentencePieceProcessor()
train_lazy_data = LazyList(group_data)
test_lazy_data = LazyList(test_data) if test_data is not None else None
if model and os.path.exists(model):
print("loading model...")
loaded_instance = self.load(model)
self.__dict__.update(loaded_instance.__dict__)
print(f"Loaded pre-existing model from {model}")
else:
print("building models")
with tempfile.NamedTemporaryFile(mode="w", encoding="utf8", delete=False) as temp_file:
for group in train_lazy_data:
for sentence in group:
temp_file.write(sentence.strip() + "\n")
temp_dataset_path = temp_file.name
# Use a persistent directory instead of temp_dir to avoid premature cleanup
os.makedirs("temp_spm_dir", exist_ok=True)
temp_spm_prefix = os.path.join("temp_spm_dir", "spm")
try:
print("training sentence piece (could take a while on massive datasets)")
spm.SentencePieceTrainer.Train(
input=temp_dataset_path,
model_prefix=temp_spm_prefix,
vocab_size=vocab_size,
model_type="bpe",
pad_id=0, unk_id=1, bos_id=2, eos_id=3,
max_sentence_length=self.max_len,
input_sentence_size=8_388_608,
train_extremely_large_corpus=True
)
# Ensure files are copied to the final location
shutil.copy(f"{temp_spm_prefix}.model", f"{spm_model_prefix}.model")
shutil.copy(f"{temp_spm_prefix}.vocab", f"{spm_model_prefix}.vocab")
if not self.sp.Load(f"{spm_model_prefix}.model"):
raise RuntimeError("Failed to load SentencePiece model.")
finally:
os.remove(temp_dataset_path)
# Only remove temp_spm_dir after successful save, handled in save method
self.vocab_size = self.sp.GetPieceSize()
self.encoder = ThoughtEncoder(
vocab_size=self.vocab_size, d_model=d_model, max_thoughts=max_thoughts,
nhead=encoder_nhead, num_layers=encoder_layers, dropout=dropout,
termination_threshold=termination_threshold, max_len=self.max_len
).to(self.device)
self.decoder = ThoughtDecoder(
vocab_size=self.vocab_size, d_model=d_model, num_layers=decoder_layers,
nhead=decoder_nhead, dropout=dropout, max_len=self.max_len
).to(self.device)
self.vocab_size = self.sp.GetPieceSize()
self.encoder = ThoughtEncoder(
vocab_size=self.vocab_size, d_model=d_model, max_thoughts=max_thoughts,
nhead=encoder_nhead, num_layers=encoder_layers, dropout=dropout,
termination_threshold=termination_threshold, max_len=self.max_len
).to(self.device)
self.decoder = ThoughtDecoder(
vocab_size=self.vocab_size, d_model=d_model, num_layers=decoder_layers,
nhead=decoder_nhead, dropout=dropout, max_len=self.max_len
).to(self.device)
optimizer = optim.Adam(list(self.encoder.parameters()) + list(self.decoder.parameters()),
lr=learning_rate, weight_decay=weight_decay)
scheduler = optim.lr_scheduler.CosineAnnealingLR(optimizer, T_max=num_epochs)
criterion = nn.CrossEntropyLoss(ignore_index=self.sp.pad_id())
scaler = GradScaler('cuda')
best_val_loss = float('inf')
patience_counter = 0
best_batch = 0 # Track best batch number
print("Beginning training")
try:
for epoch in range(num_epochs):
total_train_loss = 0.0
num_train_batches = (len(train_lazy_data) + batch_size - 1) // batch_size
for i in range(0, len(train_lazy_data), batch_size):
batch_idx = i // batch_size + 1
raw_batch = train_lazy_data[i:i + batch_size]
tokenized_batch = [
[[self.sp.bos_id()] + self.sp.EncodeAsIds(s) + [self.sp.eos_id()] for s in group]
for group in raw_batch
]
group_batch = group_collate_fn(tokenized_batch)
batch_loss = 0.0
with autocast('cuda'):
for group_tensor in group_batch:
if group_tensor.shape[1] <= self.max_len and group_tensor.numel() > 0:
group_tensor = group_tensor.to(self.device)
force_single_vector = random.random() < single_vector_prob
thought_vectors = self.encoder(group_tensor, force_single_vector)
output_logits = self.decoder(thought_vectors, group_tensor[:, :-1])
mask = (group_tensor[:, 1:] != self.sp.pad_id()).float()
loss = criterion(output_logits.reshape(-1, self.vocab_size),
group_tensor[:, 1:].reshape(-1))
loss = (loss * mask.reshape(-1)).sum() / mask.sum().clamp(min=1.0)
if torch.isnan(loss):
print(f"NaN detected: mask_sum={mask.sum().item()}, loss_pre_mask={loss.item()}")
print(f"group_tensor={group_tensor}\n\n")
print(f"thought vectors={thought_vectors}\n\n")
print("Raising keyboard interrupt to allow preservation.")
raise KeyboardInterrupt
if not force_single_vector:
loss += length_penalty * thought_vectors.shape[1]
batch_loss += loss / accum_steps
scaler.scale(batch_loss).backward()
total_train_loss += batch_loss.item() * accum_steps
if (batch_idx % accum_steps) == 0:
scaler.unscale_(optimizer)
torch.nn.utils.clip_grad_norm_(
list(self.encoder.parameters()) + list(self.decoder.parameters()), max_norm=0.5)
scaler.step(optimizer)
scaler.update()
optimizer.zero_grad()
torch.cuda.empty_cache()
# Per group of batches validation
if (batch_idx + 1) % batches_per_val == 0 or batch_idx == num_train_batches:
if test_lazy_data is None:
avg_val_loss = batch_loss.item() # Use train loss if no val
else:
self.encoder.eval()
self.decoder.eval()
total_val_loss = 0.0
num_val_batches = min(val_batches,
(len(test_lazy_data) + batch_size - 1) // batch_size) # Quick val subset
with torch.no_grad():
val_indices = list(range(0, len(test_lazy_data), batch_size))[:num_val_batches]
for i in val_indices:
raw_batch = test_lazy_data[i:i + batch_size]
tokenized_batch = [
[[self.sp.bos_id()] + self.sp.EncodeAsIds(s) + [self.sp.eos_id()] for s in group]
for group in raw_batch
]
group_batch = group_collate_fn(tokenized_batch)
batch_val_loss = 0.0
for group_tensor in group_batch:
if group_tensor.shape[1] <= self.max_len and group_tensor.numel() > 0:
group_tensor = group_tensor.to(self.device)
thought_vectors = self.encoder(group_tensor)
output_logits = self.decoder(thought_vectors, group_tensor[:, :-1])
val_loss = criterion(output_logits.reshape(-1, self.vocab_size),
group_tensor[:, 1:].reshape(-1))
val_loss += length_penalty * thought_vectors.shape[1]
batch_val_loss += val_loss.item()
total_val_loss += batch_val_loss
print(f" Input: {self.sp.DecodeIds(group_tensor[0].tolist())}")
print(f"Output: {self.sp.DecodeIds(output_logits.argmax(-1)[0].tolist())}\n")
avg_val_loss = total_val_loss / num_val_batches if num_val_batches > 0 else float('inf')
print(f"Val Loss: {avg_val_loss:.4f}\n\n")
# Per-batch patience
if avg_val_loss < best_val_loss:
best_val_loss = avg_val_loss
best_batch = batch_idx + epoch * num_train_batches # Total batch count
if patience_counter > 0:
patience_counter -= 1
self.save(save_path, spm_model_prefix)
else:
patience_counter += 1
if patience_counter >= patience:
print(
f"Early stopping triggered at batch {batch_idx} (total {best_batch + patience}), best val loss: {best_val_loss:.4f}")
self.load(save_path) # Roll back to best
self.save(save_path, spm_model_prefix) # Save and clean up
return # Exit training
print(f"Batch {batch_idx}/{num_train_batches} - Loss: {batch_loss.item():.4f}")
print(f" Input: {self.sp.DecodeIds(group_tensor[0].tolist())}")
print(f"Output: {self.sp.DecodeIds(output_logits.argmax(-1)[0].tolist())}\n")
self.encoder.train()
self.decoder.train()
avg_train_loss = total_train_loss / num_train_batches if num_train_batches > 0 else float('inf')
print(f"Epoch {epoch + 1}/{num_epochs} - Train Loss: {avg_train_loss:.4f}\n")
scheduler.step()
except KeyboardInterrupt:
while True:
saving = input("Would you like to save the best (b), current (c), or no (n) model: ")
saving = saving.strip().lower()
if saving in ("best", "b"):
break
elif saving in ("current", "c", "curr"):
self.save(save_path, spm_model_prefix)
break
elif saving in ("n", "no", ""):
if os.path.exists(save_path):
os.remove(save_path)
break
else:
print("invalid input, type 'b', 'c', or 'n'")
print(f"Model saved to {save_path}")
def encode(self, text: Union[str, List[str]], force_single_vector: bool = False) -> torch.Tensor:
"""Encodes text into thought vectors.
Args:
text (Union[str, List[str]]): Input text as a string or list of strings.
force_single_vector (bool, optional): If True, generates only one thought vector.
Defaults to False.
Returns:
torch.Tensor: Thought vectors, shape (batch_size, num_thoughts, d_model).
"""
if self.encoder is None or self.sp is None:
raise RuntimeError("Model not trained or loaded.")
if isinstance(text, str):
text = [text]
tokenized = [[self.sp.bos_id()] + self.sp.EncodeAsIds(t) + [self.sp.eos_id()] for t in text]
input_tokens = nn.utils.rnn.pad_sequence(
[torch.tensor(seq, dtype=torch.long) for seq in tokenized],
batch_first=True, padding_value=self.sp.pad_id()
).to(self.device)
return self.encoder(input_tokens, force_single_vector)
def decode(self, thought_vectors: torch.Tensor, max_length: int = 50, beam_width: int = 0,
temperature: float = 1.0) -> List[str]:
"""Decodes thought vectors into text sequences.
Args:
thought_vectors (torch.Tensor): Thought vectors, shape (batch_size, num_thoughts, d_model).
max_length (int, optional): Maximum length of generated sequences. Defaults to 50.
beam_width (Optional[int], optional): Beam width for beam search; if 0, uses greedy
decoding. Defaults to 5.
temperature (float, optional): Temperature for softmax sampling in greedy decoding.
Defaults to 1.0.
Returns:
List[str]: Decoded text sequences.
"""
if self.decoder is None or self.sp is None:
raise RuntimeError("Model not trained or loaded.")
thought_vectors = thought_vectors.to(self.device)
batch_size = thought_vectors.size(0)
if beam_width > 1:
return self._beam_search_decode(thought_vectors, max_length, beam_width)
target_tokens = torch.full((batch_size, 1), self.sp.bos_id(), dtype=torch.long, device=self.device)
for _ in range(max_length - 1):
logits = self.decoder(thought_vectors, target_tokens)
logits = logits[:, -1, :] / temperature
next_token = logits.softmax(dim=-1).multinomial(1)
target_tokens = torch.cat([target_tokens, next_token], dim=1)
if (next_token == self.sp.eos_id()).all():
break
return [self.sp.DecodeIds(seq.tolist()) for seq in target_tokens]
def _beam_search_decode(self, thought_vectors: torch.Tensor, max_length: int, beam_width: int) -> List[str]:
"""Performs beam search decoding of thought vectors.
Args:
thought_vectors (torch.Tensor): Thought vectors, shape (batch_size, num_thoughts, d_model).
max_length (int): Maximum length of generated sequences.
beam_width (int): Number of beams to maintain during search.
Returns:
List[str]: Decoded text sequences, one per batch item.
"""
batch_size = thought_vectors.size(0)
start_tokens = torch.full((batch_size, 1), self.sp.bos_id(), dtype=torch.long, device=self.device)
beams = [[] for _ in range(batch_size)]
finished = [[] for _ in range(batch_size)]
for i in range(batch_size):
beams[i].append((0.0, start_tokens[i].unsqueeze(0)))
for _ in range(max_length - 1):
new_beams = [[] for _ in range(batch_size)]
for i in range(batch_size):
for score, seq in beams[i]:
if seq[:, -1].item() == self.sp.eos_id():
finished[i].append((score, seq))
continue
logits = self.decoder(thought_vectors[i:i + 1], seq)[:, -1, :]
probs, next_tokens = logits.softmax(dim=-1).topk(beam_width, dim=-1)
for p, t in zip(probs[0], next_tokens[0]):
new_score = score - p.log().item()
new_seq = torch.cat([seq, t.unsqueeze(0).unsqueeze(-1)], dim=1)
new_beams[i].append((new_score, new_seq))
beams[i] = sorted(new_beams[i], key=lambda x: x[0])[:beam_width]
if all(len(finished[i]) >= beam_width for i in range(batch_size)):
break
results = []
for i in range(batch_size):
combined = sorted(finished[i] + beams[i], key=lambda x: x[0])
results.append(self.sp.DecodeIds(combined[0][1].squeeze(0).tolist()) if combined else self.sp.DecodeIds(
beams[i][0][1].squeeze(0).tolist()))
return results
def save(self, path: str, spm_model_prefix: str) -> None:
if self.encoder is None or self.decoder is None:
raise RuntimeError("No model to save.")
if not os.path.exists(f"{spm_model_prefix}.model") or not os.path.exists(f"{spm_model_prefix}.vocab"):
raise FileNotFoundError(f"SentencePiece files ({spm_model_prefix}.model or .vocab) not found.")
with tarfile.open(path, "w") as tar:
torch.save({
'encoder_state_dict': self.encoder.state_dict(),
'decoder_state_dict': self.decoder.state_dict(),
'vocab_size': self.vocab_size,
'd_model': self.d_model,
'max_thoughts': self.max_thoughts,
'encoder_nhead': self.encoder_nhead,
'decoder_nhead': self.decoder_nhead,
'encoder_layers': self.encoder_layers,
'decoder_layers': self.decoder_layers,
'dropout': self.dropout,
'max_len': self.max_len,
'termination_threshold': self.termination_threshold,
'spm_model_path': f"{spm_model_prefix}.model"
}, "model.pth")
tar.add("model.pth")
tar.add(f"{spm_model_prefix}.model")
tar.add(f"{spm_model_prefix}.vocab")
@classmethod
def load(cls, path: str) -> 'ThoughtVectors':
"""Loads a trained model from a tar file.
Args:
path (str): Path to the tar file containing the model.
Returns:
ThoughtVectors: Loaded model instance.
"""
translator = cls()
translator.device = torch.device("cuda" if torch.cuda.is_available() else "cpu") # load based on available device, not model train device
with tarfile.open(path, "r") as tar:
tar.extractall()
checkpoint = torch.load("model.pth", map_location=torch.device(translator.device))
translator.sp = spm.SentencePieceProcessor()
if not translator.sp.Load("spm.model"):
raise RuntimeError("Failed to load SentencePiece model.")
translator.vocab_size = checkpoint['vocab_size']
translator.d_model = checkpoint['d_model']
translator.max_thoughts = checkpoint['max_thoughts']
translator.encoder_nhead = checkpoint['encoder_nhead']
translator.decoder_nhead = checkpoint['decoder_nhead']
translator.encoder_layers = checkpoint['encoder_layers']
translator.decoder_layers = checkpoint['decoder_layers']
translator.dropout = checkpoint['dropout']
translator.max_len = checkpoint['max_len']
translator.termination_threshold = checkpoint['termination_threshold']
translator.encoder = ThoughtEncoder(
vocab_size=translator.vocab_size, d_model=translator.d_model, max_thoughts=translator.max_thoughts,
nhead=translator.encoder_nhead, num_layers=translator.encoder_layers, dropout=translator.dropout,
termination_threshold=translator.termination_threshold, max_len=translator.max_len
).to(translator.device)
translator.decoder = ThoughtDecoder(
vocab_size=translator.vocab_size, d_model=translator.d_model, num_layers=translator.decoder_layers,
nhead=translator.decoder_nhead, dropout=translator.dropout, max_len=translator.max_len
).to(translator.device)
translator.encoder.load_state_dict(checkpoint['encoder_state_dict'])
translator.decoder.load_state_dict(checkpoint['decoder_state_dict'])
return translator
def _add_temp_file(self, filepath: str):
"""Add a file to the set of temporary files to clean up on close."""
self._temp_files.add(os.path.abspath(filepath))
def _cleanup(self):
"""Remove all tracked temporary files."""
for filepath in self._temp_files:
if os.path.exists(filepath):
try:
os.remove(filepath)
except OSError as e:
print(f"Failed to clean up {filepath}: {e}")
# Clear the set after cleanup
self._temp_files.clear()
# example usage
if __name__ == "__main__":
tv = ThoughtVectors()
tv.train(
group_data="train.csv",
test_data="val.csv",
model="thought_vectors_prototype.tar",
num_epochs=2048,
batch_size=256,
batches_per_val=32,
val_batches=16,
accum_steps=1,
learning_rate=2e-4,
weight_decay=1e-5,
length_penalty=0.001,
single_vector_prob=0.1,
save_path="thought_vectors_prototype-0.2.0.tar",
spm_model_prefix="spm",
vocab_size=8192,
d_model=512,
encoder_nhead=8,
decoder_nhead=8,
encoder_layers=4,
decoder_layers=4,
max_thoughts=16,
dropout=0.1,
max_len=256,
termination_threshold=0.8,
patience=10
)
thought_vectors = tv.encode("AI is smart")
generated_text_greedy = tv.decode(thought_vectors, temperature=0.7, beam_width=0)
generated_text_beam = tv.decode(thought_vectors, beam_width=5)
print(f"Greedy decoding: {generated_text_greedy}")
print(f"Beam search decoding: {generated_text_beam}")