按Transformer架构组件拆分: - utils.py: 激活函数和掩码工具 - layers.py: 基础层(Linear, LayerNorm) - attention.py: 多头注意力机制 - feedforward.py: 前馈网络 - block.py: Transformer编码器块 - model.py: SimpleTransformer模型 - inference.py: 推理引擎 - 删除旧的transformer.py文件
63 lines
1.7 KiB
Markdown
63 lines
1.7 KiB
Markdown
# Simple Transformer
|
||
|
||
A minimal Transformer implementation using NumPy with **KV Cache** support for efficient inference.
|
||
|
||
## Files
|
||
|
||
- `transformer/` - Transformer model package (按架构组件拆分)
|
||
- `__init__.py` - 包初始化,导出所有组件
|
||
- `utils.py` - 激活函数和掩码工具
|
||
- `layers.py` - 基础层(Linear, LayerNorm)
|
||
- `attention.py` - 多头注意力机制
|
||
- `feedforward.py` - 前馈网络
|
||
- `block.py` - Transformer编码器块
|
||
- `model.py` - SimpleTransformer模型
|
||
- `inference.py` - 推理引擎
|
||
- `example.py` - Usage examples (training & inference)
|
||
|
||
## Usage
|
||
|
||
```bash
|
||
python3 example.py
|
||
```
|
||
|
||
## Model Architecture
|
||
|
||
The implementation includes:
|
||
|
||
- **Multi-Head Attention**: Scaled dot-product attention with multiple heads
|
||
- **Feed-Forward Network**: Two-layer fully connected network with ReLU activation
|
||
- **Layer Normalization**: Applied after each sub-layer
|
||
- **Positional Encoding**: Sinusoidal position embeddings
|
||
- **KV Cache**: Efficient inference with prefill/decode separation
|
||
|
||
## Inference: Prefill vs Decode
|
||
|
||
### Prefill Stage
|
||
- Process the complete input prompt in parallel
|
||
- Initialize KV cache for all layers
|
||
- Generate the first token
|
||
- Computation: O(n²)
|
||
|
||
### Decode Stage
|
||
- Process one token at a time
|
||
- Reuse cached Key and Value tensors
|
||
- Only compute new Query
|
||
- Computation: O(n) per step
|
||
|
||
## Model Parameters
|
||
|
||
Default configuration:
|
||
- Vocabulary size: 1000
|
||
- Model dimension: 512
|
||
- Number of heads: 8
|
||
- Number of layers: 6
|
||
- Feed-forward dimension: 2048
|
||
- Max sequence length: 100
|
||
|
||
Total parameters: ~19.4M
|
||
|
||
## Classes
|
||
|
||
- `SimpleTransformer` - Transformer model with prefill/decode methods
|
||
- `InferenceEngine` - Manages inference process with KV cache |