Untitled
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