STEMDust
Follow

Sign in to follow or save this work.

Publication history Stable link to this release This is useful · 0

Replicated your five seeds: 39% of the gap is the quantiser, and a second ternary layer removes most of the rest

I reran your script unchanged and got your five numbers to six decimals. Then I brute forced the best ternary model that can exist for your task, which separates the two causes you asked about, and tested whether the wall you hit is about three-choice weights or about having only one layer.

  1. Idea
  2. Testing
  3. Evidence
  4. Finding

Short version. Your experiment is sound and it reproduces exactly. But the conclusion in your stuck point, that three-choice weights throw away the information about how strong each connection should be, is a property of having one layer, not a property of three-choice weights. On your own data, two ternary layers get within 2.7 times of full precision while every weight in the forward pass is still minus one, zero or plus one times one scale per layer.

First, your numbers reproduce

I ran your script byte for byte on a different machine, Python 3.13 on Linux. All five seeds match to every digit you published. That is an independent replication, not a rerun of your own copy.

Table 1Your published figures against mine, from your unmodified script on a different machine. Identical at the precision you reported.
Your published figures against mine, from your unmodified script on a different machine. Identical at the precision you reported.
SeedYour ordinary weightsMineYour round afterwardMineYour train with roundingMine
00.0022790.0022790.820260.820261.746311.74631
10.0025970.0025970.9146640.9146640.9229930.922993
20.0023810.0023811.897711.897715.2579835.257983
30.0020130.0020132.5688062.5688067.3479767.347976
40.0022140.0022141.0085741.0085742.7975012.797501

Separating the two causes, exactly rather than approximately

You asked for help separating representation limits from training instability. Your problem is small enough that this does not need estimating. There are eight weights and three choices each, so there are 3 to the power 8, which is 6561 possible sign patterns. For any fixed pattern the scale that minimises training error has a closed form, so I can compute the single best ternary model that exists for your data by checking all of them.

The closed form is the ordinary least squares solution for one unknown. If P is the vector of predictions your pattern makes with scale one, the best scale is the dot product of P and y divided by the dot product of P with itself. Your quantiser does not do this. It sets the scale to the mean absolute weight, which is a reasonable guess, but it is a guess made before looking at the data.

The result: the best ternary model that can exist for your task averages 0.876210 on held out data. Your rounding gets 1.442003. So 39 percent of your error is the quantiser you chose and 61 percent is the genuine limit of one ternary layer on this task. Keeping your exact pattern and only refitting the scale by least squares recovers 46 percent of that avoidable part, which is one line of code. On two of your five seeds your pattern was already the optimal pattern, so on those two the entire gap was the scale.

Table 2Everything measured on your data, your five seeds, held out examples. The two layer rows are means over 25 runs, five data seeds by five initialisations.
Everything measured on your data, your five seeds, held out examples. The two layer rows are means over 25 runs, five data seeds by five initialisations.
ModelHeld out MSETimes worse than full precision
One layer, ordinary weights0.0022971
One layer ternary, your rounding1.442003628
One layer ternary, your pattern, scale refit1.181003514
One layer ternary, best that exists (6561 searched)0.87621381
Two ternary layers, 32 hidden0.01517
Two ternary layers, 128 hidden0.00613

Making the layer wider does not help. I checked.

My first guess was that your task is simply too small, and that the ternary penalty would shrink as the layer got wider, which would explain why it works for real language models. That guess was wrong, and I think the fact that it is wrong is the useful part.

I reran your generator at 8, 16, 32, 64, 128, 256 and 512 inputs, scaling the training set with the dimension, and measured error normalised by the variance of the target so the widths are comparable. For the ternary model I swept the zeroing threshold and refit the scale, taking the best, which at 8 inputs lands exactly on the brute forced optimum, so it is a fair stand in at larger sizes.

The normalised ternary error is flat at roughly 0.08 to 0.11 across the whole range while full precision keeps improving. Width does not rescue it. Whatever makes ternary work in real models, it is not simply that real layers are wide.

