高效解包!5 个三进制位紧密放入 8 位字节,效率达 99.06% 如何用 8 位字节打包三进制数三进制数的每一位有 3 种可能的值这 3 种值可以代表任何内容。最近有人被一个问题吸引试图把 [BitNet b1.58] 的三进制权重打包使其尽可能接近每个三进制数位 log(3) / log(2) 比特的理论理想值。这里把“三进制数位”称为“三进制位trit”就像把“二进制数位”称为“二进制位bit”一样。块大小由于目标是实现快速的 **并行** 解包三进制位的块不能无限大。需要找到一个合适的“块”大小理想情况下这个大小既能保证信息密度高效又能在当前硬件上方便使用。为了找到合适的块大小要找到一个 3 的幂使其下一个 2 的幂与之非常接近。通过列表展示不同三进制位数量对应的 3 的幂、所需比特数、2 的幂以及每个三进制位的比特数1 个三进制位3 的幂为 3所需 2 比特2 的幂为 4每个三进制位 2 比特2 个三进制位3 的幂为 9所需 4 比特2 的幂为 16每个三进制位 2 比特3 个三进制位3 的幂为 27所需 5 比特2 的幂为 32每个三进制位 1.666... 比特4 个三进制位3 的幂为 81所需 7 比特2 的幂为 128每个三进制位 1.75 比特5 个三进制位3 的幂为 243所需 8 比特2 的幂为 256每个三进制位 1.6 比特非常幸运的是5 个三进制位刚好能紧密地放入 8 位字节中每个三进制位只需 1.6 比特。与完美打包相比这种方式的效率达到了 99.06%。每个三进制位 1.6 比特这种打包方案的基本思路很简单就是用三进制数位组成一个数。给出了 Python 代码示例def pack_number(digits: list[int], base: int) - int: number 0 for digit in digits: assert digit base number number * base number number digit return number将三进制位打包成字节的过程与之类似。快速乘法解包虽然可以通过反复取余和除法来提取一个数的各个数位但除法和取模操作在 SIMD 编程中通常对整数不支持。解决这个问题的方法显然是换一种方式看待数字。提出疑问如果不用取模来提取最低有效位而是用乘法来提取最高有效位是不是会更好呢答案是定点数可以解决这个问题。通过图示展示相关操作当用这个 8 位字节乘以 3 时可以很容易地从得到的 10 位数字的前两位中提取出三进制位。在使用 SIMD 解包时这种方法比取模操作方便得多。在将三进制位打包成字节的过程中唯一涉及除法的地方就在这里。这里假设打包操作的频率低于解包操作在大语言模型LLM权重的场景中确实如此。给出了将三进制位打包成字节和从字节解包出三进制位的 Python 代码示例# 接收一个包含 -1、0、1 的列表并将其打包成字节def pack_trits(digits: list[int]) - bytearray: assert len(digits) % 5 0 # 这里不处理填充 n_bytes len(digits) // 5 packed bytearray() for i in range(n_bytes): b 0 for j in range(5): digit digits[5*i j] digit max(-1, min(digit, 1)) # 限制在 -1 和 1 之间 digit 1 # 从 -1、0、1 转换为 0、1、2 b * 3 b digit b ((b * 256) (243 - 1)) // 243 packed.append(b) return packeddef unpack_trits(packed: bytes) - list[int]: trits: list[int] [] for byte in packed: b byte for i in range(5): b b * 3 trit b 8 trits.append(trit - 1) # 0、1、2 转换为 -1、0、1 b b 0xFF return trits为了验证这个方法是否可行写了一个 C 程序来检查是否真的没有信息损失并给出了 C 程序代码。编译并运行该程序对于能放入 8 位字节的 243 个三进制数都得到了 PASS 的结果。llama.cpp 中用于 TriLMs 和 BitNet b1.58 的三进制类型就采用了这种技术。提醒大家可以关注相关技术的应用和发展。