|
| 1 | +# TTQ Implementation Verification |
| 2 | + |
| 3 | +## Paper Reference |
| 4 | +**Trained Ternary Quantization** (Zhu et al., ICLR 2017) |
| 5 | +arXiv: https://arxiv.org/abs/1612.01064 |
| 6 | + |
| 7 | +## Algorithm Summary (from paper) |
| 8 | + |
| 9 | +### Forward Pass |
| 10 | + |
| 11 | +**Quantization (Eq. 1):** |
| 12 | +``` |
| 13 | +W_t = { +Wp if W > delta |
| 14 | + { -Wn if W < -delta |
| 15 | + { 0 otherwise |
| 16 | +``` |
| 17 | + |
| 18 | +**Initialization (Eq. 2, Section 3.1):** |
| 19 | +- Threshold: `delta = 0.7 * E[|W|]` |
| 20 | +- Positive scale: `Wp = E[|W|]` |
| 21 | +- Negative scale: `Wn = E[|W|]` |
| 22 | + |
| 23 | +Where `E[|W|]` = mean of absolute weight values |
| 24 | + |
| 25 | +### Backward Pass |
| 26 | +Straight-through estimator (STE): gradients flow as if no quantization |
| 27 | + |
| 28 | +### Key Innovation |
| 29 | +Unlike fixed ternary {-1, 0, +1}, TTQ learns three FP32 parameters per layer: |
| 30 | +- `Wp` (positive scale) |
| 31 | +- `Wn` (negative scale) |
| 32 | +- `delta` (threshold) |
| 33 | + |
| 34 | +## Our Implementation |
| 35 | + |
| 36 | +### Files |
| 37 | +- `bitnet/nn/ttq_quantization.py` - Core quantization function |
| 38 | +- `bitnet/nn/ttq_linear.py` - Linear layer with TTQ |
| 39 | +- `bitnet/nn/ttq_conv2d.py` - Conv2d layer with TTQ |
| 40 | +- `tests/test_ttq_layers.py` - Test suite |
| 41 | + |
| 42 | +### Key Design Decisions |
| 43 | + |
| 44 | +**1. Positivity Constraint:** |
| 45 | +- Paper assumes Wp, Wn, delta > 0 but doesn't specify enforcement |
| 46 | +- We use `F.softplus()` to ensure positivity while maintaining gradients |
| 47 | +- Returns tuple `(quantized, wp_pos, wn_pos)` for consistent scaling |
| 48 | + |
| 49 | +**2. Activation Quantization:** |
| 50 | +- TTQ paper only specifies weight quantization |
| 51 | +- We use BitNet's activation quantization (`quantize_activations` + `dequantize`) |
| 52 | +- This allows fair comparison: both methods quantize weights AND activations |
| 53 | +- Beta for dequantization: `beta = (wp_pos + wn_pos) / 2` |
| 54 | + |
| 55 | +**3. Initialization:** |
| 56 | +- Wp, Wn = `mean(abs(weight))` ✓ Matches paper |
| 57 | +- delta = `0.7 * mean(abs(weight))` ✓ Matches paper |
| 58 | + |
| 59 | +## Verification Checklist |
| 60 | + |
| 61 | +- [x] Quantization logic matches Eq. 1 |
| 62 | +- [x] Threshold comparison: `W > delta` and `W < -delta` |
| 63 | +- [x] Three learnable parameters: wp, wn, delta |
| 64 | +- [x] Initialization: Wp = Wn = E[|W|] |
| 65 | +- [x] Initialization: delta = 0.7 * E[|W|] |
| 66 | +- [x] Straight-through estimator for gradients |
| 67 | +- [x] Positivity enforcement (softplus) |
| 68 | +- [x] Consistent scale usage in quantization and dequantization |
| 69 | +- [x] Test suite covers shapes, gradients, initialization, stability |
| 70 | + |
| 71 | +## Differences from Pure TTQ |
| 72 | + |
| 73 | +1. **Activation Quantization:** We add BitNet-style activation quantization (8-bit) |
| 74 | + - Reason: Fair comparison (both methods quantize weights + activations) |
| 75 | + - Impact: More realistic for deployment |
| 76 | + |
| 77 | +2. **Positivity Enforcement:** We use softplus, paper doesn't specify |
| 78 | + - Reason: Prevent training instability from negative scales |
| 79 | + - Impact: Minor, gradients still flow |
| 80 | + |
| 81 | +## Bugs Fixed |
| 82 | + |
| 83 | +1. **Double softplus application:** quantization used softplus(wp), dequantization used softplus(softplus(wp)) |
| 84 | + - Fixed: Return wp_pos, wn_pos from ttq_quantize |
| 85 | + |
| 86 | +2. **Wrong initialization:** Used std(W) instead of mean(|W|) for delta |
| 87 | + - Fixed: Both use `weight.abs().mean()` |
| 88 | + |
| 89 | +3. **NaN losses:** Parameters could go negative without constraints |
| 90 | + - Fixed: Softplus enforcement |
| 91 | + |
| 92 | +## Expected Behavior |
| 93 | + |
| 94 | +- **Accuracy:** Should achieve ~0.5-1.5% better than BitNet+Recipe (based on literature) |
| 95 | +- **Complexity:** Requires 2 FP32 params per layer (vs BitNet+Recipe's 1 FP32 layer) |
| 96 | +- **Trade-off:** Better accuracy, more deployment complexity |
| 97 | + |
| 98 | +## Test Results |
| 99 | + |
| 100 | +All 9 tests pass: |
| 101 | +- Forward pass shapes |
| 102 | +- Gradient flow (wp, wn get gradients) |
| 103 | +- Correct initialization (E[|W|] and 0.7*E[|W|]) |
| 104 | +- Numerical stability (no NaN in 10 training steps) |
| 105 | +- Various kernel sizes |
0 commit comments