SparseGPT: Language Models Can Be Accurately Pruned in One-Shot
arxiv.org
arxiv.org
- Existing pruners were written for models that are order-of-magnitudes smaller than any in the modern GPT family. They grow in linear time with the amount of input parameters so they're unequipped to work on current architectures. The best existing pruner performs takes 4.3h for a 1.3B model
- The core issue to scale is time to calculate the Hessian during prune analysis (effectively a matrix of second-order derivatives, famously computationally intense to calculate)
- They follow the existing literature and use a local approach to each layer. By doing this (and doing it well), it can preserve the input/output contract for surrounding layers, which makes the whole thing paralellizable across machines
- Their solution approximates reconstruction loss by approximating a quadratic loss and then running a OBS update (with a few other optimizations on ordering and iteration on the side)
I'm particularly excited for these smaller models, mostly for inference efficiency gains in realtime applications. The general con of weight pruning is they still require incredibly large training clusters / investment in training resources upfront to get the original parameter weight. But if the lottery ticket hypothesis holds true, this might be the best way we have at the moment to get models with same performance and lower longterm operational costs.
It might be able to provide performance similar to fine-tuning but without the weight skew that you'll necessarily see in parameter values.
As part of that project I constructed an API that took a small dataset and a model, launched a K8s pod and ran something like this from the paper:
> The pruning defense works as follows: the defender exercises the DNN received from the attacker with clean inputs from the validation dataset, D_valid, and records the average activation of each neuron. The defender then iteratively prunes neurons from the DNN in increasing order of average activations and records the accuracy of the pruned network in each iteration. The defense terminates when the accuracy on the validation dataset drops below a pre-determined threshold. We note that pruning has been proposed in prior work for n
Obviously this wasn't on transformers but the idea is similar.
For those that, like me, didn't know the reference: https://arxiv.org/abs/1803.03635
Implementation q - can torch or other inference runtimes take advantage of the memory savings delivered by a sparsification like this? Or do you need a special implementation to not malloc out all the memory implied by each tensor layer?
Note that this is still desirable for inference because you want the most possible training on whatever model you can actually fit in your memory.
Like the sibling comment said - the proportion of training tokens to parameter size is very important, and there's a certain threshold needed to be met for it to be "fully trained".
Usually you have a fixed amount of compute (budget/time essentially) - and in that case you want to pick the largest parameter count that you can fully train, and not the largest parameter count your hardware can support and then train that for less time.
tl;dr - Small models with training over the chinchilla threshold can out perform large models that are undertrained
EDIT: Figure 2 page 5, and Table 3 page 8 - might be worth checking out.
The was, for a minute, ignored, because the PaLM paper came out very shortly thereafter which seemed to show, pretty conclusively, that there are unusual and exciting emergent behaviours coming out of much larger models, (PaLM is 540B parameters), and so that was hotter news.
In the meantime, some really smart folks looked at the Chinchilla curve, and were like "hmm. One way to think about this is to see that if you are willing to put a LOT more compute in upfront on a model, then the inference costs go down in some sub-linear function."
Llama's architectural instincts are that if you're going to give away a model, and it is going to get run on the edge, it might make sense to spend a whole, whole lot of compute, once, training something past what the paper considered optimal, and well into the point where the paper thought of it as "not worth it", precisely because the entire world might be able to run it if you can get something good and much smaller.
Conclusively, OPT and LLMs from its era are significantly 'under-trained' compared even to GPT-3, itself undertrained by something like an order of magnitude from where the Chinchilla paper implies they should be.
I guess I made up the phrase over and under-trained; their might be some other way to talk about it elsewhere. Sorry! :)
Hopefully this lowers the cost of doing instruct fine tuning on the larger models, and we see a Vicuna like model based on LLaMA 65B soon. This is exciting folks.
Quantizing/pruning/deduplicating/compressing models and embeddings is still a vast orchard of low hanging fruit.
I personally think there are still quite a few multiple-orders-of-magnitude scale opportunities to accelerate inference, and we are fortunate to have strong economic incentives aligned with the problem.
It always felt weird that we have to sleep, it doesn't seem to give any evolutionary advantages.
It may be hard to pin point exactly what advantage, but as we do it, it must have given us an advantage!
- Animals that have peaks of energy use outcompete animals that have a steady-state energy use. Catch the animal, then rest and recover. For any given amount of energy, this means we can recruit more in a smaller window compared to an animal that plods along with no recuperative phase.
- Many things happen when you're sleeping. Rather than having everything running 24/7, having different phases means we can specialise action and recovery. Since the time is already driven by energy demands, many parts of our body and mind leverage it for different purposes.
If they can develop new methods to “overtrain” these models they will get more bang out of the smaller parameter model buck.
Confirm your spot: https://neuralmagic.com/unlock-faster-and-more-efficient-lan...
> adjective: analog
> relating to or using signals or information represented by a continuously variable physical quantity such as spatial position, voltage, etc.
The expected and observed change is virtually none. That's the whole point!
Notably, quantizing weights from 16bit weights to 4bit weights (reducing the size by 75%) also has almost no change in output quality when using modern algorithms like GPTQ.
A 0.01% loss in quality for a 4x speed up and 4x less VRAM/RAM requirement.
90GB models now fit and run on a $600 consumer video card with quality so similar the difference is only detectable on hours long automated tests with tens of thousands of iterations.
> at minimal loss of accuracy
Suggesting that there is a lot of redundancy in the weights.
As some other commentator stated, there's currently a lot of low hanging fruit in optimizing NN.