Untitled

 avatar
unknown
plain_text
9 months ago
10 kB
15
Indexable
"""Utility functions for KV-Map."""

from typing import Dict, Tuple, Optional, Any
import jax
import jax.numpy as jnp
import numpy as np

from MaxText import max_logging
from MaxText import max_utils
from MaxText.globals import EPS


def extract_hidden_states(
    model,
    params: Dict,
    input_tokens: jnp.ndarray,
    input_positions: jnp.ndarray,
    segment_ids: Optional[jnp.ndarray] = None,
    layer_idx: int = -1,
) -> jnp.ndarray:
  """Extract hidden states from a model at a specific layer."""
  outputs = model.apply(
      params,
      input_tokens,
      input_positions,
      decoder_segment_ids=segment_ids,
      enable_dropout=False,
      mutable='intermediates'
  )
  
  if isinstance(outputs, tuple):
    logits, intermediates = outputs
    # Extract the desired layer's hidden states
    if 'intermediates' in intermediates:
      layer_key = f'decoder_layer_{layer_idx}' if layer_idx >= 0 else 'final_layer'
      if layer_key in intermediates['intermediates']:
        return intermediates['intermediates'][layer_key]
  
  # Fallback: return embedding if intermediates not available
  max_logging.log("Warning: Could not extract intermediate hidden states")
  return jnp.zeros((input_tokens.shape[0], input_tokens.shape[1], params.get('hidden_dim', 512)))


def extract_kv_from_attention(
    model,
    params: Dict,
    input_tokens: jnp.ndarray,
    input_positions: jnp.ndarray,
    segment_ids: Optional[jnp.ndarray] = None,
) -> Dict[Tuple[int, int], Tuple[jnp.ndarray, jnp.ndarray]]:
  """Extract key-value pairs from all attention layers and heads.
  """
  kv_dict = {}
  
  # Run model with mutable intermediates to capture K/V
  outputs = model.apply(
      params,
      input_tokens,
      input_positions,
      decoder_segment_ids=segment_ids,
      enable_dropout=False,
      mutable='intermediates'
  )
  
  if isinstance(outputs, tuple):
    _, intermediates = outputs
    if 'intermediates' in intermediates:
      # Parse intermediates for attention K/V
      for key, value in intermediates['intermediates'].items():
        if 'attention' in key and 'kv' in key:
          # Expected format: 'layer_X_head_Y_kv'
          parts = key.split('_')
          if len(parts) >= 4:
            try:
              layer = int(parts[1])
              head = int(parts[3])
              if isinstance(value, tuple) and len(value) == 2:
                kv_dict[(layer, head)] = value
            except (ValueError, IndexError):
              continue
  
  return kv_dict


def compute_attention_weights(
    query: jnp.ndarray,
    key: jnp.ndarray,
    segment_ids: Optional[jnp.ndarray] = None,
    scale: Optional[float] = None,
) -> jnp.ndarray:
  """Compute attention weights from query and key.
  
  Args:
    query: Query tensor, shape [batch, num_heads, seq_len_q, head_dim]
    key: Key tensor, shape [batch, num_heads, seq_len_k, head_dim]
    segment_ids: Optional segment IDs for masking
    scale: Optional scale factor (default: 1/sqrt(head_dim))
    
  Returns:
    Attention weights, shape [batch, num_heads, seq_len_q, seq_len_k]
  """
  head_dim = query.shape[-1]
  if scale is None:
    scale = 1.0 / jnp.sqrt(head_dim)
  
  # Compute attention scores
  scores = jnp.einsum('bhqd,bhkd->bhqk', query, key) * scale
  
  # Apply causal mask
  seq_len_q = query.shape[2]
  seq_len_k = key.shape[2]
  causal_mask = jnp.tril(jnp.ones((seq_len_q, seq_len_k)))
  scores = jnp.where(causal_mask, scores, -1e10)
  
  # Apply segment mask if provided
  if segment_ids is not None:
    segment_mask = segment_ids[:, None, :, None] == segment_ids[:, None, None, :]
    scores = jnp.where(segment_mask, scores, -1e10)
  
  # Softmax
  weights = jax.nn.softmax(scores, axis=-1)
  
  return weights


