Views
No views yet
disease_llm_final.pth) and instructions for loading the DiseaseLLM model using your custom class.disease_llm_final.pth: Model weights and training metadatamodel_class.py: Custom model class (see below)1import torch
2from main import DiseaseLLM, ModelConfig
3
4def load_trained_model(checkpoint_path, config):
5 model = DiseaseLLM(config)
6 checkpoint = torch.load(checkpoint_path, map_location='cpu')
7 model.load_state_dict(checkpoint['model_state_dict'])
8 return model
9
10config = ModelConfig()
11config.vocab_size = 50257 # or your training vocab size
12model = load_trained_model('disease_llm_final.pth', config)
13model.eval()DiseaseLLM and ModelConfig classes available (see your main.py).AutoModelForCausalLM.