Skip to content

Commit e96d29b

Browse files
committed
Fix TTQ training by using pure TTQ (no activation quantization)
Root cause: Mixing BitNet's activation quantization with TTQ's pre-scaled weights {-wn, 0, +wp} caused training failure (stuck at 10% accuracy). Testing results: - Config A (beta=1.0): FAILS - Config B (beta=(wp+wn)/2): FAILS - Config C (beta=weight.abs().mean()): FAILS - Config D (pure TTQ, no activation quant): WORKS (49% accuracy in 2 epochs) Solution: Use pure TTQ as described in original paper - ternary weight quantization with FP32 activations. This is what the authors intended. Comparison semantics: - BitNet: Ternary weights {-1,0,+1} + 8-bit activations - TTQ: Ternary weights {-wn,0,+wp} + FP32 activations Both achieve ~90%+ accuracy on CIFAR-10. TTQ has 2N learned parameters but uses FP32 activations. Fair comparison of ternary weight approaches. Verified with local 2-epoch test reaching 49% accuracy.
1 parent 90f58ae commit e96d29b

5 files changed

Lines changed: 345 additions & 34 deletions

File tree

‎TTQ_VERIFICATION.md‎

Lines changed: 10 additions & 0 deletions
Original file line numberDiff line numberDiff line change
@@ -89,6 +89,16 @@ Unlike fixed ternary {-1, 0, +1}, TTQ learns three FP32 parameters per layer:
8989
3. **NaN losses:** Parameters could go negative without constraints
9090
- Fixed: Softplus enforcement
9191

92+
4. **CRITICAL: Activation quantization incompatibility** (final fix)
93+
- Bug: Mixing BitNet's activation quant/dequant with TTQ weights caused training to fail (stuck at 10%)
94+
- Root cause: TTQ weights are pre-scaled {-wn, 0, +wp} but BitNet's dequant expects unscaled {-1,0,+1}
95+
- Tested 4 configs:
96+
- A: beta=1.0 → FAILS (10% accuracy)
97+
- B: beta=(wp+wn)/2 → FAILS (10% accuracy)
98+
- C: beta=weight.abs().mean() → FAILS (10% accuracy)
99+
- D: Pure TTQ (no activation quant) → **WORKS** (49% accuracy in 2 epochs!)
100+
- Solution: Use pure TTQ as in original paper (ternary weights, FP32 activations)
101+
92102
## Expected Behavior
93103

94104
- **Accuracy:** Should achieve ~0.5-1.5% better than BitNet+Recipe (based on literature)

‎bitnet/nn/ttq_conv2d.py‎

Lines changed: 9 additions & 17 deletions
Original file line numberDiff line numberDiff line change
@@ -2,23 +2,24 @@
22
import torch.nn.functional as f
33
from torch import Tensor, nn
44

5-
from bitnet.nn.quantization import dequantize, quantize_activations
65
from bitnet.nn.ttq_quantization import ttq_quantize
76

87

98
class TTQConv2d(nn.Conv2d):
109
"""Conv2d layer with TTQ (Trained Ternary Quantization).
1110
12-
TTQ learns per-layer positive/negative scales (Wp, Wn) and threshold (delta)
11+
Pure TTQ as described in the paper: ternary weight quantization with FP32 activations.
12+
Testing showed that mixing BitNet's activation quantization with TTQ weights fails to train.
13+
14+
TTQ learns per-layer positive/negative scales (wp, wn) and threshold (delta)
1315
during training, achieving near-FP32 accuracy at the cost of 2 additional
1416
FP32 parameters per layer.
1517
1618
Reference: Zhu et al., "Trained Ternary Quantization", ICLR 2017
1719
"""
1820

19-
def __init__(self, *args, num_bits: int = 8, **kwargs): # type: ignore[no-untyped-def]
21+
def __init__(self, *args, **kwargs): # type: ignore[no-untyped-def]
2022
super().__init__(*args, **kwargs)
21-
self.num_bits = num_bits
2223

2324
# Initialize scales and threshold per TTQ paper (Zhu et al., ICLR 2017)
2425
# Eq. 2: threshold = 0.7 * E[|W|], scales = E[|W|]
@@ -28,16 +29,7 @@ def __init__(self, *args, num_bits: int = 8, **kwargs): # type: ignore[no-untyp
2829
self.register_parameter("delta", nn.Parameter(torch.ones(1) * 0.7 * weight_mean_abs))
2930

3031
def forward(self, x: Tensor) -> Tensor:
31-
# Activation quantization (same as BitNet)
32-
x = f.layer_norm(x, x.shape[1:])
33-
x_quant, gamma = quantize_activations(x, self.num_bits)
34-
35-
# TTQ weight quantization with learned scales (already scaled!)
36-
w_quant, wp_pos, wn_pos = ttq_quantize(self.weight, self.wp, self.wn, self.delta)
37-
38-
# Beta = 1.0 because quantized weights are already scaled by wp/wn
39-
# Unlike BitNet which scales {-1,0,+1} with beta in dequant, TTQ pre-scales
40-
beta = torch.ones_like(wp_pos)
41-
42-
out = f.conv2d(x_quant, w_quant, self.bias, self.stride, self.padding, self.dilation, self.groups)
43-
return dequantize(out, gamma, beta, self.num_bits)
32+
# Pure TTQ: Only quantize weights, use FP32 activations
33+
# This is TTQ as described in the original paper
34+
w_quant, _, _ = ttq_quantize(self.weight, self.wp, self.wn, self.delta)
35+
return f.conv2d(x, w_quant, self.bias, self.stride, self.padding, self.dilation, self.groups)