def compute_kl_divergence(
    p: jnp.ndarray,
    q: jnp.ndarray,
    eps: float = EPS
) -> jnp.ndarray:
  # Clip to avoid log(0)
  p_safe = jnp.clip(p, eps, 1.0)
  q_safe = jnp.clip(q, eps, 1.0)
  
  # Compute KL divergence
  kl = jnp.sum(p_safe * jnp.log(p_safe / q_safe))
  
  return kl


def compute_js_divergence(
    p: jnp.ndarray,
    q: jnp.ndarray,
    eps: float = EPS
) -> jnp.ndarray:
  m = 0.5 * (p + q)
  js = 0.5 * compute_kl_divergence(p, m, eps) + 0.5 * compute_kl_divergence(q, m, eps)
  return js


def create_synthetic_kv_data(
    batch_size: int,
    context_len: int,
    query_len: int,
    vocab_size: int,
    rng: jax.random.PRNGKey,
) -> Dict[str, jnp.ndarray]:
  rng1, rng2, rng3 = jax.random.split(rng, 3)
  
  # Generate random token sequences
  context_inputs = jax.random.randint(
      rng1, (batch_size, context_len), 0, vocab_size
  )
  query_inputs = jax.random.randint(
      rng2, (batch_size, query_len), 0, vocab_size
  )
  
  # Positions are sequential
  context_positions = jnp.arange(context_len)[None, :].repeat(batch_size, axis=0)
  query_positions = jnp.arange(context_len, context_len + query_len)[None, :].repeat(batch_size, axis=0)
  
  # Targets are shifted query inputs
  targets = jnp.roll(query_inputs, shift=-1, axis=1)
  # Set last target to 0 (will be masked out)
  targets = targets.at[:, -1].set(0)
  
  # All weights are 1 except the last position
  weights = jnp.ones((batch_size, query_len))
  weights = weights.at[:, -1].set(0.0)
  
  return {
      'context_inputs': context_inputs,
      'context_positions': context_positions,
      'query_inputs': query_inputs,
      'query_positions': query_positions,
      'targets': targets,
      'weights': weights,
  }


def merge_context_and_query(
    context_tokens: jnp.ndarray,
    query_tokens: jnp.ndarray,
    pad_token_id: int = 0,
) -> Tuple[jnp.ndarray, jnp.ndarray, jnp.ndarray]:
  batch_size = context_tokens.shape[0]
  context_len = context_tokens.shape[1]
  query_len = query_tokens.shape[1]
  total_len = context_len + query_len
  
  # Concatenate tokens
  merged_tokens = jnp.concatenate([context_tokens, query_tokens], axis=1)
  
  # Create position indices
  merged_positions = jnp.arange(total_len)[None, :].repeat(batch_size, axis=0)
  
  # Create segment IDs
  segment_ids = jnp.concatenate([
      jnp.ones((batch_size, context_len), dtype=jnp.int32),
      jnp.ones((batch_size, query_len), dtype=jnp.int32) * 2,
  ], axis=1)
  
  return merged_tokens, merged_positions, segment_ids


