Views
No views yet
config.json (hyperparameters) and model.safetensors (weights).| Dataset | input_d_data | patch_size | d_embed | n_layer | n_head | max_seq_len | Parameters |
|---|---|---|---|---|---|---|---|
| MSL | 55 | 2 | 512 | 6 | 8 | 512 | ~19M |
| SMAP | 25 | 4 | 512 | 6 | 8 | 512 | ~19M |
| SWaT | 50 | 14 | 512 | 6 | 8 | 512 | ~19M |
| WADI | 122 | 8 | 512 | 6 | 8 | 512 | ~19M |
1import json
2from pathlib import Path
3
4import torch
5from safetensors.torch import load_file
6
7from models.anomaly_transformer import get_anomaly_transformer
8
9
10def load_model(dataset_dir: str) -> torch.nn.Module:
11 """Load an AnomalyBERT model from config + safetensors."""
12 dataset_path = Path(dataset_dir)
13
14 with open(dataset_path / 'config.json') as f:
15 config = json.load(f)
16
17 model = get_anomaly_transformer(
18 input_d_data=config['input_d_data'],
19 output_d_data=config['output_d_data'],
20 patch_size=config['patch_size'],
21 d_embed=config['d_embed'],
22 hidden_dim_rate=config['hidden_dim_rate'],
23 max_seq_len=config['max_seq_len'],
24 positional_encoding=config['positional_encoding'],
25 relative_position_embedding=config['relative_position_embedding'],
26 transformer_n_layer=config['transformer_n_layer'],
27 transformer_n_head=config['transformer_n_head'],
28 dropout=config['dropout'],
29 )
30
31 state_dict = load_file(str(dataset_path / 'model.safetensors'))
32 model.load_state_dict(state_dict)
33 model.eval()
34 return model
35
36
37# Example: load the MSL model
38model = load_model('MSL')
39
40# Inference
41# x shape: (batch, patch_size * max_seq_len, input_d_data)
42x = torch.randn(1, 1024, 55)
43with torch.no_grad():
44 output = model(x)
45# output shape: (batch, patch_size * max_seq_len, output_d_data)├── MSL/
│ ├── config.json
│ └── model.safetensors
├── SMAP/
│ ├── config.json
│ └── model.safetensors
├── SWaT/
│ ├── config.json
│ └── model.safetensors
├── WADI/
│ ├── config.json
│ └── model.safetensors
├── convert_to_hf.py # Conversion script (.pt -> safetensors)
├── inspect_pt.py # Checkpoint inspection script
└── verify_conversion.py # Conversion verification script1@article{jeong2023anomalybert,
2 title={AnomalyBERT: Self-Supervised Transformer for Time Series Anomaly Detection using Data Degradation Scheme},
3 author={Jeong, Yungi and Yang, Eunseok and Ryu, Jung Hyun and Park, Imseong and Kang, Myungjoo},
4 journal={arXiv preprint arXiv:2305.04468},
5 year={2023}
6}