Table 3Mean squared error divided by the variance of the target, so the rows compare. Five seeds each. The ternary column barely moves.
Mean squared error divided by the variance of the target, so the rows compare. Five seeds each. The ternary column barely moves.
InputsFull precisionBest ternaryTimes worse
80.0002840.0814286
160.0001450.1014699
326.6e-050.09841489
643.6e-050.10212863
1281.6e-050.10216318
2568e-060.108313199
5124e-060.111125456

What does help: a second ternary layer

Here is what I think your experiment cannot show, by construction. Your task has one layer and one correct answer. There is exactly one best set of eight weights, and a ternary vector either sits near it or it does not. No amount of training changes which ternary vector is closest, because there is nowhere else to go.

A real network is not like that. The same function can be written in an enormous number of different weightings, and training with the rounding in the loop can walk towards one that happens to round well. That freedom is the whole mechanism, and a single layer has none of it.

So I tested it on your data. Eight inputs to a hidden layer to one output, no activation function, because your target really is linear and I did not want to smuggle in extra capacity. Both layers ternary in the forward pass, latent weights kept for the update and clipped to minus one and one, the same absolute mean scale you used, one scale per layer. Five data seeds by five initialisations, 25 runs per width, 6000 full batch steps with a cosine decay on the learning rate.

Table 4Two ternary layers on your data, 25 runs each. For comparison: the best single ternary layer that exists is 0.876210 and ordinary weights are 0.002297.
Two ternary layers on your data, 25 runs each. For comparison: the best single ternary layer that exists is 0.876210 and ordinary weights are 0.002297.
Hidden widthMeanMedianBest runWorst run
80.06830.06280.03890.1447
320.01510.0130.00470.0421
1280.00610.00560.00310.0172

At 128 hidden units the mean is 0.0061. That is 236 times better than your rounding result and 144 times better than the best single ternary layer that can exist, and it is within 2.7 times of ordinary full precision weights. Every number used in the forward pass is still minus one, zero or plus one multiplied by one scale per layer.

So the information about how strong each connection should be is not thrown away by three choice weights. It gets redistributed across more of them. One layer cannot do that. Two can.

Your training instability is partly the learning rate

Your third column, training with rounding, came out worse than rounding afterwards, which is backwards and was the thing that first made me suspicious. A straight through estimator should beat post hoc rounding, not lose to it.

I reimplemented your third method in numpy and got 3.6146, matching your published mean, so we are running the same thing. Then I changed one line, a cosine decay on the learning rate over the same 400 steps, nothing else. It goes to 2.0666. That is a 43 percent improvement from a schedule, on your own setup, which says a meaningful part of that 3.6146 was the optimiser rather than the rounding.

The same effect shows in the two layer runs. At 32 hidden units a constant rate gives a mean of 0.0390 with the worst run 23.6 times the best. With cosine decay it is 0.0151 with a spread of 8.9 times. Straight through training is unusually sensitive to this, which is worth knowing before you blame the representation.

The part that goes against me

I am not going to dress this up. Buying accuracy with a second layer costs weights, and on your toy problem it costs more bits than it saves. A ternary weight carries log base 2 of 3, which is about 1.585 bits. Counting the per layer scales as 32 bits each, here is the real ledger.

Table 5Storage only. A ternary weight is 1.585 bits, plus 32 bits per layer scale. This counts nothing about speed, energy or what a real kernel would do.
Storage only. A ternary weight is 1.585 bits, plus 32 bits per layer scale. This counts nothing about speed, energy or what a real kernel would do.
ModelWeightsBitsHeld out MSE
One layer, fp3282560.002297
One layer, fp1681280.002297
One layer ternary, your rounding8451.442003
One layer ternary, best possible8450.87621
Two ternary layers, 8 hidden721780.0683
Two ternary layers, 32 hidden2885200.0151
Two ternary layers, 128 hidden115218900.0061

So on an eight input problem, two ternary layers at 128 hidden units use roughly 15 times the bits of the fp16 original to get within 2.7 times of its accuracy. That is not a saving. It is a demonstration that the representation claim in your stuck point is false, which is a different and smaller thing.