def compute_cache_metrics(
    kv_ideal: Dict[Tuple[int, int], Tuple[jnp.ndarray, jnp.ndarray]],
    kv_adapted: Dict[Tuple[int, int], Tuple[jnp.ndarray, jnp.ndarray]],
):
  key_mses = []
  value_mses = []
  key_cosines = []
  value_cosines = []
  
  for (layer, head), (k_ideal, v_ideal) in kv_ideal.items():
    if (layer, head) not in kv_adapted:
      continue
    
    k_adapted, v_adapted = kv_adapted[(layer, head)]
    
    # Compute MSE
    key_mse = jnp.mean((k_ideal - k_adapted) ** 2)
    value_mse = jnp.mean((v_ideal - v_adapted) ** 2)
    key_mses.append(key_mse)
    value_mses.append(value_mse)
    
    # Compute cosine similarity
    k_ideal_norm = k_ideal / (jnp.linalg.norm(k_ideal, axis=-1, keepdims=True) + EPS)
    k_adapted_norm = k_adapted / (jnp.linalg.norm(k_adapted, axis=-1, keepdims=True) + EPS)
    key_cosine = jnp.mean(jnp.sum(k_ideal_norm * k_adapted_norm, axis=-1))
    
    v_ideal_norm = v_ideal / (jnp.linalg.norm(v_ideal, axis=-1, keepdims=True) + EPS)
    v_adapted_norm = v_adapted / (jnp.linalg.norm(v_adapted, axis=-1, keepdims=True) + EPS)
    value_cosine = jnp.mean(jnp.sum(v_ideal_norm * v_adapted_norm, axis=-1))
    
    key_cosines.append(key_cosine)
    value_cosines.append(value_cosine)
  
  return {
      'mean_key_mse': float(jnp.mean(jnp.array(key_mses))) if key_mses else 0.0,
      'mean_value_mse': float(jnp.mean(jnp.array(value_mses))) if value_mses else 0.0,
      'mean_key_cosine': float(jnp.mean(jnp.array(key_cosines))) if key_cosines else 0.0,
      'mean_value_cosine': float(jnp.mean(jnp.array(value_cosines))) if value_cosines else 0.0,
  }


def get_adapter_parameter_count(params: Dict) -> int:
  """Count total number of parameters in adapter system.
  
  Args:
    params: Adapter parameters pytree
    
  Returns:
    Total number of parameters
  """
  return sum(x.size for x in jax.tree_util.tree_leaves(params))


def estimate_memory_usage(
    num_layers: int,
    num_heads: int,
    batch_size: int,
    seq_len: int,
    d_k: int,
    dtype: jnp.dtype = jnp.float32,
) -> Dict[str, float]:
  """Estimate memory usage for KV-Map components.
  """
  bytes_per_element = jnp.dtype(dtype).itemsize
  
  # KV cache size per layer/head
  kv_cache_per_head = 2 * batch_size * seq_len * d_k * bytes_per_element
  total_kv_cache = num_layers * num_heads * kv_cache_per_head
  
  adapter_params_per_head = (512 * 512 + 512 * d_k * 2 + d_k * 256 + 256 * d_k * 2) * bytes_per_element
  total_adapter_params = num_layers * num_heads * adapter_params_per_head
  
  return {
      'kv_cache_gb': total_kv_cache / (1024 ** 3),
      'adapter_params_gb': total_adapter_params / (1024 ** 3),
      'total_estimated_gb': (total_kv_cache + total_adapter_params) / (1024 ** 3),
  }


def validate_cache_shapes(
    kv_dict: Dict[Tuple[int, int], Tuple[jnp.ndarray, jnp.ndarray]],
    expected_num_layers: int,
    expected_num_heads: int,
    expected_seq_len: int,
    expected_d_k: int,
) -> bool:
  """Validate that KV cache has expected shapes.
  """
  expected_keys = {(l, h) for l in range(expected_num_layers) for h in range(expected_num_heads)}
  
  if set(kv_dict.keys()) != expected_keys:
    max_logging.log(f"Warning: Missing or extra keys in KV cache. "
                   f"Expected {len(expected_keys)}, got {len(kv_dict)}")
    return False
  
  for (layer, head), (k, v) in kv_dict.items():
    if k.shape[-2:] != (expected_seq_len, expected_d_k):
      max_logging.log(f"Warning: Key shape mismatch at layer {layer}, head {head}. "
                     f"Expected [..., {expected_seq_len}, {expected_d_k}], got {k.shape}")
      return False
    
    if v.shape[-2:] != (expected_seq_len, expected_d_k):
      max_logging.log(f"Warning: Value shape mismatch at layer {layer}, head {head}. "
                     f"Expected [..., {expected_seq_len}, {expected_d_k}], got {v.shape}")
      return False
  
  return True

Editor is loading...
Leave a Comment