Cross-Attention - Multimodal Fusion

# Cross-Attention - Multimodal Fusion

## Introduction & Motivation

Cross-attention: attend between different sequences. Enable feature fusion across modalities. Applications: vision-language models, video understanding.

Motivation: Combine information from different sources or modalities.

Applications: Image-text models, video captioning, visual question answering.

---

## Core Concepts & Theory

### Source Attention

Query from one sequence, key-value from another.

### Modality Fusion

Combine cross-modal representations.

### Interaction Matrix

Relate elements across modalities.

### Asymmetric Processing

Different roles for query and key-value.

---

## Mathematical Formulation

Cross-Attention:
$$ ext{CrossAttn}(Q, K, V) = ext{softmax}(\frac{QK^T}{\sqrt{d_k}}) V$$

Where Q comes from one modality, K,V from another.

Bidirectional Fusion:
$$F = ext{CrossAttn}(A, B) + ext{CrossAttn}(B, A)$$

---

## Advanced Theory & Extensions

### Multi-Modal Fusion

Combine 3+ modalities.

### Temporal Cross-Attention

Attend across time steps.

### Hierarchical Fusion

Multi-level cross-attention.

---

## Computational Considerations

Cross-attention: O(T_q × T_kv × D).

Asymmetric lengths: More flexible.

Fusion overhead: Linear in sequence lengths.

---

## Practical Implementation Strategies

### Modality Alignment

Project to shared space first.

### Layer-wise Fusion

Cross-attention at each layer.

### Gating Mechanisms

Control information flow.

---

## Benchmark Datasets & Evaluation

COCO Captions: Image-text matching.

Visual QA: Vision-language understanding.

MSR-VTT: Video understanding.

---

## Key Challenges & Limitations

### Modality Gap

Bridging different representations.

### Alignment Quality

Matching elements across modalities.

### Computational Cost

O(T_q × T_kv).

---

## Hyperparameter Tuning

Query dimension: 64-256.

Key-value dimension: 64-256.

Number of heads: 8-16.

---

## Real-World Applications & Case Studies

Image Captioning: Vision-to-text generation.

Visual QA: Answer questions about images.

Video Understanding: Temporal and spatial fusion.

---

## Integration with Other Methods

Cross-attention + self-attention for end-to-end fusion.

---

## Summary & Key Takeaways

Cross-attention enables effective multimodal fusion.

Principles:
1. Asymmetric: Query from one modality.
2. Fusion: Combine different representations.
3. Interaction: Relate across modalities.
4. Flexibility: Variable lengths.
5. Scalability: Practical for multiple modalities.

---

## Appendix: Practical Labs

### Lab 1: Cross-Attention Computation

import numpy as np

def cross_attention(query, key, value):
 """Compute cross-attention between sequences"""
 d_k = query.shape[-1]
 
 # Query from one modality, key-value from another
 scores = query @ key.T / np.sqrt(d_k)
 
 # Attention weights
 attention = np.exp(scores) / np.sum(np.exp(scores), axis=-1, keepdims=True)
 
 # Apply to values
 output = attention @ value
 
 return output, attention

np.random.seed(42)
q = np.random.randn(10, 64) # Query from modality A
k = np.random.randn(20, 64) # Key from modality B
v = np.random.randn(20, 64) # Value from modality B
out, attn = cross_attention(q, k, v)
assert out.shape == (10, 64)
assert attn.shape == (10, 20)
print("✓ Cross-attention working")

### Lab 2: Bidirectional Fusion

import numpy as np

def bidirectional_cross_attention(modality_a, modality_b):
 """Fuse two modalities bidirectionally"""
 # A → B
 fusion_ab = modality_a @ modality_b.T
 
 # B → A
 fusion_ba = modality_b @ modality_a.T
 
 # Combined
 combined = fusion_ab + fusion_ba.T
 
 return combined

np.random.seed(42)
mod_a = np.random.randn(10, 256)
mod_b = np.random.randn(15, 256)
combined = bidirectional_cross_attention(mod_a, mod_b)
assert combined.shape == (10, 15)
print("✓ Bidirectional fusion working")

### Lab 3: Multi-Modal Alignment

import numpy as np

def align_modalities(visual_feat, text_feat):
 """Align visual and text features"""
 # Project to shared space
 visual_proj = visual_feat @ np.random.randn(256, 128)
 text_proj = text_feat @ np.random.randn(256, 128)
 
 # Normalize
 visual_proj = visual_proj / (np.linalg.norm(visual_proj, axis=1, keepdims=True) + 1e-8)
 text_proj = text_proj / (np.linalg.norm(text_proj, axis=1, keepdims=True) + 1e-8)
 
 # Similarity
 similarity = visual_proj @ text_proj.T
 
 return similarity

np.random.seed(42)
vis = np.random.randn(10, 256)
txt = np.random.randn(8, 256)
sim = align_modalities(vis, txt)
assert sim.shape == (10, 8)
print("✓ Multi-modal alignment working")

### Lab 4: Gated Cross-Attention

import numpy as np

def gated_cross_attention(query, key, value, gate=None):
 """Cross-attention with gating mechanism"""
 d_k = query.shape[-1]
 
 # Standard cross-attention
 scores = query @ key.T / np.sqrt(d_k)
 attention = np.exp(scores) / np.sum(np.exp(scores), axis=-1, keepdims=True)
 output = attention @ value
 
 # Apply gate if provided
 if gate is not None:
 output = output * gate[:, np.newaxis]
 
 return output

np.random.seed(42)
q = np.random.randn(10, 64)
k = np.random.randn(20, 64)
v = np.random.randn(20, 64)
gate = np.random.rand(10) # Attention gate
out = gated_cross_attention(q, k, v, gate)
assert out.shape == (10, 64)
print("✓ Gated cross-attention working")

---

Go deeper with CFSGPT

Get AI-powered deep-dives, save terms, and run advanced simulations — free account.

Create Free Account