from torch.utils.data import Dataset, DataLoader
class TextDataset(Dataset):
def __init__(self, text, seq_length):
self.chars = sorted(list(set(text)))
self.vocab_size = len(self.chars)
self.char_to_idx = {char: idx for idx, char in enumerate(self.chars)}
self.idx_to_char = {idx: char for idx, char in enumerate(self.chars)}
self.seq_length = seq_length
self.text_encoded = [self.char_to_idx[char] for char in text]
for i in range(0, len(self.text_encoded) - seq_length):
self.data.append((torch.tensor(self.text_encoded[i:i+seq_length], dtype=torch.long),
torch.tensor(self.text_encoded[i+1:i+seq_length+1], dtype=torch.long)))
def __getitem__(self, index):
def get_random_start(self):
return torch.tensor([self.char_to_idx[random.choice(self.chars)]], dtype=torch.long)
class LSTMTextGenerator(nn.Module):
def __init__(self, vocab_size, embedding_dim, hidden_dim, n_layers):
super(LSTMTextGenerator, self).__init__()
self.hidden_dim = hidden_dim
self.embedding = nn.Embedding(vocab_size, embedding_dim)
self.lstm = nn.LSTM(embedding_dim, hidden_dim, n_layers, batch_first=True)
self.fc = nn.Linear(hidden_dim, vocab_size)
def forward(self, x, hidden=None):
embeds = self.embedding(x)
lstm_out, hidden = self.lstm(embeds, hidden)
output = self.fc(lstm_out)
def init_hidden(self, batch_size, device):
torch.zeros(self.lstm.num_layers, batch_size, self.hidden_dim).to(device),
torch.zeros(self.lstm.num_layers, batch_size, self.hidden_dim).to(device)
def train(model, dataloader, criterion, optimizer, device):
for inputs, targets in dataloader:
inputs, targets = inputs.to(device), targets.to(device)
hidden = model.init_hidden(inputs.size(0), device)
outputs, hidden = model(inputs, hidden)
loss = criterion(outputs.reshape(-1, outputs.shape[-1]), targets.reshape(-1))
total_loss += loss.item()
return total_loss / len(dataloader)
def generate_text(model, start_token, num_generate, device, temperature=1.0):
hidden = model.init_hidden(1, device) # 初始化隐藏状态和细胞状态
for _ in range(num_generate):
input_tensor = generated[-1].unsqueeze(0).to(device) # 维度为 (1, embedding_dim)
output, hidden = model(input_tensor, (hidden[0].squeeze(0), hidden[1].squeeze(0)))
output = output[:, -1, :].div(temperature).exp()
predicted_idx = torch.multinomial(output, 1)
generated = torch.cat((generated, predicted_idx), dim=1)
return generated.cpu().numpy().flatten()
if __name__ == "__main__":
text = "Hello, this is an example of text used for LSTM text generation. Let's generate some text!"
device = torch.device("cuda" if torch.cuda.is_available() else "cpu")
dataset = TextDataset(text, seq_length)
data_loader = DataLoader(dataset, batch_size=batch_size, shuffle=True)
model = LSTMTextGenerator(dataset.vocab_size, embedding_dim, hidden_dim, n_layers).to(device)
criterion = nn.CrossEntropyLoss()
optimizer = torch.optim.Adam(model.parameters(), lr=0.001)
for epoch in range(1, n_epochs + 1):
loss = train(model, data_loader, criterion, optimizer, device)
print(f"Epoch: {epoch}/{n_epochs}, Loss: {loss:.4f}")
start_token = dataset.get_random_start().unsqueeze(0).to(device)
generated_text = generate_text(model, start_token, num_generate=200, device=device)
generated_chars = [dataset.idx_to_char[int(idx)] for idx in generated_text]
print(''.join(generated_chars))