Q-TensorFormer / src /hf_model.py
Premchandyadav369
Implement top-tier research features: PID Dual Controller, GQA, Roofline Analyzer, Early-Exit, Meyer-Wallach Entanglement, and Interactive Visual Dashboard
0431133
Raw History Blame Contribute Delete
9.13 kB
"""
Hugging Face Transformers Integration for Q-TensorFormer.
Provides first-class Hugging Face ecosystem compatibility:
- QTensorFormerConfig (inherits PretrainedConfig)
- QTensorFormerForCausalLM (inherits PreTrainedModel, GenerationMixin)
- AutoConfig & AutoModelForCausalLM registration
- Generation support with past_key_values and Adaptive KV Cache
- save_pretrained and from_pretrained serialization
"""
import torch
import torch.nn as nn
import torch.nn.functional as F
import math
from typing import Optional, Tuple, Dict, List, Union
from dataclasses import dataclass
from transformers import PretrainedConfig, PreTrainedModel, GenerationMixin
from transformers.modeling_outputs import CausalLMOutputWithPast
from transformers import AutoConfig, AutoModelForCausalLM
from .blocks import HybridBlock
from .kv_cache import AdaptiveKVCache, KVPrecision
from .resource_allocator import AllocationBudget
class QTensorFormerConfig(PretrainedConfig):
"""
Configuration class for Q-TensorFormer.
"""
model_type = "qtensorformer"
def __init__(
self,
vocab_size: int = 10000,
d_model: int = 128,
n_layers: int = 2,
n_heads: int = 4,
n_kv_heads: Optional[int] = None,
ff_multiplier: int = 4,
max_seq_len: int = 128,
dropout: float = 0.1,
tt_rank: int = 8,
tt_min_rank: int = 2,
use_quantum: bool = True,
n_qubits: int = 4,
backend_type: str = "classical_surrogate",
default_preset: str = "balanced",
hysteresis_tau: float = 0.15,
tie_word_embeddings: bool = True,
enable_early_exit: bool = False,
early_exit_threshold: float = 0.20,
**kwargs,
):
super().__init__(tie_word_embeddings=tie_word_embeddings, **kwargs)
self.vocab_size = vocab_size
self.d_model = d_model
self.n_layers = n_layers
self.n_heads = n_heads
self.n_kv_heads = n_kv_heads if n_kv_heads is not None else n_heads
self.ff_multiplier = ff_multiplier
self.max_seq_len = max_seq_len
self.dropout = dropout
self.tt_rank = tt_rank
self.tt_min_rank = tt_min_rank
self.use_quantum = use_quantum
self.n_qubits = n_qubits
self.backend_type = backend_type
self.default_preset = default_preset
self.hysteresis_tau = hysteresis_tau
self.enable_early_exit = enable_early_exit
self.early_exit_threshold = early_exit_threshold
class PositionalEncoding(nn.Module):
def __init__(self, d_model: int, max_len: int = 512, dropout: float = 0.1):
super().__init__()
self.dropout = nn.Dropout(dropout)
pe = torch.zeros(max_len, d_model)
pos = torch.arange(0, max_len).float().unsqueeze(1)
div = torch.exp(torch.arange(0, d_model, 2).float() * (-math.log(10000.0) / d_model))
pe[:, 0::2] = torch.sin(pos * div)
pe[:, 1::2] = torch.cos(pos * div)
self.register_buffer("pe", pe.unsqueeze(0))
def forward(self, x: torch.Tensor, start_pos: int = 0) -> torch.Tensor:
seq_len = x.size(1)
return self.dropout(x + self.pe[:, start_pos:start_pos + seq_len, :])
class QTensorFormerForCausalLM(PreTrainedModel, GenerationMixin):
"""
Q-TensorFormer Causal Language Model.
Compatible with Hugging Face generate() and AutoModelForCausalLM.
"""
config_class = QTensorFormerConfig
base_model_prefix = "qtensorformer"
_tied_weights_keys = ["lm_head.weight"]
def __init__(self, config: QTensorFormerConfig):
super().__init__(config)
self.config = config
self.embedding = nn.Embedding(config.vocab_size, config.d_model)
self.pos_encoding = PositionalEncoding(config.d_model, config.max_seq_len, config.dropout)
self.blocks = nn.ModuleList([
HybridBlock(
d_model=config.d_model,
n_heads=config.n_heads,
n_kv_heads=config.n_kv_heads,
ff_multiplier=config.ff_multiplier,
tt_rank=config.tt_rank,
tt_min_rank=config.tt_min_rank,
use_quantum=config.use_quantum,
n_qubits=config.n_qubits,
backend_type=config.backend_type,
dropout=config.dropout,
max_seq_len=config.max_seq_len,
hysteresis_tau=config.hysteresis_tau,
default_preset=config.default_preset,
enable_early_exit=config.enable_early_exit,
early_exit_threshold=config.early_exit_threshold,
vocab_size=config.vocab_size,
)
for _ in range(config.n_layers)
])
self.ln_f = nn.LayerNorm(config.d_model)
self.lm_head = nn.Linear(config.d_model, config.vocab_size, bias=False)
self.lm_head.weight = self.embedding.weight # Weight tying
self.post_init()
def get_input_embeddings(self):
return self.embedding
def set_input_embeddings(self, value):
self.embedding = value
def get_output_embeddings(self):
return self.lm_head
def set_output_embeddings(self, new_embeddings):
self.lm_head = new_embeddings
def prepare_inputs_for_generation(
self,
input_ids: torch.Tensor,
past_key_values: Optional[List[AdaptiveKVCache]] = None,
attention_mask: Optional[torch.Tensor] = None,
**kwargs,
) -> Dict:
# If past_key_values are present, only pass the latest single token
if past_key_values is not None and past_key_values[0].seq_len > 0:
input_ids = input_ids[:, -1:]
return {
"input_ids": input_ids,
"past_key_values": past_key_values,
"attention_mask": attention_mask,
"use_cache": kwargs.get("use_cache", True),
}
def forward(
self,
input_ids: torch.Tensor,
attention_mask: Optional[torch.Tensor] = None,
past_key_values: Optional[List[AdaptiveKVCache]] = None,
labels: Optional[torch.Tensor] = None,
use_cache: bool = True,
output_attentions: Optional[bool] = None,
output_hidden_states: Optional[bool] = None,
return_dict: Optional[bool] = None,
preset: Optional[str] = None,
budget: Optional[AllocationBudget] = None,
force_classical: bool = False,
) -> Union[Tuple, CausalLMOutputWithPast]:
return_dict = return_dict if return_dict is not None else self.config.use_return_dict
B, T = input_ids.shape
start_pos = 0
if past_key_values is not None and past_key_values[0].seq_len > 0:
start_pos = past_key_values[0].seq_len
# Initialize KV caches if requested and not provided
if use_cache and past_key_values is None:
past_key_values = [
AdaptiveKVCache(max_capacity=self.config.max_seq_len)
for _ in range(self.config.n_layers)
]
x = self.embedding(input_ids)
x = self.pos_encoding(x, start_pos=start_pos)
all_stats = []
early_exit_logits = None
for i, block in enumerate(self.blocks):
layer_kv = past_key_values[i] if past_key_values is not None else None
x, stats = block(
x,
mask=attention_mask,
kv_cache=layer_kv,
budget=budget,
preset=preset or self.config.default_preset,
force_classical=force_classical,
)
all_stats.append(stats)
if stats.get("early_exit_triggered", False) and "early_exit_logits" in stats:
early_exit_logits = stats["early_exit_logits"]
break
if early_exit_logits is not None:
logits = early_exit_logits
else:
x = self.ln_f(x)
logits = self.lm_head(x)
loss = None
if labels is not None:
shift_logits = logits[..., :-1, :].contiguous()
shift_labels = labels[..., 1:].contiguous()
loss = F.cross_entropy(
shift_logits.view(-1, shift_logits.size(-1)),
shift_labels.view(-1),
ignore_index=-100,
)
if not return_dict:
output = (logits, past_key_values)
return ((loss,) + output) if loss is not None else output
return CausalLMOutputWithPast(
loss=loss,
logits=logits,
past_key_values=past_key_values,
hidden_states=None,
attentions=None,
)
def set_preset(self, preset_name: str):
"""Set deployment preset: 'full', 'balanced', 'latency', 'memory', 'energy', 'edge', 'classical_only'."""
self.config.default_preset = preset_name.lower()
# Register with Hugging Face Auto classes
try:
AutoConfig.register("qtensorformer", QTensorFormerConfig)
AutoModelForCausalLM.register(QTensorFormerConfig, QTensorFormerForCausalLM)
except Exception:
pass