MedSLM uses a modern GPT-style transformer with several architectural improvements over the standard GPT-2 design:
1import torch
2import json
3from safetensors.torch import load_file
4from transformers import AutoTokenizer
5
6# Load config
7with open("config.json") as f:
8 config_dict = json.load(f)
9
10# Reconstruct model (requires the MedSLM class definition)
11config = MedSLMConfig(**{k: v for k, v in config_dict.items()
12 if k in MedSLMConfig.__dataclass_fields__})
13model = MedSLM(config)
14
15# Load weights
16state_dict = load_file("model.safetensors")
17model.load_state_dict(state_dict)
18model.eval()
19
20# Load tokenizer
21tokenizer = AutoTokenizer.from_pretrained("tokenizer/")
1prompt = "The patient presented with acute myocardial infarction"
2input_ids = torch.tensor(tokenizer.encode(prompt)).unsqueeze(0)
3
4output = model.generate(input_ids, max_new_tokens=200, temperature=0.8, top_k=50, top_p=0.9)
5print(tokenizer.decode(output.squeeze().tolist()))
1# Load optimizer state
2optimizer_state = torch.load("optimizer.pt")
3optimizer.load_state_dict(optimizer_state)