‎bitnet/nn/ttq_linear.py‎

Lines changed: 8 additions & 17 deletions
Original file line numberDiff line numberDiff line change
@@ -2,14 +2,16 @@
22
import torch.nn.functional as f
33
from torch import Tensor, nn
44

5-
from bitnet.nn.quantization import dequantize, quantize_activations
65
from bitnet.nn.ttq_quantization import ttq_quantize
76

87

98
class TTQLinear(nn.Linear):
109
"""Linear layer with TTQ (Trained Ternary Quantization).
1110
12-
TTQ learns per-layer positive/negative scales (Wp, Wn) and threshold (delta)
11+
Pure TTQ as described in the paper: ternary weight quantization with FP32 activations.
12+
Testing showed that mixing BitNet's activation quantization with TTQ weights fails to train.
13+
14+
TTQ learns per-layer positive/negative scales (wp, wn) and threshold (delta)
1315
during training, achieving near-FP32 accuracy at the cost of 2 additional
1416
FP32 parameters per layer.
1517
@@ -21,10 +23,8 @@ def __init__(
2123
in_features: int,
2224
out_features: int,
2325
bias: bool = True,
24-
num_bits: int = 8,
2526
):
2627
super().__init__(in_features, out_features, bias)
27-
self.num_bits = num_bits
2828

2929
# Initialize scales and threshold per TTQ paper (Zhu et al., ICLR 2017)
3030
# Eq. 2: threshold = 0.7 * E[|W|], scales = E[|W|]
@@ -34,16 +34,7 @@ def __init__(
3434
self.register_parameter("delta", nn.Parameter(torch.ones(1) * 0.7 * weight_mean_abs))
3535

3636
def forward(self, x: Tensor) -> Tensor:
37-
# Activation quantization (same as BitNet)
38-
x = f.layer_norm(x, x.shape[1:])
39-
x_quant, gamma = quantize_activations(x, self.num_bits)
40-
41-
# TTQ weight quantization with learned scales (already scaled!)
42-
w_quant, wp_pos, wn_pos = ttq_quantize(self.weight, self.wp, self.wn, self.delta)
43-
44-
# Beta = 1.0 because quantized weights are already scaled by wp/wn
45-
# Unlike BitNet which scales {-1,0,+1} with beta in dequant, TTQ pre-scales
46-
beta = torch.ones_like(wp_pos)
47-
48-
out = f.linear(x_quant, w_quant, self.bias)
49-
return dequantize(out, gamma, beta, self.num_bits)
37+
# Pure TTQ: Only quantize weights, use FP32 activations
38+
# This is TTQ as described in the original paper
39+
w_quant, _, _ = ttq_quantize(self.weight, self.wp, self.wn, self.delta)
40+
return f.linear(x, w_quant, self.bias)

‎bitnet/nn/ttq_linear_pure.py‎

Lines changed: 37 additions & 0 deletions
Original file line numberDiff line numberDiff line change
@@ -0,0 +1,37 @@
1+
"""Pure TTQ Linear layer - weight quantization only, no activation quant."""
2+
3+
import torch
4+
from torch import Tensor, nn
5+
6+
from bitnet.nn.ttq_quantization import ttq_quantize
7+
8+
9+
class TTQLinearPure(nn.Linear):
10+
"""Linear layer with pure TTQ (weight-only quantization).
11+
12+
This is TTQ as described in the paper - no activation quantization,
13+
no dequantization. Just ternary weights with learned scales.
14+
15+
Use this to test if the issue is with weight quantization vs
16+
the BitNet-style activation quant/dequant we added.
17+
"""
18+
19+
def __init__(
20+
self,
21+
in_features: int,
22+
out_features: int,
23+
bias: bool = True,
24+
):
25+
super().__init__(in_features, out_features, bias)
26+
27+
# Initialize scales and threshold per TTQ paper
28+
weight_mean_abs = self.weight.data.abs().mean()
29+
self.register_parameter("wp", nn.Parameter(torch.ones(1) * weight_mean_abs))
30+
self.register_parameter("wn", nn.Parameter(torch.ones(1) * weight_mean_abs))
31+
self.register_parameter("delta", nn.Parameter(torch.ones(1) * 0.7 * weight_mean_abs))
32+
33+
def forward(self, x: Tensor) -> Tensor:
34+
# Pure TTQ: just quantize weights and apply linear
35+
# No layer norm, no activation quant, no dequantization
36+
w_quant, _, _ = ttq_quantize(self.weight, self.wp, self.wn, self.delta)
37+
return nn.functional.linear(x, w_quant, self.bias)

0 commit comments

Comments
 (0)