Views
No views yet
1import torch
2from soothsayer import Soothsayer
3
4model = Soothsayer(input_dim=768, hidden_dim=256, n_heads=4,
5 n_layers=3, ff_dim=512, dropout=0.1)
6model.load_state_dict(torch.load("soothsayer_grid/lejepa_gedicnn_k5.pt"))
7model.eval()
8
9# embeddings: (batch, seq_len, 768) from GEDI-CNN penultimate layer
10# hours: (batch, seq_len) elapsed hours
11# lengths: (batch,) actual sequence lengths
12pred_next, death_logit, ttd_pred = model(embeddings, hours, lengths)
13death_prob = torch.sigmoid(death_logit)