Hezu06 commited on
Commit
c688f4d
·
1 Parent(s): 6ec9e8e

update model AI

Browse files
Files changed (1) hide show
  1. main.py +13 -10
main.py CHANGED
@@ -29,17 +29,20 @@ class HeartRateData(BaseModel):
29
  # ==========================================
30
  # Cấu trúc class phải y hệt lúc train
31
  class BioVibeAI(nn.Module):
32
- def __init__(self, vocab_size, embed_size=256, hidden_size=512, num_layers=2):
33
- super(BioVibeAI, self).__init__()
34
- self.embedding = nn.Embedding(vocab_size, embed_size)
35
- self.lstm = nn.LSTM(embed_size, hidden_size, num_layers, batch_first=True)
36
- self.fc = nn.Linear(hidden_size, vocab_size)
37
-
 
 
38
  def forward(self, x):
39
- embedded = self.embedding(x)
40
- out, _ = self.lstm(embedded)
41
- out = self.fc(out[:, -1, :])
42
- return out
 
43
 
44
  # Nạp từ điển
45
  with open('mapping_dict.pkl', 'rb') as f:
 
29
  # ==========================================
30
  # Cấu trúc class phải y hệt lúc train
31
  class BioVibeAI(nn.Module):
32
+ def __init__(self, vocab_size, embed=128, hidden=256):
33
+ super().__init__()
34
+
35
+ self.embedding = nn.Embedding(vocab_size, embed)
36
+ self.lstm = nn.LSTM(embed, hidden, num_layers=2, batch_first=True, dropout=0.3)
37
+ self.norm = nn.LayerNorm(hidden)
38
+ self.fc = nn.Linear(hidden, vocab_size)
39
+
40
  def forward(self, x):
41
+ x = self.embedding(x)
42
+ out, _ = self.lstm(x)
43
+ out = out[:, -1, :]
44
+ out = self.norm(out)
45
+ return self.fc(out)
46
 
47
  # Nạp từ điển
48
  with open('mapping_dict.pkl', 'rb') as f: