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")---