Supporting half-precision floats is really annoying (2021)
futhark-lang.org
futhark-lang.org
The details of bfloat16 are very simple: Build a standard ieee754 float32, then take only the upper 16 bits. It's a LOT easier than trying to support half-precision float, a form that has so little demand that most platforms don't even support it.
It’s also a native type on the CPU side on all Apple Silicon and iOS devices.
That seems like a them problem. Nobody needs an abstraction that isn't like /any/ of the concrete implementations below it. Especially when the ML version of half-floats is different from the GPU one.
You mean like bignums?
I would say that that’s almost exactly analogous to fp16s swapping from one representation to another as they’re moved between the GPU and a CPU that either does or doesn’t have native hardware support for them.
Another good analogy might be to x86 real-mode segmented “far” pointers, vs. protected-mode “flat” pointers, in cases where a given memory address is expressible under both representations.
Bignums are straightforwardly O(N) based on size - which is different from O(1) and can definitely lead to performance and security bugs - but it is consistent.
My understanding is that they don't train with half precision, but quantisation allows even integer values to be used for inference?
Yet the reasoning still makes sense. Can training take place at half precision or even if you intend to run inference at halfp do you want doublep for training?
You gain more by being able to run gradient descent faster than by having higher-precision floats.
I don't know much about futhark nor cuda includes, but if futhark supports cuda, wasn't this already an issue before adding f16 support, to find the headers of other main functionality of cuda? Or is f16 the first thing in cuda that needs a header and other things like regular float computations work without any includes at all?
Further, at compile-time the usual CPATH/LIBRARY_PATH environment variables are respected, which I think is not the case for the run-time NVRTC compiler.