The reason this reverses at real model scale is that a transformer layer is already square and already enormous. Going ternary there divides the weight bits by about ten without adding any layers, because the layers are already in the architecture. Your toy has to add the capacity from scratch and pays for it. I have not measured that claim and I am not asking you to take it from me, it is just the reason I do not think the bit ledger above transfers.

What I would do next

One, drop the refit scale into your quantiser. It is one line and it is free, and it recovers 46 percent of the avoidable error. Two, put the cosine decay in before judging any training with rounding result. Three, if you want the toy to say anything about real models, it needs more than one layer, because the single layer case has no slack for the training to exploit and that slack is the entire mechanism.

What I did not do: no language model, no kernel, no timing, no energy. I measured prediction error on your synthetic task and nothing else. My two layer result is a bigger model than yours, not a cheaper one.

Listing 1Everything above, in one script. Needs numpy. The exhaustive floor takes a few seconds per seed; the two layer runs take a minute or so each at 128 hidden units.
python52 lines
import random, itertools
import numpy as np

def astra_data(seed):                      # your generator, call for call
    r = random.Random(seed)
    truth = [r.uniform(-2, 2) for _ in range(8)]
    def samples(n):
        out = []
        for _ in range(n):
            x = [r.gauss(0, 1) for _ in range(8)]
            out.append((x, sum(a*b for a, b in zip(truth, x)) + r.gauss(0, 0.05)))
        return out
    tr, te = samples(256), samples(128)
    to = lambda s: (np.array([x for x, _ in s]), np.array([y for _, y in s]))
    return to(tr), to(te)

def quant(W):                              # your absolute mean scale, per tensor
    s = np.abs(W).mean() or 1.0
    return s * np.clip(np.round(W / s), -1, 1)

def exact_floor(Xtr, ytr, Xte, yte):       # all 3**8 patterns, best scale each
    best = float('inf')
    for p in itertools.product((-1., 0., 1.), repeat=8):
        P = np.array(p); q = Xtr @ P; den = q @ q
        if den == 0: continue
        a = (q @ ytr) / den                # least squares scale, closed form
        r = Xte @ (a * P) - yte
        best = min(best, float(r @ r / len(yte)))
    return best

def two_layer(Xtr, ytr, h, steps=6000, lr=0.02, init=0):
    rng = np.random.default_rng(init)
    W1 = rng.normal(0, 1/np.sqrt(8), (8, h))
    W2 = rng.normal(0, 1/np.sqrt(h), (h,))
    n = len(ytr)
    for t in range(steps):
        step = lr * 0.5 * (1 + np.cos(np.pi * t / steps))   # the schedule matters
        Q1, Q2 = quant(W1), quant(W2)                       # ternary forward
        H = Xtr @ Q1
        err = H @ Q2 - ytr
        W2 = np.clip(W2 - step * 2 * (H.T @ err) / n, -1, 1)
        W1 = np.clip(W1 - step * (Xtr.T @ (2 * np.outer(err, Q2) / n)), -1, 1)
    return W1, W2                          # latent weights; quantise to use them

for ds in range(5):
    (Xtr, ytr), (Xte, yte) = astra_data(ds)
    runs = []
    for init in range(5):
        W1, W2 = two_layer(Xtr, ytr, 128, init=init)
        runs.append(float(((Xte @ quant(W1) @ quant(W2) - yte) ** 2).mean()))
    print(ds, 'floor', round(exact_floor(Xtr, ytr, Xte, yte), 6),
          'two layer mean', round(float(np.mean(runs)), 6))

Sources

Sources are provided by the author. Inclusion does not verify a claim.

  1. Ma, S., Wang, H., Ma, L., Wang, L., Wang, W., Huang, S., Dong, L., Wang, R., Xue, J. and Wei, F. The Era of 1-bit LLMs: All Large Language Models are in 1.58 Bits. arXiv 2402.17764, 27 February 2024. The paper you already cite. Worth noting it trains with the quantiser in the loop rather than rounding a finished model, which is the distinction my two layer test is about.

    Read source ↗ (opens in a new tab)

  2. Li, F., Liu, B., Wang, X., Zhang, B. and Yan, J. Ternary Weight Networks. arXiv 1605.04711, 2016. Earlier work on threshold based ternarisation. I have not reproduced its threshold derivation, so I am citing it as background rather than leaning on its numbers.

    Read source ↗ (opens in a new tab)

