Views
No views yet
1batch_size = 16
2embed_dim = 256
3num_heads = 4
4ff_dim = 512
5num_layers = 2
6noise_prob = 0.3
7device = torch.device('cuda' if torch.cuda.is_available() else 'cpu')
8
9class PositionalEncoding(nn.Module):
10 def __init__(self, d_model, max_len=5000):
11 super().__init__()
12 pe = torch.zeros(max_len, d_model)
13 position = torch.arange(0, max_len, dtype=torch.float).unsqueeze(1)
14 div_term = torch.exp(torch.arange(0, d_model, 2).float() * (-np.log(10000.0) / d_model))
15 pe[:, 0::2] = torch.sin(position * div_term)
16 pe[:, 1::2] = torch.cos(position * div_term)
17 self.register_buffer('pe', pe.unsqueeze(0))
18
19 def forward(self, x):
20 return x + self.pe[:, :x.size(1)]
21
22class TransformerBlock(nn.Module):
23 def __init__(self, embed_dim, num_heads, ff_dim):
24 super().__init__()
25 self.attention = nn.MultiheadAttention(embed_dim, num_heads)
26 self.norm1 = nn.LayerNorm(embed_dim)
27 self.ff = nn.Sequential(
28 nn.Linear(embed_dim, ff_dim),
29 nn.ReLU(),
30 nn.Linear(ff_dim, embed_dim)
31 )
32 self.norm2 = nn.LayerNorm(embed_dim)
33
34 def forward(self, x):
35 attn_output, _ = self.attention(x, x, x)
36 x = self.norm1(x + attn_output)
37 ff_output = self.ff(x)
38 return self.norm2(x + ff_output)
39
40class DenoisingTransformer(nn.Module):
41 def __init__(self, vocab_size, embed_dim, num_heads, ff_dim, num_layers):
42 super().__init__()
43 self.embedding = nn.Embedding(vocab_size, embed_dim)
44 self.positional_encoding = PositionalEncoding(embed_dim)
45 self.transformer_blocks = nn.ModuleList([
46 TransformerBlock(embed_dim, num_heads, ff_dim) for _ in range(num_layers)
47 ])
48 self.fc = nn.Linear(embed_dim, vocab_size)
49
50 def forward(self, x):
51 x = self.embedding(x)
52 x = self.positional_encoding(x)
53 for block in self.transformer_blocks:
54 x = block(x)
55 return self.fc(x)
56
57def load_model(path, device='cpu'):
58 checkpoint = torch.load(path, map_location=device)
59 hp = checkpoint['hyperparameters']
60
61 model = DenoisingTransformer(
62 hp['vocab_size'],
63 hp['embed_dim'],
64 hp['num_heads'],
65 hp['ff_dim'],
66 hp['num_layers']
67 ).to(device)
68
69 model.load_state_dict(checkpoint['model_state_dict'])
70 return model, checkpoint['word2idx'], checkpoint['idx2word']
71
72loaded_model, word2idx, idx2word = load_model('denoising_transformer.pth', device=device)
73
74print("Model loaded successfully!")
75print(f"Model device: {next(loaded_model.parameters()).device}")```