//! GPT decoder-only transformer implementation //! //! This module provides a complete GPT architecture with proper causal attention, //! embeddings, and language modeling heads organized into logical components for //! maintainability. //! //! ## Architecture Overview //! //! GPT (Generative Pre-trained Transformer) consists of: //! - Token and positional embeddings //! - Multi-head causal self-attention layers //! - Position-wise feed-forward networks //! - Layer normalization and residual connections //! - Language modeling head for next token prediction //! //! ## Usage //! //! ```rust,ignore //! use crate::architectures::gpt_old::{GPTConfig, GPTLMHeadModel}; //! use rtx_tensor::Device; //! //! let config = GPTConfig::gpt2_small(); //! let device = Device::cuda(0).unwrap_or(Device::default()); //! let model = GPTLMHeadModel::new(config, &device)?; //! ``` // Re-export from the main gpt module for backward compatibility pub use super::gpt::{ FeedForward, GPTBlock, GPTConfig, GPTLMHeadModel, GPTModel, LayerNorm, MultiHeadAttention, PositionalEmbedding, TokenEmbedding, }; // Legacy compatibility - keep some key types at this level for backward compatibility use crate::{Result, TransformerError}; use rtx_tensor::{Device, Tensor}; // Legacy functions for backward compatibility /// Create a simple GPT model for testing purposes pub fn create_simple_gpt_model(config: GPTConfig, device: &Device) -> Result { GPTModel::new(config, device) } /// Legacy helper to create language modeling model pub fn create_gpt_for_lm(config: GPTConfig, device: &Device) -> Result { GPTLMHeadModel::new(config, device) } // Additional convenience functions for common GPT operations /// Load GPT configuration from preset pub fn load_gpt_config(model_name: &str) -> Result { match model_name { "gpt2" | "gpt2-small" => Ok(GPTConfig::gpt2_small()), "gpt2-medium" => Ok(GPTConfig::gpt2_medium()), "gpt2-large" => Ok(GPTConfig::gpt2_large()), "gpt2-xl" => Ok(GPTConfig::gpt2_xl()), _ => { // Try parsing as GPT-3 size if let Some(size) = model_name.strip_prefix("gpt3-") { GPTConfig::gpt3(size) } else { Err(TransformerError::config(format!( "Unknown GPT model: {model_name}" ))) } } } } /// Validate GPT configuration parameters pub fn validate_gpt_config(config: &GPTConfig) -> Result<()> { if !config.hidden_size.is_multiple_of(config.num_heads) { return Err(TransformerError::config(format!( "Hidden size ({}) must be divisible by number of heads ({})", config.hidden_size, config.num_heads ))); } if config.num_layers == 0 { return Err(TransformerError::config( "Number of layers must be greater than 0".to_string(), )); } if config.vocab_size == 0 { return Err(TransformerError::config( "Vocabulary size must be greater than 0".to_string(), )); } Ok(()) } /// Estimate the number of parameters in a GPT model #[must_use] pub fn estimate_gpt_parameters(config: &GPTConfig) -> usize { // Token embeddings let token_embedding_params = config.vocab_size * config.hidden_size; // Position embeddings let pos_embedding_params = config.max_sequence_length * config.hidden_size; // Transformer layers let attention_params_per_layer = 4 * config.hidden_size * config.hidden_size; // Q, K, V, O let ff_params_per_layer = 2 * config.hidden_size * config.intermediate_size; // Up and down let ln_params_per_layer = 4 * config.hidden_size; // 2 layer norms per layer let params_per_layer = attention_params_per_layer + ff_params_per_layer + ln_params_per_layer; let total_layer_params = config.num_layers * params_per_layer; // Final layer norm let final_ln_params = config.hidden_size; // Language modeling head (tied embeddings don't add extra params) let lm_head_params = if true { // assuming tied embeddings 0 } else { config.hidden_size * config.vocab_size }; token_embedding_params + pos_embedding_params + total_layer_params + final_ln_params + lm_head_params } /// Create position IDs for a given sequence length pub fn create_position_ids(seq_len: usize, device: &Device) -> Result { let position_ids_data: Vec = (0..seq_len).map(|i| i as f32).collect(); Tensor::from_data(position_ids_data, [seq_len], device) .map_err(|e| TransformerError::tensor_op(e.to_string())) } /// Generate text with basic sampling (simplified implementation) pub fn generate_text( model: &GPTLMHeadModel, input_ids: &Tensor, max_new_tokens: usize, temperature: f32, ) -> Result { let device = input_ids.device(); let mut current_ids = input_ids.clone(); for _ in 0..max_new_tokens { // Get logits for current sequence let logits = model.forward(¤t_ids, None)?; // Extract logits for the last position let batch_size = logits.shape().dims()[0]; let seq_len = logits.shape().dims()[1]; let vocab_size = logits.shape().dims()[2]; let logits_data = logits.to_cpu()?; // For simplicity, just take the most likely token (greedy decoding) let mut next_tokens = Vec::new(); for batch_idx in 0..batch_size { let start_idx = (batch_idx * seq_len + (seq_len - 1)) * vocab_size; let last_logits = &logits_data[start_idx..start_idx + vocab_size]; let max_idx = last_logits .iter() .enumerate() .max_by(|(_, a), (_, b)| a.total_cmp(b)) .map_or(0, |(idx, _)| idx); next_tokens.push(max_idx as f32); } // Create next token tensor let next_token_tensor = Tensor::from_data(next_tokens, [batch_size, 1], device)?; // Concatenate with current sequence current_ids = Tensor::cat(&[current_ids, next_token_tensor], 1)?; } Ok(current_ids) } // Tests for backward compatibility #[cfg(all(test, feature = "disabled_tests"))] mod tests { use super::*; use rtx_tensor::Device; #[test] fn test_gpt_config_creation() { let config = GPTConfig::gpt2_small(); assert_eq!(config.hidden_size, 768); assert_eq!(config.num_heads, 12); assert_eq!(config.num_layers, 12); } #[test] fn test_gpt_config_validation() { let config = GPTConfig::gpt2_small(); assert!(validate_gpt_config(&config).is_ok()); } #[test] fn test_parameter_estimation() { let config = GPTConfig::gpt2_small(); let param_count = estimate_gpt_parameters(&config); // GPT-2 small should have approximately 117M parameters assert!(param_count > 100_000_000 && param_count < 130_000_000); } #[test] fn test_load_gpt_config() { assert!(load_gpt_config("gpt2").is_ok()); assert!(load_gpt_config("gpt2-medium").is_ok()); assert!(load_gpt_config("unknown-model").is_err()); } }