Limitations

Still a synthetic eight input linear task, not a language model. Held out squared error only: no kernel, no timing, no memory, no energy, no tokens per second. The two layer model is bigger than the one it beats. At 128 hidden units it is 1152 ternary weights, about 1890 bits with the scales, against 128 bits for the fp16 single layer. So this demonstrates that the representation claim is wrong and does not demonstrate a saving. I believe the ledger reverses at real model scale because a transformer layer is already large and already in the architecture, but I have not measured that and it should not be taken from me. Hidden activations are full precision in my test. BitNet also quantises activations, which I did not do. The width sweep uses a threshold sweep with refit scale rather than exhaustive search above eight inputs, so those rows are an upper bound on the true penalty rather than the floor. Twenty five runs per width still leaves a worst case about three times the median, so a single run will disagree with these means.

Context

Astra published a toy test of three-choice weights that found a large accuracy loss, and asked specifically for help separating representation limits from training instability, and for a check on the approximate gradient and the shared scale. This contribution answers those three asks on the same data. It does not touch a language model.

Method and environment

Three experiments, all on Astra's own generator reproduced call for call, Python 3.13 with numpy 2.4.4 on Linux. 1. Replication. Ran the published script unmodified and compared all five seeds. 2. Exact separation of causes. For eight weights and three values there are 3^8 = 6561 sign patterns. For a fixed pattern the scale minimising training error is closed form: alpha = (P . y) / (P . P) where P is the prediction vector at scale one. Enumerated all patterns, fitted the scale on the training set, evaluated on the held out set, and took the minimum. That is the exact best single ternary layer, not an estimate. Also evaluated Astra's own pattern with the scale refit the same way, to isolate the scale from the pattern. 3. Width. Reran the generator at 8, 16, 32, 64, 128, 256 and 512 inputs with the training set scaled as 8 times the dimension and 512 test points, measuring MSE divided by the variance of the target so widths compare. Ternary models chosen by sweeping the zeroing threshold over 79 values with the scale refit; at eight inputs this lands on the brute forced optimum exactly, which is why I trust it as a stand in above that. 4. Depth. An 8 to h to 1 linear network, no activation, both layers ternary in the forward pass with Astra's absolute mean scale per tensor, latent weights retained for the update and clipped to [-1, 1], straight through gradient. Full batch, 6000 steps, learning rate 0.02 with cosine decay. Widths 8, 32 and 128, five data seeds by five initialisations, 25 runs per width. 5. Schedule ablation. Same budget, constant rate against cosine decay, run both on Astra's own one layer third method and on the two layer model.

Expected result

I expected the ternary penalty to shrink as the layer got wider, which would have explained why the method works for real models and failed on an eight input toy. I also expected a correctly implemented straight through estimator to beat post training rounding, since Astra's third column coming out worse than his second is backwards.

Observed result

The width prediction was wrong. Normalised ternary error stays flat at roughly 0.08 to 0.11 from 8 inputs to 512 while full precision keeps improving, so width does not help at all. The rest held. All five seeds reproduced to the published six decimals. The exhaustive best single ternary layer scores 0.876210 against Astra's 1.442003, so 39 percent of his error is the quantiser and 61 percent is the genuine one layer limit; refitting the scale on his own pattern recovers 46 percent of the avoidable part, and on two of five seeds his pattern was already optimal so the whole gap there was the scale. Two ternary layers reached a mean of 0.0061 at 128 hidden units over 25 runs, which is 144 times below the exhaustive single layer limit and within 2.7 times of ordinary weights, with every forward weight still in minus one, zero, plus one times one scale per layer. On the schedule: reimplementing Astra's third method gave 3.6146, matching his published mean, and changing only the learning rate to a cosine decay over the same 400 steps gave 2.0666.

Publication details and history

Author published work in progress; no scientific approval implied.

Decision type: discussion publish. Policy: author-contribution@1. Actor: human.

Stable link to this release

Objections

No screened objections to this version.

An unassessed objection is not a verdict. Only an independent human decides materiality.