""" RiadDiseaseGPT Model - Custom Implementation """ import torch import torch.nn as nn from typing import Optional from transformers import PreTrainedModel, PretrainedConfig from transformers.modeling_outputs import CausalLMOutputWithPast class RiadDiseaseGPTConfig(PretrainedConfig): """Configuration class for RiadDiseaseGPT""" model_type = "riad_disease_gpt" def __init__( self, vocab_size: int = 2500, hidden_size: int = 384, num_hidden_layers: int = 6, num_attention_heads: int = 6, intermediate_size: int = 1536, max_position_embeddings: int = 512, pad_token_id: int = 3, bos_token_id: int = 1, eos_token_id: int = 2, unk_token_id: int = 0, **kwargs ): super().__init__( pad_token_id=pad_token_id, bos_token_id=bos_token_id, eos_token_id=eos_token_id, unk_token_id=unk_token_id, **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 RiadDiseaseGPTModel(PreTrainedModel): """Main model class""" config_class = RiadDiseaseGPTConfig _tied_weights_keys = ["lm_head.weight"] def __init__(self, config: RiadDiseaseGPTConfig): super().__init__(config) self.config = config self.tok_emb = nn.Embedding(config.vocab_size, config.hidden_size) self.pos_emb = nn.Embedding(config.max_position_embeddings, config.hidden_size) self.drop = nn.Dropout(0.1) encoder_layer = nn.TransformerEncoderLayer( d_model=config.hidden_size, nhead=config.num_attention_heads, dim_feedforward=config.intermediate_size, batch_first=True, dropout=0.1, activation='gelu' ) self.transformer = nn.TransformerEncoder(encoder_layer, num_layers=config.num_hidden_layers) self.ln_f = nn.LayerNorm(config.hidden_size) self.lm_head = nn.Linear(config.hidden_size, config.vocab_size, bias=False) # Weight tying self.lm_head.weight = self.tok_emb.weight self.apply(self._init_weights) def get_input_embeddings(self): return self.tok_emb def set_input_embeddings(self, value): self.tok_emb = value def get_output_embeddings(self): return self.lm_head def set_output_embeddings(self, new_embeddings): self.lm_head = new_embeddings def _init_weights(self, module): if isinstance(module, nn.Linear): torch.nn.init.normal_(module.weight, mean=0.0, std=0.02) if module.bias is not None: torch.nn.init.zeros_(module.bias) elif isinstance(module, nn.Embedding): torch.nn.init.normal_(module.weight, mean=0.0, std=0.02) def _causal_mask(self, seq_len, device): return torch.triu(torch.ones(seq_len, seq_len, device=device) * float('-inf'), diagonal=1) def forward( self, input_ids: torch.LongTensor, attention_mask: Optional[torch.Tensor] = None, labels: Optional[torch.LongTensor] = None, **kwargs ) -> CausalLMOutputWithPast: b, t = input_ids.size() device = input_ids.device tok = self.tok_emb(input_ids) pos = self.pos_emb(torch.arange(0, t, device=device).unsqueeze(0)) hidden_states = self.drop(tok + pos) causal_mask = self._causal_mask(t, device) hidden_states = self.transformer(hidden_states, mask=causal_mask, is_causal=True) hidden_states = self.ln_f(hidden_states) logits = self.lm_head(hidden_states) loss = None if labels is not None: shift_logits = logits[..., :-1, :].contiguous() shift_labels = labels[..., 1:].contiguous() loss_fct = nn.CrossEntropyLoss(ignore_index=self.config.pad_token_id) loss = loss_fct(shift_logits.view(-1, shift_logits.size(-1)), shift_labels.view(-1)) return CausalLMOutputWithPast(loss=loss, logits=logits) def generate_text( self, tokenizer, prompt: str, max_length: int = 150, temperature: float = 0.7, top_k: int = 40, device: str = 'cpu' ) -> str: """Text generation method""" self.eval() device = torch.device(device) formatted_prompt = f"Q: {prompt.strip()}\nA:" input_ids = tokenizer.encode(formatted_prompt, return_tensors='pt').to(device) with torch.no_grad(): for _ in range(max_length): if input_ids.size(1) >= self.config.max_position_embeddings: input_ids = input_ids[:, -self.config.max_position_embeddings+1:] outputs = self(input_ids) logits = outputs.logits[0, -1, :] / temperature if top_k > 0: v, _ = torch.topk(logits, min(top_k, logits.size(-1))) logits[logits < v[-1]] = float('-inf') probs = torch.softmax(logits, dim=-1) next_token = torch.multinomial(probs, num_samples=1) if next_token.item() == tokenizer.eos_token_id: break input_ids = torch.cat([input_ids, next_token.unsqueeze(0)], dim=1) generated_text = tokenizer.decode(input_ids[0], skip_special_tokens=True) if 'a:' in generated_text.lower(): answer = generated_text.lower().split('a:')[-1].strip() if answer: return answer[0].upper() + answer[1:] return "I don't know."