""" Custom Disease GPT Model for Hugging Face """ import torch import torch.nn as nn from transformers import PreTrainedModel, PretrainedConfig class DiseaseGPTConfig(PretrainedConfig): model_type = "disease-gpt" def __init__( self, vocab_size=2500, hidden_size=256, num_hidden_layers=6, num_attention_heads=8, intermediate_size=1024, max_position_embeddings=512, **kwargs ): super().__init__(**kwargs) self.vocab_size = vocab_size self.hidden_size = hidden_size self.num_hidden_layers = num_hidden_layers self.num_attention_heads = num_attention_heads self.intermediate_size = intermediate_size self.max_position_embeddings = max_position_embeddings class DiseaseGPTModel(PreTrainedModel): config_class = DiseaseGPTConfig def __init__(self, config): super().__init__(config) self.config = config # Your model architecture here self.embedding = nn.Embedding(config.vocab_size, config.hidden_size) self.pos_embedding = nn.Embedding(config.max_position_embeddings, config.hidden_size) # Transformer layers encoder_layer = nn.TransformerEncoderLayer( d_model=config.hidden_size, nhead=config.num_attention_heads, dim_feedforward=config.intermediate_size, batch_first=True ) self.transformer = nn.TransformerEncoder(encoder_layer, config.num_hidden_layers) self.fc_out = nn.Linear(config.hidden_size, config.vocab_size) def forward(self, input_ids, attention_mask=None): positions = torch.arange(0, input_ids.size(1), device=input_ids.device).unsqueeze(0) x = self.embedding(input_ids) + self.pos_embedding(positions) if attention_mask is not None: attention_mask = attention_mask.float() attention_mask = (1.0 - attention_mask) * -10000.0 x = self.transformer(x, src_key_padding_mask=attention_mask) logits = self.fc_out(x) return {"logits": logits} def generate_text(self, tokenizer, prompt, max_length=100, temperature=0.8): """Generate text from prompt""" self.eval() tokens = tokenizer.encode(prompt) tokens = torch.tensor([tokens]) with torch.no_grad(): for _ in range(max_length): outputs = self.forward(tokens) logits = outputs["logits"][:, -1, :] / temperature probs = torch.softmax(logits, dim=-1) next_token = torch.multinomial(probs, 1) tokens = torch.cat([tokens, next_token], dim=1) if next_token.item() == tokenizer.eos_token_id: break return tokenizer.decode(tokens[0].tolist())