SDNQ add 15, 13, 11 and 9 bit support

This commit is contained in:
Disty0
2026-03-11 03:39:02 +03:00
parent e5f9707e48
commit d9e628574a
15 changed files with 867 additions and 455 deletions
+73 -3
View File
@@ -13,9 +13,14 @@ dtype_dict = {
"int16": {"min": -32768, "max": 32767, "num_bits": 16, "sign": 1, "exponent": 0, "mantissa": 15, "target_dtype": torch.int16, "torch_dtype": torch.int16, "storage_dtype": torch.int16, "is_unsigned": False, "is_integer": True, "is_packed": False},
"int8": {"min": -128, "max": 127, "num_bits": 8, "sign": 1, "exponent": 0, "mantissa": 7, "target_dtype": torch.int8, "torch_dtype": torch.int8, "storage_dtype": torch.int8, "is_unsigned": False, "is_integer": True, "is_packed": False},
### Custom Integers
"int15": {"min": -16384, "max": 16383, "num_bits": 15, "sign": 1, "exponent": 0, "mantissa": 14, "target_dtype": "int15", "torch_dtype": torch.int16, "storage_dtype": torch.int16, "is_unsigned": False, "is_integer": True, "is_packed": True},
"int14": {"min": -8192, "max": 8191, "num_bits": 14, "sign": 1, "exponent": 0, "mantissa": 13, "target_dtype": "int14", "torch_dtype": torch.int16, "storage_dtype": torch.int16, "is_unsigned": False, "is_integer": True, "is_packed": True},
"int13": {"min": -4096, "max": 4095, "num_bits": 13, "sign": 1, "exponent": 0, "mantissa": 12, "target_dtype": "int13", "torch_dtype": torch.int16, "storage_dtype": torch.int16, "is_unsigned": False, "is_integer": True, "is_packed": True},
"int12": {"min": -2048, "max": 2047, "num_bits": 12, "sign": 1, "exponent": 0, "mantissa": 11, "target_dtype": "int12", "torch_dtype": torch.int16, "storage_dtype": torch.int16, "is_unsigned": False, "is_integer": True, "is_packed": True},
"int11": {"min": -1024, "max": 1023, "num_bits": 11, "sign": 1, "exponent": 0, "mantissa": 10, "target_dtype": "int11", "torch_dtype": torch.int16, "storage_dtype": torch.int16, "is_unsigned": False, "is_integer": True, "is_packed": True},
"int10": {"min": -512, "max": 511, "num_bits": 10, "sign": 1, "exponent": 0, "mantissa": 9, "target_dtype": "int10", "torch_dtype": torch.int16, "storage_dtype": torch.int16, "is_unsigned": False, "is_integer": True, "is_packed": True},
"int9": {"min": -256, "max": 255, "num_bits": 9, "sign": 1, "exponent": 0, "mantissa": 8, "target_dtype": "int9", "torch_dtype": torch.int16, "storage_dtype": torch.int16, "is_unsigned": False, "is_integer": True, "is_packed": True},
#
"int7": {"min": -64, "max": 63, "num_bits": 7, "sign": 1, "exponent": 0, "mantissa": 6, "target_dtype": "int7", "torch_dtype": torch.int8, "storage_dtype": torch.uint8, "is_unsigned": False, "is_integer": True, "is_packed": True},
"int6": {"min": -32, "max": 31, "num_bits": 6, "sign": 1, "exponent": 0, "mantissa": 5, "target_dtype": "int6", "torch_dtype": torch.int8, "storage_dtype": torch.uint8, "is_unsigned": False, "is_integer": True, "is_packed": True},
"int5": {"min": -16, "max": 15, "num_bits": 5, "sign": 1, "exponent": 0, "mantissa": 4, "target_dtype": "int5", "torch_dtype": torch.int8, "storage_dtype": torch.uint8, "is_unsigned": False, "is_integer": True, "is_packed": True},
@@ -27,9 +32,14 @@ dtype_dict = {
"uint16": {"min": 0, "max": 65535, "num_bits": 16, "sign": 0, "exponent": 0, "mantissa": 16, "target_dtype": torch.uint16, "torch_dtype": torch.uint16, "storage_dtype": torch.uint16, "is_unsigned": True, "is_integer": True, "is_packed": False},
"uint8": {"min": 0, "max": 255, "num_bits": 8, "sign": 0, "exponent": 0, "mantissa": 8, "target_dtype": torch.uint8, "torch_dtype": torch.uint8, "storage_dtype": torch.uint8, "is_unsigned": True, "is_integer": True, "is_packed": False},
### Custom Unsigned Integers
"uint15": {"min": 0, "max": 32768, "num_bits": 15, "sign": 0, "exponent": 0, "mantissa": 15, "target_dtype": "uint15", "torch_dtype": torch.int16, "storage_dtype": torch.int16, "is_unsigned": True, "is_integer": True, "is_packed": True},
"uint14": {"min": 0, "max": 16384, "num_bits": 14, "sign": 0, "exponent": 0, "mantissa": 14, "target_dtype": "uint14", "torch_dtype": torch.int16, "storage_dtype": torch.int16, "is_unsigned": True, "is_integer": True, "is_packed": True},
"uint13": {"min": 0, "max": 8192, "num_bits": 13, "sign": 0, "exponent": 0, "mantissa": 13, "target_dtype": "uint13", "torch_dtype": torch.int16, "storage_dtype": torch.int16, "is_unsigned": True, "is_integer": True, "is_packed": True},
"uint12": {"min": 0, "max": 4096, "num_bits": 12, "sign": 0, "exponent": 0, "mantissa": 12, "target_dtype": "uint12", "torch_dtype": torch.int16, "storage_dtype": torch.int16, "is_unsigned": True, "is_integer": True, "is_packed": True},
"uint11": {"min": 0, "max": 2048, "num_bits": 11, "sign": 0, "exponent": 0, "mantissa": 11, "target_dtype": "uint11", "torch_dtype": torch.int16, "storage_dtype": torch.int16, "is_unsigned": True, "is_integer": True, "is_packed": True},
"uint10": {"min": 0, "max": 1024, "num_bits": 10, "sign": 0, "exponent": 0, "mantissa": 10, "target_dtype": "uint10", "torch_dtype": torch.int16, "storage_dtype": torch.int16, "is_unsigned": True, "is_integer": True, "is_packed": True},
"uint9": {"min": 0, "max": 512, "num_bits": 9, "sign": 0, "exponent": 0, "mantissa": 9, "target_dtype": "uint9", "torch_dtype": torch.int16, "storage_dtype": torch.int16, "is_unsigned": True, "is_integer": True, "is_packed": True},
#
"uint7": {"min": 0, "max": 127, "num_bits": 7, "sign": 0, "exponent": 0, "mantissa": 7, "target_dtype": "uint7", "torch_dtype": torch.uint8, "storage_dtype": torch.uint8, "is_unsigned": True, "is_integer": True, "is_packed": True},
"uint6": {"min": 0, "max": 63, "num_bits": 6, "sign": 0, "exponent": 0, "mantissa": 6, "target_dtype": "uint6", "torch_dtype": torch.uint8, "storage_dtype": torch.uint8, "is_unsigned": True, "is_integer": True, "is_packed": True},
"uint5": {"min": 0, "max": 31, "num_bits": 5, "sign": 0, "exponent": 0, "mantissa": 5, "target_dtype": "uint5", "torch_dtype": torch.uint8, "storage_dtype": torch.uint8, "is_unsigned": True, "is_integer": True, "is_packed": True},
@@ -50,24 +60,48 @@ dtype_dict = {
"float16_e4m11fn": {"min": -511.875, "max": 511.875, "num_bits": 16, "sign": 1, "exponent": 4, "mantissa": 11, "min_normal": 0.007816314697265625, "target_dtype": "fp16", "torch_dtype": torch.float32, "storage_dtype": torch.uint16, "is_unsigned": False, "is_integer": False, "is_packed": True},
"float16_e5m10fn": {"min": -131008.0, "max": 131008.0, "num_bits": 16, "sign": 1, "exponent": 5, "mantissa": 10, "min_normal": 3.0547380447387695e-05, "target_dtype": "fp16", "torch_dtype": torch.float32, "storage_dtype": torch.uint16, "is_unsigned": False, "is_integer": False, "is_packed": True},
#
"float15_e1m13fn": {"min": -3.999755859375, "max": 3.999755859375, "num_bits": 15, "sign": 1, "exponent": 1, "mantissa": 13, "min_normal": 1.0001220703125, "target_dtype": "fp15", "torch_dtype": torch.float32, "storage_dtype": torch.int16, "is_unsigned": False, "is_integer": False, "is_packed": True},
"float15_e2m12fn": {"min": -7.9990234375, "max": 7.9990234375, "num_bits": 15, "sign": 1, "exponent": 2, "mantissa": 12, "min_normal": 0.5001220703125, "target_dtype": "fp15", "torch_dtype": torch.float32, "storage_dtype": torch.int16, "is_unsigned": False, "is_integer": False, "is_packed": True},
"float15_e3m11fn": {"min": -31.9921875, "max": 31.9921875, "num_bits": 15, "sign": 1, "exponent": 3, "mantissa": 11, "min_normal": 0.12506103515625, "target_dtype": "fp15", "torch_dtype": torch.float32, "storage_dtype": torch.int16, "is_unsigned": False, "is_integer": False, "is_packed": True},
"float15_e4m10fn": {"min": -511.75, "max": 511.75, "num_bits": 15, "sign": 1, "exponent": 4, "mantissa": 10, "min_normal": 0.00782012939453125, "target_dtype": "fp15", "torch_dtype": torch.float32, "storage_dtype": torch.int16, "is_unsigned": False, "is_integer": False, "is_packed": True},
"float15_e5m9fn": {"min": -130944.0, "max": 130944.0, "num_bits": 15, "sign": 1, "exponent": 5, "mantissa": 9, "min_normal": 3.057718276977539e-05, "target_dtype": "fp15", "torch_dtype": torch.float32, "storage_dtype": torch.int16, "is_unsigned": False, "is_integer": False, "is_packed": True},
#
"float14_e1m12fn": {"min": -3.99951171875, "max": 3.99951171875, "num_bits": 14, "sign": 1, "exponent": 1, "mantissa": 12, "min_normal": 1.000244140625, "target_dtype": "fp14", "torch_dtype": torch.float32, "storage_dtype": torch.int16, "is_unsigned": False, "is_integer": False, "is_packed": True},
"float14_e2m11fn": {"min": -7.998046875, "max": 7.998046875, "num_bits": 14, "sign": 1, "exponent": 2, "mantissa": 11, "min_normal": 0.500244140625, "target_dtype": "fp14", "torch_dtype": torch.float32, "storage_dtype": torch.int16, "is_unsigned": False, "is_integer": False, "is_packed": True},
"float14_e3m10fn": {"min": -31.984375, "max": 31.984375, "num_bits": 14, "sign": 1, "exponent": 3, "mantissa": 10, "min_normal": 0.1251220703125, "target_dtype": "fp14", "torch_dtype": torch.float32, "storage_dtype": torch.int16, "is_unsigned": False, "is_integer": False, "is_packed": True},
"float14_e4m9fn": {"min": -511.5, "max": 511.5, "num_bits": 14, "sign": 1, "exponent": 4, "mantissa": 9, "min_normal": 0.0078277587890625, "target_dtype": "fp14", "torch_dtype": torch.float32, "storage_dtype": torch.int16, "is_unsigned": False, "is_integer": False, "is_packed": True},
"float14_e5m8fn": {"min": -130816.0, "max": 130816.0, "num_bits": 14, "sign": 1, "exponent": 5, "mantissa": 8, "min_normal": 3.063678741455078e-05, "target_dtype": "fp14", "torch_dtype": torch.float32, "storage_dtype": torch.int16, "is_unsigned": False, "is_integer": False, "is_packed": True},
#
"float13_e1m11fn": {"min": -3.9990234375, "max": 3.9990234375, "num_bits": 13, "sign": 1, "exponent": 1, "mantissa": 11, "min_normal": 1.00048828125, "target_dtype": "fp13", "torch_dtype": torch.float32, "storage_dtype": torch.int16, "is_unsigned": False, "is_integer": False, "is_packed": True},
"float13_e2m10fn": {"min": -7.99609375, "max": 7.99609375, "num_bits": 13, "sign": 1, "exponent": 2, "mantissa": 10, "min_normal": 0.50048828125, "target_dtype": "fp13", "torch_dtype": torch.float32, "storage_dtype": torch.int16, "is_unsigned": False, "is_integer": False, "is_packed": True},
"float13_e3m9fn": {"min": -31.96875, "max": 31.96875, "num_bits": 13, "sign": 1, "exponent": 3, "mantissa": 9, "min_normal": 0.125244140625, "target_dtype": "fp13", "torch_dtype": torch.float32, "storage_dtype": torch.int16, "is_unsigned": False, "is_integer": False, "is_packed": True},
"float13_e4m8fn": {"min": -511.0, "max": 511.0, "num_bits": 13, "sign": 1, "exponent": 4, "mantissa": 8, "min_normal": 0.007843017578125, "target_dtype": "fp13", "torch_dtype": torch.float32, "storage_dtype": torch.int16, "is_unsigned": False, "is_integer": False, "is_packed": True},
"float13_e5m7fn": {"min": -130560.0, "max": 130560.0, "num_bits": 13, "sign": 1, "exponent": 5, "mantissa": 7, "min_normal": 3.075599670410156e-05, "target_dtype": "fp13", "torch_dtype": torch.float32, "storage_dtype": torch.int16, "is_unsigned": False, "is_integer": False, "is_packed": True},
#
"float12_e1m10fn": {"min": -3.998046875, "max": 3.998046875, "num_bits": 12, "sign": 1, "exponent": 1, "mantissa": 10, "min_normal": 1.0009765625, "target_dtype": "fp12", "torch_dtype": torch.float32, "storage_dtype": torch.int16, "is_unsigned": False, "is_integer": False, "is_packed": True},
"float12_e2m9fn": {"min": -7.9921875, "max": 7.9921875, "num_bits": 12, "sign": 1, "exponent": 2, "mantissa": 9, "min_normal": 0.5009765625, "target_dtype": "fp12", "torch_dtype": torch.float32, "storage_dtype": torch.int16, "is_unsigned": False, "is_integer": False, "is_packed": True},
"float12_e3m8fn": {"min": -31.9375, "max": 31.9375, "num_bits": 12, "sign": 1, "exponent": 3, "mantissa": 8, "min_normal": 0.12548828125, "target_dtype": "fp12", "torch_dtype": torch.float32, "storage_dtype": torch.int16, "is_unsigned": False, "is_integer": False, "is_packed": True},
"float12_e4m7fn": {"min": -510.0, "max": 510.0, "num_bits": 12, "sign": 1, "exponent": 4, "mantissa": 7, "min_normal": 0.00787353515625, "target_dtype": "fp12", "torch_dtype": torch.float32, "storage_dtype": torch.int16, "is_unsigned": False, "is_integer": False, "is_packed": True},
"float12_e5m6fn": {"min": -130048.0, "max": 130048.0, "num_bits": 12, "sign": 1, "exponent": 5, "mantissa": 6, "min_normal": 3.0994415283203125e-05, "target_dtype": "fp12", "torch_dtype": torch.float32, "storage_dtype": torch.int16, "is_unsigned": False, "is_integer": False, "is_packed": True},
#
"float11_e1m9fn": {"min": -3.99609375, "max": 3.99609375, "num_bits": 11, "sign": 1, "exponent": 1, "mantissa": 9, "min_normal": 1.001953125, "target_dtype": "fp11", "torch_dtype": torch.float32, "storage_dtype": torch.int16, "is_unsigned": False, "is_integer": False, "is_packed": True},
"float11_e2m8fn": {"min": -7.984375, "max": 7.984375, "num_bits": 11, "sign": 1, "exponent": 2, "mantissa": 8, "min_normal": 0.501953125, "target_dtype": "fp11", "torch_dtype": torch.float32, "storage_dtype": torch.int16, "is_unsigned": False, "is_integer": False, "is_packed": True},
"float11_e3m7fn": {"min": -31.875, "max": 31.875, "num_bits": 11, "sign": 1, "exponent": 3, "mantissa": 7, "min_normal": 0.1259765625, "target_dtype": "fp11", "torch_dtype": torch.float32, "storage_dtype": torch.int16, "is_unsigned": False, "is_integer": False, "is_packed": True},
"float11_e4m6fn": {"min": -508.0, "max": 508.0, "num_bits": 11, "sign": 1, "exponent": 4, "mantissa": 6, "min_normal": 0.0079345703125, "target_dtype": "fp11", "torch_dtype": torch.float32, "storage_dtype": torch.int16, "is_unsigned": False, "is_integer": False, "is_packed": True},
"float11_e5m5fn": {"min": -129024.0, "max": 129024.0, "num_bits": 11, "sign": 1, "exponent": 5, "mantissa": 5, "min_normal": 3.147125244140625e-05, "target_dtype": "fp11", "torch_dtype": torch.float32, "storage_dtype": torch.int16, "is_unsigned": False, "is_integer": False, "is_packed": True},
#
"float10_e1m8fn": {"min": -3.9921875, "max": 3.9921875, "num_bits": 10, "sign": 1, "exponent": 1, "mantissa": 8, "min_normal": 1.00390625, "target_dtype": "fp10", "torch_dtype": torch.float32, "storage_dtype": torch.int16, "is_unsigned": False, "is_integer": False, "is_packed": True},
"float10_e2m7fn": {"min": -7.96875, "max": 7.96875, "num_bits": 10, "sign": 1, "exponent": 2, "mantissa": 7, "min_normal": 0.50390625, "target_dtype": "fp10", "torch_dtype": torch.float32, "storage_dtype": torch.int16, "is_unsigned": False, "is_integer": False, "is_packed": True},
"float10_e3m6fn": {"min": -31.75, "max": 31.75, "num_bits": 10, "sign": 1, "exponent": 3, "mantissa": 6, "min_normal": 0.126953125, "target_dtype": "fp10", "torch_dtype": torch.float32, "storage_dtype": torch.int16, "is_unsigned": False, "is_integer": False, "is_packed": True},
"float10_e4m5fn": {"min": -504.0, "max": 504.0, "num_bits": 10, "sign": 1, "exponent": 4, "mantissa": 5, "min_normal": 0.008056640625, "target_dtype": "fp10", "torch_dtype": torch.float32, "storage_dtype": torch.int16, "is_unsigned": False, "is_integer": False, "is_packed": True},
"float10_e5m4fn": {"min": -126976.0, "max": 126976.0, "num_bits": 10, "sign": 1, "exponent": 5, "mantissa": 4, "min_normal": 3.24249267578125e-05, "target_dtype": "fp10", "torch_dtype": torch.float32, "storage_dtype": torch.int16, "is_unsigned": False, "is_integer": False, "is_packed": True},
#
"float9_e1m7fn": {"min": -3.984375, "max": 3.984375, "num_bits": 9, "sign": 1, "exponent": 1, "mantissa": 7, "min_normal": 1.0078125, "target_dtype": "fp9", "torch_dtype": torch.float32, "storage_dtype": torch.int16, "is_unsigned": False, "is_integer": False, "is_packed": True},
"float9_e2m6fn": {"min": -7.9375, "max": 7.9375, "num_bits": 9, "sign": 1, "exponent": 2, "mantissa": 6, "min_normal": 0.5078125, "target_dtype": "fp9", "torch_dtype": torch.float32, "storage_dtype": torch.int16, "is_unsigned": False, "is_integer": False, "is_packed": True},
"float9_e3m5fn": {"min": -31.5, "max": 31.5, "num_bits": 9, "sign": 1, "exponent": 3, "mantissa": 5, "min_normal": 0.12890625, "target_dtype": "fp9", "torch_dtype": torch.float32, "storage_dtype": torch.int16, "is_unsigned": False, "is_integer": False, "is_packed": True},
"float9_e4m4fn": {"min": -496.0, "max": 496.0, "num_bits": 9, "sign": 1, "exponent": 4, "mantissa": 4, "min_normal": 0.00830078125, "target_dtype": "fp9", "torch_dtype": torch.float32, "storage_dtype": torch.int16, "is_unsigned": False, "is_integer": False, "is_packed": True},
"float9_e5m3fn": {"min": -122880.0, "max": 122880.0, "num_bits": 9, "sign": 1, "exponent": 5, "mantissa": 3, "min_normal": 3.4332275390625e-05, "target_dtype": "fp9", "torch_dtype": torch.float32, "storage_dtype": torch.int16, "is_unsigned": False, "is_integer": False, "is_packed": True},
#
"float8_e1m6fn": {"min": -3.96875, "max": 3.96875, "num_bits": 8, "sign": 1, "exponent": 1, "mantissa": 6, "min_normal": 1.015625, "target_dtype": "fp8", "torch_dtype": torch.float32, "storage_dtype": torch.uint8, "is_unsigned": False, "is_integer": False, "is_packed": True},
"float8_e2m5fn": {"min": -7.875, "max": 7.875, "num_bits": 8, "sign": 1, "exponent": 2, "mantissa": 5, "min_normal": 0.515625, "target_dtype": "fp8", "torch_dtype": torch.float32, "storage_dtype": torch.uint8, "is_unsigned": False, "is_integer": False, "is_packed": True},
"float8_e3m4fn": {"min": -31.0, "max": 31.0, "num_bits": 8, "sign": 1, "exponent": 3, "mantissa": 4, "min_normal": 0.1328125, "target_dtype": "fp8", "torch_dtype": torch.float32, "storage_dtype": torch.uint8, "is_unsigned": False, "is_integer": False, "is_packed": True},
@@ -106,24 +140,48 @@ dtype_dict = {
"float16_e4m12fnu": {"min": 0, "max": 511.9375, "num_bits": 16, "sign": 0, "exponent": 4, "mantissa": 12, "min_normal": 0.007814407348632812, "target_dtype": "fp16", "torch_dtype": torch.float32, "storage_dtype": torch.uint16, "is_unsigned": True, "is_integer": False, "is_packed": True},
"float16_e5m11fnu": {"min": 0, "max": 131040.0, "num_bits": 16, "sign": 0, "exponent": 5, "mantissa": 11, "min_normal": 3.053247928619385e-05, "target_dtype": "fp16", "torch_dtype": torch.float32, "storage_dtype": torch.uint16, "is_unsigned": True, "is_integer": False, "is_packed": True},
#
"float15_e1m14fnu": {"min": 0, "max": 3.9998779296875, "num_bits": 15, "sign": 0, "exponent": 1, "mantissa": 14, "min_normal": 1.00006103515625, "target_dtype": "fp15", "torch_dtype": torch.float32, "storage_dtype": torch.int16, "is_unsigned": True, "is_integer": False, "is_packed": True},
"float15_e2m13fnu": {"min": 0, "max": 7.99951171875, "num_bits": 15, "sign": 0, "exponent": 2, "mantissa": 13, "min_normal": 0.50006103515625, "target_dtype": "fp15", "torch_dtype": torch.float32, "storage_dtype": torch.int16, "is_unsigned": True, "is_integer": False, "is_packed": True},
"float15_e3m12fnu": {"min": 0, "max": 31.99609375, "num_bits": 15, "sign": 0, "exponent": 3, "mantissa": 12, "min_normal": 0.125030517578125, "target_dtype": "fp15", "torch_dtype": torch.float32, "storage_dtype": torch.int16, "is_unsigned": True, "is_integer": False, "is_packed": True},
"float15_e4m11fnu": {"min": 0, "max": 511.875, "num_bits": 15, "sign": 0, "exponent": 4, "mantissa": 11, "min_normal": 0.007816314697265625, "target_dtype": "fp15", "torch_dtype": torch.float32, "storage_dtype": torch.int16, "is_unsigned": True, "is_integer": False, "is_packed": True},
"float15_e5m10fnu": {"min": 0, "max": 131008.0, "num_bits": 15, "sign": 0, "exponent": 5, "mantissa": 10, "min_normal": 3.0547380447387695e-05, "target_dtype": "fp15", "torch_dtype": torch.float32, "storage_dtype": torch.int16, "is_unsigned": True, "is_integer": False, "is_packed": True},
#
"float14_e1m13fnu": {"min": 0, "max": 3.999755859375, "num_bits": 14, "sign": 0, "exponent": 1, "mantissa": 13, "min_normal": 1.0001220703125, "target_dtype": "fp14", "torch_dtype": torch.float32, "storage_dtype": torch.int16, "is_unsigned": True, "is_integer": False, "is_packed": True},
"float14_e2m12fnu": {"min": 0, "max": 7.9990234375, "num_bits": 14, "sign": 0, "exponent": 2, "mantissa": 12, "min_normal": 0.5001220703125, "target_dtype": "fp14", "torch_dtype": torch.float32, "storage_dtype": torch.int16, "is_unsigned": True, "is_integer": False, "is_packed": True},
"float14_e3m11fnu": {"min": 0, "max": 31.9921875, "num_bits": 14, "sign": 0, "exponent": 3, "mantissa": 11, "min_normal": 0.12506103515625, "target_dtype": "fp14", "torch_dtype": torch.float32, "storage_dtype": torch.int16, "is_unsigned": True, "is_integer": False, "is_packed": True},
"float14_e4m10fnu": {"min": 0, "max": 511.75, "num_bits": 14, "sign": 0, "exponent": 4, "mantissa": 10, "min_normal": 0.00782012939453125, "target_dtype": "fp14", "torch_dtype": torch.float32, "storage_dtype": torch.int16, "is_unsigned": True, "is_integer": False, "is_packed": True},
"float14_e5m9fnu": {"min": 0, "max": 130944.0, "num_bits": 14, "sign": 0, "exponent": 5, "mantissa": 9, "min_normal": 3.057718276977539e-05, "target_dtype": "fp14", "torch_dtype": torch.float32, "storage_dtype": torch.int16, "is_unsigned": True, "is_integer": False, "is_packed": True},
#
"float13_e1m12fnu": {"min": 0, "max": 3.99951171875, "num_bits": 13, "sign": 0, "exponent": 1, "mantissa": 12, "min_normal": 1.000244140625, "target_dtype": "fp13", "torch_dtype": torch.float32, "storage_dtype": torch.int16, "is_unsigned": True, "is_integer": False, "is_packed": True},
"float13_e2m11fnu": {"min": 0, "max": 7.998046875, "num_bits": 13, "sign": 0, "exponent": 2, "mantissa": 11, "min_normal": 0.500244140625, "target_dtype": "fp13", "torch_dtype": torch.float32, "storage_dtype": torch.int16, "is_unsigned": True, "is_integer": False, "is_packed": True},
"float13_e3m10fnu": {"min": 0, "max": 31.984375, "num_bits": 13, "sign": 0, "exponent": 3, "mantissa": 10, "min_normal": 0.1251220703125, "target_dtype": "fp13", "torch_dtype": torch.float32, "storage_dtype": torch.int16, "is_unsigned": True, "is_integer": False, "is_packed": True},
"float13_e4m9fnu": {"min": 0, "max": 511.5, "num_bits": 13, "sign": 0, "exponent": 4, "mantissa": 9, "min_normal": 0.0078277587890625, "target_dtype": "fp13", "torch_dtype": torch.float32, "storage_dtype": torch.int16, "is_unsigned": True, "is_integer": False, "is_packed": True},
"float13_e5m8fnu": {"min": 0, "max": 130816.0, "num_bits": 13, "sign": 0, "exponent": 5, "mantissa": 8, "min_normal": 3.063678741455078e-05, "target_dtype": "fp13", "torch_dtype": torch.float32, "storage_dtype": torch.int16, "is_unsigned": True, "is_integer": False, "is_packed": True},
#
"float12_e1m11fnu": {"min": 0, "max": 3.9990234375, "num_bits": 12, "sign": 0, "exponent": 1, "mantissa": 11, "min_normal": 1.00048828125, "target_dtype": "fp12", "torch_dtype": torch.float32, "storage_dtype": torch.int16, "is_unsigned": True, "is_integer": False, "is_packed": True},
"float12_e2m10fnu": {"min": 0, "max": 7.99609375, "num_bits": 12, "sign": 0, "exponent": 2, "mantissa": 10, "min_normal": 0.50048828125, "target_dtype": "fp12", "torch_dtype": torch.float32, "storage_dtype": torch.int16, "is_unsigned": True, "is_integer": False, "is_packed": True},
"float12_e3m9fnu": {"min": 0, "max": 31.96875, "num_bits": 12, "sign": 0, "exponent": 3, "mantissa": 9, "min_normal": 0.125244140625, "target_dtype": "fp12", "torch_dtype": torch.float32, "storage_dtype": torch.int16, "is_unsigned": True, "is_integer": False, "is_packed": True},
"float12_e4m8fnu": {"min": 0, "max": 511.0, "num_bits": 12, "sign": 0, "exponent": 4, "mantissa": 8, "min_normal": 0.007843017578125, "target_dtype": "fp12", "torch_dtype": torch.float32, "storage_dtype": torch.int16, "is_unsigned": True, "is_integer": False, "is_packed": True},
"float12_e5m7fnu": {"min": 0, "max": 130560.0, "num_bits": 12, "sign": 0, "exponent": 5, "mantissa": 7, "min_normal": 3.075599670410156e-05, "target_dtype": "fp12", "torch_dtype": torch.float32, "storage_dtype": torch.int16, "is_unsigned": True, "is_integer": False, "is_packed": True},
#
"float11_e1m10fnu": {"min": 0, "max": 3.998046875, "num_bits": 11, "sign": 0, "exponent": 1, "mantissa": 10, "min_normal": 1.0009765625, "target_dtype": "fp11", "torch_dtype": torch.float32, "storage_dtype": torch.int16, "is_unsigned": True, "is_integer": False, "is_packed": True},
"float11_e2m9fnu": {"min": 0, "max": 7.9921875, "num_bits": 11, "sign": 0, "exponent": 2, "mantissa": 9, "min_normal": 0.5009765625, "target_dtype": "fp11", "torch_dtype": torch.float32, "storage_dtype": torch.int16, "is_unsigned": True, "is_integer": False, "is_packed": True},
"float11_e3m8fnu": {"min": 0, "max": 31.9375, "num_bits": 11, "sign": 0, "exponent": 3, "mantissa": 8, "min_normal": 0.12548828125, "target_dtype": "fp11", "torch_dtype": torch.float32, "storage_dtype": torch.int16, "is_unsigned": True, "is_integer": False, "is_packed": True},
"float11_e4m7fnu": {"min": 0, "max": 510.0, "num_bits": 11, "sign": 0, "exponent": 4, "mantissa": 7, "min_normal": 0.00787353515625, "target_dtype": "fp11", "torch_dtype": torch.float32, "storage_dtype": torch.int16, "is_unsigned": True, "is_integer": False, "is_packed": True},
"float11_e5m6fnu": {"min": 0, "max": 130048.0, "num_bits": 11, "sign": 0, "exponent": 5, "mantissa": 6, "min_normal": 3.0994415283203125e-05, "target_dtype": "fp11", "torch_dtype": torch.float32, "storage_dtype": torch.int16, "is_unsigned": True, "is_integer": False, "is_packed": True},
#
"float10_e1m9fnu": {"min": 0, "max": 3.99609375, "num_bits": 10, "sign": 0, "exponent": 1, "mantissa": 9, "min_normal": 1.001953125, "target_dtype": "fp10", "torch_dtype": torch.float32, "storage_dtype": torch.int16, "is_unsigned": True, "is_integer": False, "is_packed": True},
"float10_e2m8fnu": {"min": 0, "max": 7.984375, "num_bits": 10, "sign": 0, "exponent": 2, "mantissa": 8, "min_normal": 0.501953125, "target_dtype": "fp10", "torch_dtype": torch.float32, "storage_dtype": torch.int16, "is_unsigned": True, "is_integer": False, "is_packed": True},
"float10_e3m7fnu": {"min": 0, "max": 31.875, "num_bits": 10, "sign": 0, "exponent": 3, "mantissa": 7, "min_normal": 0.1259765625, "target_dtype": "fp10", "torch_dtype": torch.float32, "storage_dtype": torch.int16, "is_unsigned": True, "is_integer": False, "is_packed": True},
"float10_e4m6fnu": {"min": 0, "max": 508.0, "num_bits": 10, "sign": 0, "exponent": 4, "mantissa": 6, "min_normal": 0.0079345703125, "target_dtype": "fp10", "torch_dtype": torch.float32, "storage_dtype": torch.int16, "is_unsigned": True, "is_integer": False, "is_packed": True},
"float10_e5m5fnu": {"min": 0, "max": 129024.0, "num_bits": 10, "sign": 0, "exponent": 5, "mantissa": 5, "min_normal": 3.147125244140625e-05, "target_dtype": "fp10", "torch_dtype": torch.float32, "storage_dtype": torch.int16, "is_unsigned": True, "is_integer": False, "is_packed": True},
#
"float9_e1m8fnu": {"min": 0, "max": 3.9921875, "num_bits": 9, "sign": 0, "exponent": 1, "mantissa": 8, "min_normal": 1.00390625, "target_dtype": "fp9", "torch_dtype": torch.float32, "storage_dtype": torch.int16, "is_unsigned": True, "is_integer": False, "is_packed": True},
"float9_e2m7fnu": {"min": 0, "max": 7.96875, "num_bits": 9, "sign": 0, "exponent": 2, "mantissa": 7, "min_normal": 0.50390625, "target_dtype": "fp9", "torch_dtype": torch.float32, "storage_dtype": torch.int16, "is_unsigned": True, "is_integer": False, "is_packed": True},
"float9_e3m6fnu": {"min": 0, "max": 31.75, "num_bits": 9, "sign": 0, "exponent": 3, "mantissa": 6, "min_normal": 0.126953125, "target_dtype": "fp9", "torch_dtype": torch.float32, "storage_dtype": torch.int16, "is_unsigned": True, "is_integer": False, "is_packed": True},
"float9_e4m5fnu": {"min": 0, "max": 504.0, "num_bits": 9, "sign": 0, "exponent": 4, "mantissa": 5, "min_normal": 0.008056640625, "target_dtype": "fp9", "torch_dtype": torch.float32, "storage_dtype": torch.int16, "is_unsigned": True, "is_integer": False, "is_packed": True},
"float9_e5m4fnu": {"min": 0, "max": 126976.0, "num_bits": 9, "sign": 0, "exponent": 5, "mantissa": 4, "min_normal": 3.24249267578125e-05, "target_dtype": "fp9", "torch_dtype": torch.float32, "storage_dtype": torch.int16, "is_unsigned": True, "is_integer": False, "is_packed": True},
#
"float8_e1m7fnu": {"min": 0, "max": 3.984375, "num_bits": 8, "sign": 0, "exponent": 1, "mantissa": 7, "min_normal": 1.0078125, "target_dtype": "fp8", "torch_dtype": torch.float32, "storage_dtype": torch.uint8, "is_unsigned": True, "is_integer": False, "is_packed": True},
"float8_e2m6fnu": {"min": 0, "max": 7.9375, "num_bits": 8, "sign": 0, "exponent": 2, "mantissa": 6, "min_normal": 0.5078125, "target_dtype": "fp8", "torch_dtype": torch.float32, "storage_dtype": torch.uint8, "is_unsigned": True, "is_integer": False, "is_packed": True},
"float8_e3m5fnu": {"min": 0, "max": 31.5, "num_bits": 8, "sign": 0, "exponent": 3, "mantissa": 5, "min_normal": 0.12890625, "target_dtype": "fp8", "torch_dtype": torch.float32, "storage_dtype": torch.uint8, "is_unsigned": True, "is_integer": False, "is_packed": True},
@@ -166,9 +224,13 @@ dtype_dict = {
dtype_dict["fp32"] = dtype_dict["float32"]
dtype_dict["bf16"] = dtype_dict["bfloat16"]
dtype_dict["fp16"] = dtype_dict["float16"]
dtype_dict["fp15"] = dtype_dict["float15_e5m9fn"]
dtype_dict["fp14"] = dtype_dict["float14_e5m8fn"]
dtype_dict["fp13"] = dtype_dict["float13_e5m7fn"]
dtype_dict["fp12"] = dtype_dict["float12_e5m6fn"]
dtype_dict["fp11"] = dtype_dict["float11_e5m5fn"]
dtype_dict["fp10"] = dtype_dict["float10_e5m4fn"]
dtype_dict["fp9"] = dtype_dict["float9_e4m4fn"]
dtype_dict["fp8"] = dtype_dict["float8_e4m3fn"]
dtype_dict["fp7"] = dtype_dict["float7_e3m3fn"]
dtype_dict["fp6"] = dtype_dict["float6_e3m2fn"]
@@ -228,14 +290,22 @@ weights_dtype_order = [
"uint7", "float7_e1m6fnu", "float7_e2m5fnu", "float7_e3m4fnu", "float7_e4m3fnu", "float7_e5m2fnu",
"int8", "float8_e4m3fn", "float8_e5m2", "float8_e1m6fn", "float8_e2m5fn", "float8_e3m4fn", "float8_e4m3fn_sdnq", "float8_e5m2fn",
"uint8", "float8_e1m7fnu", "float8_e2m6fnu", "float8_e3m5fnu", "float8_e4m4fnu", "float8_e5m3fnu",
"int10", "float10_e1m8fn", "float10_e2m7fn", "float10_e3m6fn", "float10_e4m5fn_sdnq", "float10_e5m4fn",
"int9", "float9_e1m7fn", "float9_e2m6fn", "float9_e3m5fn", "float9_e4m4fn", "float9_e5m3fn",
"uint9", "float9_e1m8fnu", "float9_e2m7fnu", "float9_e3m6fnu", "float9_e4m5fnu", "float9_e5m4fnu",
"int10", "float10_e1m8fn", "float10_e2m7fn", "float10_e3m6fn", "float10_e4m5fn", "float10_e5m4fn",
"uint10", "float10_e1m9fnu", "float10_e2m8fnu", "float10_e3m7fnu", "float10_e4m6fnu", "float10_e5m5fnu",
"int12", "float12_e1m10fn", "float12_e2m9fn", "float12_e3m8fn", "float12_e4m7fn_sdnq", "float12_e5m6fn",
"int11", "float11_e1m9fn", "float11_e2m8fn", "float11_e3m7fn", "float11_e4m6fn", "float11_e5m5fn",
"uint11", "float11_e1m10fnu", "float11_e2m9fnu", "float11_e3m8fnu", "float11_e4m7fnu", "float11_e5m6fnu",
"int12", "float12_e1m10fn", "float12_e2m9fn", "float12_e3m8fn", "float12_e4m7fn", "float12_e5m6fn",
"uint12", "float12_e1m11fnu", "float12_e2m10fnu", "float12_e3m9fnu", "float12_e4m8fnu", "float12_e5m7fnu",
]
weights_dtype_order_fp32 = weights_dtype_order + [
"int14", "float14_e1m12fn", "float14_e2m11fn", "float14_e3m10fn", "float14_e4m9fn_sdnq", "float14_e5m8fn",
"int13", "float13_e1m11fn", "float13_e2m10fn", "float13_e3m9fn", "float13_e4m8fn", "float13_e5m7fn",
"uint13", "float13_e1m12fnu", "float13_e2m11fnu", "float13_e3m10fnu", "float13_e4m9fnu", "float13_e5m8fnu",
"int14", "float14_e1m12fn", "float14_e2m11fn", "float14_e3m10fn", "float14_e4m9fn", "float14_e5m8fn",
"uint14", "float14_e1m13fnu", "float14_e2m12fnu", "float14_e3m11fnu", "float14_e4m10fnu", "float14_e5m9fnu",
"int15", "float15_e1m13fn", "float15_e2m12fn", "float15_e3m11fn", "float15_e4m10fn", "float15_e5m9fn",
"uint15", "float15_e1m14fnu", "float15_e2m13fnu", "float15_e3m12fnu", "float15_e4m11fnu", "float15_e5m10fnu",
"int16", "float16", "float16_e1m14fn", "float16_e2m13fn", "float16_e3m12fn", "float16_e4m11fn", "float16_e5m10fn",
"uint16", "float16_e1m15fnu", "float16_e2m14fnu", "float16_e3m13fnu", "float16_e4m12fnu", "float16_e5m11fnu",
]
+4 -4
View File
@@ -77,12 +77,12 @@ def dequantize_packed_int_symmetric(weight: torch.ByteTensor, scale: torch.Float
@devices.inference_context()
def dequantize_packed_float_asymmetric(weight: torch.ByteTensor, scale: torch.FloatTensor, zero_point: torch.FloatTensor, shape: torch.Size, weights_dtype: str, svd_up: torch.FloatTensor = None, svd_down: torch.FloatTensor = None, dtype: torch.dtype = None, result_shape: torch.Size = None, skip_quantized_matmul: bool = False) -> torch.FloatTensor:
return dequantize_asymmetric(unpack_float(weight, shape, weights_dtype), scale, zero_point, svd_up=svd_up, svd_down=svd_down, dtype=dtype, result_shape=result_shape, skip_quantized_matmul=skip_quantized_matmul)
return dequantize_asymmetric(unpack_float(weight, weights_dtype, shape), scale, zero_point, svd_up=svd_up, svd_down=svd_down, dtype=dtype, result_shape=result_shape, skip_quantized_matmul=skip_quantized_matmul)
@devices.inference_context()
def dequantize_packed_float_symmetric(weight: torch.ByteTensor, scale: torch.FloatTensor, shape: torch.Size, weights_dtype: str, svd_up: torch.FloatTensor = None, svd_down: torch.FloatTensor = None, dtype: torch.dtype = None, result_shape: torch.Size = None, skip_quantized_matmul: bool = False, re_quantize_for_matmul: bool = False) -> torch.FloatTensor:
return dequantize_symmetric(unpack_float(weight, shape, weights_dtype), scale, svd_up=svd_up, svd_down=svd_down, dtype=dtype, result_shape=result_shape, skip_quantized_matmul=skip_quantized_matmul, re_quantize_for_matmul=re_quantize_for_matmul)
return dequantize_symmetric(unpack_float(weight, weights_dtype, shape), scale, svd_up=svd_up, svd_down=svd_down, dtype=dtype, result_shape=result_shape, skip_quantized_matmul=skip_quantized_matmul, re_quantize_for_matmul=re_quantize_for_matmul)
@devices.inference_context()
@@ -169,12 +169,12 @@ def re_quantize_matmul_packed_int_symmetric(weight: torch.ByteTensor, scale: tor
@devices.inference_context()
def re_quantize_matmul_packed_float_asymmetric(weight: torch.ByteTensor, scale: torch.FloatTensor, zero_point: torch.FloatTensor, shape: torch.Size, weights_dtype: str, matmul_dtype: str, result_shape: torch.Size, svd_up: torch.FloatTensor = None, svd_down: torch.FloatTensor = None) -> tuple[torch.Tensor, torch.FloatTensor]:
return re_quantize_matmul_asymmetric(unpack_float(weight, shape, weights_dtype), scale, zero_point, matmul_dtype, svd_up=svd_up, svd_down=svd_down, result_shape=result_shape)
return re_quantize_matmul_asymmetric(unpack_float(weight, weights_dtype, shape), scale, zero_point, matmul_dtype, svd_up=svd_up, svd_down=svd_down, result_shape=result_shape)
@devices.inference_context()
def re_quantize_matmul_packed_float_symmetric(weight: torch.ByteTensor, scale: torch.FloatTensor, shape: torch.Size, weights_dtype: str, matmul_dtype: str, result_shape: torch.Size = None, svd_up: torch.FloatTensor = None, svd_down: torch.FloatTensor = None) -> tuple[torch.Tensor, torch.FloatTensor]:
return re_quantize_matmul_symmetric(unpack_float(weight, shape, weights_dtype), scale, matmul_dtype, svd_up=svd_up, svd_down=svd_down, result_shape=result_shape)
return re_quantize_matmul_symmetric(unpack_float(weight, weights_dtype, shape), scale, matmul_dtype, svd_up=svd_up, svd_down=svd_down, result_shape=result_shape)
@devices.inference_context()
+1 -1
View File
@@ -36,7 +36,7 @@ def conv_fp16_matmul(
bias = torch.mm(torch.mm(input.to(dtype=svd_down.dtype), svd_down), svd_up)
if quantized_weight_shape is not None:
weight = unpack_float(weight, quantized_weight_shape, weights_dtype).to(dtype=torch.float16).t_()
weight = unpack_float(weight, weights_dtype, quantized_weight_shape).to(dtype=torch.float16).t_()
scale = scale.t()
elif weight.dtype != torch.float16:
weight = weight.to(dtype=torch.float16) # fp8 weights
+1 -1
View File
@@ -32,7 +32,7 @@ def conv_fp8_matmul(
svd_bias = torch.mm(torch.mm(input.to(dtype=svd_down.dtype), svd_down), svd_up)
if quantized_weight_shape is not None:
weight = unpack_float(weight, quantized_weight_shape, weights_dtype).to(dtype=torch.float8_e4m3fn).t_()
weight = unpack_float(weight, weights_dtype, quantized_weight_shape).to(dtype=torch.float8_e4m3fn).t_()
scale = scale.t()
input, input_scale = quantize_fp_mm_input(input)
input, weight = check_mats(input, weight)
@@ -36,7 +36,7 @@ def conv_fp8_matmul_tensorwise(
bias = torch.mm(torch.mm(input.to(dtype=svd_down.dtype), svd_down), svd_up)
if quantized_weight_shape is not None:
weight = unpack_float(weight, quantized_weight_shape, weights_dtype).to(dtype=torch.float8_e4m3fn).t_()
weight = unpack_float(weight, weights_dtype, quantized_weight_shape).to(dtype=torch.float8_e4m3fn).t_()
scale = scale.t()
input, scale = quantize_fp_mm_input_tensorwise(input, scale)
input, weight = check_mats(input, weight)
+1 -1
View File
@@ -21,7 +21,7 @@ def fp16_matmul(
weights_dtype: str = None,
) -> torch.FloatTensor:
if quantized_weight_shape is not None:
weight = unpack_float(weight, quantized_weight_shape, weights_dtype).to(dtype=torch.float16).t_()
weight = unpack_float(weight, weights_dtype, quantized_weight_shape).to(dtype=torch.float16).t_()
scale = scale.t()
elif weight.dtype != torch.float16:
weight = weight.to(dtype=torch.float16) # fp8 weights
+1 -1
View File
@@ -26,7 +26,7 @@ def fp8_matmul(
weights_dtype: str = None,
) -> torch.FloatTensor:
if quantized_weight_shape is not None:
weight = unpack_float(weight, quantized_weight_shape, weights_dtype).to(dtype=torch.float8_e4m3fn).t_()
weight = unpack_float(weight, weights_dtype, quantized_weight_shape).to(dtype=torch.float8_e4m3fn).t_()
scale = scale.t()
return_dtype = input.dtype
output_shape = (*input.shape[:-1], weight.shape[-1])
@@ -29,7 +29,7 @@ def fp8_matmul_tensorwise(
weights_dtype: str = None,
) -> torch.FloatTensor:
if quantized_weight_shape is not None:
weight = unpack_float(weight, quantized_weight_shape, weights_dtype).to(dtype=torch.float8_e4m3fn).t_()
weight = unpack_float(weight, weights_dtype, quantized_weight_shape).to(dtype=torch.float8_e4m3fn).t_()
scale = scale.t()
return_dtype = input.dtype
output_shape = (*input.shape[:-1], weight.shape[-1])
+12 -3
View File
@@ -165,7 +165,7 @@ def post_process_model(model):
def apply_sdnq_options_to_module(model, dtype: torch.dtype = None, dequantize_fp32: bool = None, use_quantized_matmul: bool = None):
has_children = list(model.children())
if not has_children:
if dtype is not None and getattr(model, "dtype", torch.float32) != torch.float32:
if dtype is not None and getattr(model, "dtype", torch.float32) not in {torch.float32, torch.float64}:
model = model.to(dtype=dtype)
return model
for module_name, module in model.named_children():
@@ -182,7 +182,7 @@ def apply_sdnq_options_to_module(model, dtype: torch.dtype = None, dequantize_fp
current_use_quantized_matmul = current_use_quantized_matmul and channel_size >= 32 and output_channel_size >= 32 # pylint: disable=possibly-used-before-assignment
current_use_quantized_matmul = current_use_quantized_matmul and output_channel_size % 16 == 0 and channel_size % 16 == 0 # pylint: disable=possibly-used-before-assignment
if dtype is not None and module.sdnq_dequantizer.result_dtype != torch.float32:
if dtype is not None and module.sdnq_dequantizer.result_dtype not in {torch.float32, torch.float64}:
module.sdnq_dequantizer.result_dtype = dtype
upcast_scale = bool(
@@ -194,7 +194,16 @@ def apply_sdnq_options_to_module(model, dtype: torch.dtype = None, dequantize_fp
and (not use_tensorwise_fp8_matmul or dtype_dict[module.sdnq_dequantizer.quantized_matmul_dtype]["num_bits"] == 16)
)
)
scale_dtype = torch.float32 if upcast_scale or dequantize_fp32 or (dequantize_fp32 is None and module.scale.dtype == torch.float32) else module.sdnq_dequantizer.result_dtype
if upcast_scale or dequantize_fp32:
if module.scale.dtype in {torch.float32, torch.float64}:
scale_dtype = module.scale.dtype
else:
scale_dtype = torch.float32 if module.sdnq_dequantizer.result_dtype != torch.float64 else torch.float64
elif dequantize_fp32 is None and module.scale.dtype in {torch.float32, torch.float64}:
scale_dtype = module.scale.dtype
else:
scale_dtype = module.sdnq_dequantizer.result_dtype
module.scale.data = module.scale.to(dtype=scale_dtype)
if module.zero_point is not None:
+5 -1
View File
@@ -12,9 +12,13 @@ float_bits_to_uint_dict = {
5: "uint5",
6: "uint6",
7: "uint7",
9: "uint9",
10: "uint10",
11: "uint11",
12: "uint12",
13: "uint13",
14: "uint14",
15: "uint15",
}
@@ -62,7 +66,7 @@ def pack_float(x: torch.FloatTensor, weights_dtype: str) -> torch.Tensor:
return x
def unpack_float(x: torch.Tensor, shape: torch.Size, weights_dtype: str) -> torch.FloatTensor:
def unpack_float(x: torch.Tensor, weights_dtype: str, shape: torch.Size) -> torch.FloatTensor:
exponent_bits = dtype_dict[weights_dtype]["exponent"]
mantissa_bits = dtype_dict[weights_dtype]["mantissa"]
total_bits = dtype_dict[weights_dtype]["num_bits"]
-429
View File
@@ -1,429 +0,0 @@
# pylint: disable=redefined-builtin,no-member,protected-access
import torch
from .common import dtype_dict
def pack_int(tensor: torch.Tensor, weights_dtype: str) -> torch.Tensor:
if not dtype_dict[weights_dtype]["is_unsigned"]:
tensor = tensor.sub(dtype_dict[weights_dtype]["min"])
return packed_int_function_dict[weights_dtype]["pack"](tensor.to(dtype=dtype_dict[weights_dtype]["storage_dtype"]))
def unpack_int(packed_tensor: torch.Tensor, weights_dtype: str, shape: torch.Size, dtype: torch.dtype = None) -> torch.Tensor:
packed_tensor = packed_int_function_dict[weights_dtype]["unpack"](packed_tensor, shape)
if not dtype_dict[weights_dtype]["is_unsigned"]:
packed_tensor = packed_tensor.to(dtype=dtype_dict[weights_dtype]["torch_dtype"] if dtype is None else dtype).add_(dtype_dict[weights_dtype]["min"])
return packed_tensor
def pack_uint14(tensor: torch.Tensor) -> torch.Tensor:
packed_tensor = tensor.contiguous().view(-1, 8)
packed_tensor = torch.bitwise_or(
packed_tensor[:, :7],
torch.bitwise_and(
torch.stack(
(
torch.bitwise_left_shift(packed_tensor[:, 7], 2),
torch.bitwise_left_shift(packed_tensor[:, 7], 4),
torch.bitwise_left_shift(packed_tensor[:, 7], 6),
torch.bitwise_left_shift(packed_tensor[:, 7], 8),
torch.bitwise_left_shift(packed_tensor[:, 7], 10),
torch.bitwise_left_shift(packed_tensor[:, 7], 12),
torch.bitwise_left_shift(packed_tensor[:, 7], 14),
),
dim=-1
),
49152
),
)
return packed_tensor
def pack_uint12(tensor: torch.Tensor) -> torch.Tensor:
packed_tensor = tensor.contiguous().view(-1, 4)
packed_tensor = torch.bitwise_or(
packed_tensor[:, :3],
torch.bitwise_and(
torch.stack(
(
torch.bitwise_left_shift(packed_tensor[:, 3], 4),
torch.bitwise_left_shift(packed_tensor[:, 3], 8),
torch.bitwise_left_shift(packed_tensor[:, 3], 12),
),
dim=-1
),
61440
)
)
return packed_tensor
def pack_uint10(tensor: torch.ByteTensor) -> torch.ByteTensor:
packed_tensor = tensor.contiguous().view(-1, 8)
packed_tensor = torch.cat(
(
torch.bitwise_or(packed_tensor[:, :3], torch.bitwise_left_shift(packed_tensor[:, 5:8], 10)),
torch.bitwise_or(
packed_tensor[:, 3],
torch.bitwise_or(
torch.bitwise_and(torch.bitwise_left_shift(packed_tensor[:, 5], 4), 15360),
torch.bitwise_and(torch.bitwise_left_shift(packed_tensor[:, 7], 6), 49152),
),
).unsqueeze(-1),
torch.bitwise_or(
packed_tensor[:, 4],
torch.bitwise_or(
torch.bitwise_and(torch.bitwise_left_shift(packed_tensor[:, 6], 4), 15360),
torch.bitwise_and(torch.bitwise_left_shift(packed_tensor[:, 7], 8), 49152),
),
).unsqueeze(-1),
),
dim=-1
)
return packed_tensor
def pack_uint7(tensor: torch.ByteTensor) -> torch.ByteTensor:
packed_tensor = tensor.contiguous().view(-1, 8)
packed_tensor = torch.bitwise_or(
packed_tensor[:, :7],
torch.bitwise_and(
torch.stack(
(
torch.bitwise_left_shift(packed_tensor[:, 7], 1),
torch.bitwise_left_shift(packed_tensor[:, 7], 2),
torch.bitwise_left_shift(packed_tensor[:, 7], 3),
torch.bitwise_left_shift(packed_tensor[:, 7], 4),
torch.bitwise_left_shift(packed_tensor[:, 7], 5),
torch.bitwise_left_shift(packed_tensor[:, 7], 6),
torch.bitwise_left_shift(packed_tensor[:, 7], 7),
),
dim=-1
),
128
),
)
return packed_tensor
def pack_uint6(tensor: torch.ByteTensor) -> torch.ByteTensor:
packed_tensor = tensor.contiguous().view(-1, 4)
packed_tensor = torch.bitwise_or(
packed_tensor[:, :3],
torch.bitwise_and(
torch.stack(
(
torch.bitwise_left_shift(packed_tensor[:, 3], 2),
torch.bitwise_left_shift(packed_tensor[:, 3], 4),
torch.bitwise_left_shift(packed_tensor[:, 3], 6),
),
dim=-1
),
192
)
)
return packed_tensor
def pack_uint5(tensor: torch.ByteTensor) -> torch.ByteTensor:
packed_tensor = tensor.contiguous().view(-1, 8)
packed_tensor = torch.cat(
(
torch.bitwise_or(packed_tensor[:, :3], torch.bitwise_left_shift(packed_tensor[:, 5:8], 5)),
torch.bitwise_or(
packed_tensor[:, 3],
torch.bitwise_or(
torch.bitwise_and(torch.bitwise_left_shift(packed_tensor[:, 5], 2), 96),
torch.bitwise_and(torch.bitwise_left_shift(packed_tensor[:, 7], 3), 128),
),
).unsqueeze(-1),
torch.bitwise_or(
packed_tensor[:, 4],
torch.bitwise_or(
torch.bitwise_and(torch.bitwise_left_shift(packed_tensor[:, 6], 2), 96),
torch.bitwise_and(torch.bitwise_left_shift(packed_tensor[:, 7], 4), 128),
),
).unsqueeze(-1),
),
dim=-1
)
return packed_tensor
def pack_uint4(tensor: torch.ByteTensor) -> torch.ByteTensor:
packed_tensor = tensor.contiguous().view(-1, 2)
packed_tensor = torch.bitwise_or(packed_tensor[:, 0], torch.bitwise_left_shift(packed_tensor[:, 1], 4))
return packed_tensor
def pack_uint3(tensor: torch.ByteTensor) -> torch.ByteTensor:
packed_tensor = tensor.contiguous().view(-1, 8)
packed_tensor = torch.bitwise_or(
torch.bitwise_or(packed_tensor[:, :3], torch.bitwise_left_shift(packed_tensor[:, 3:6], 3)),
torch.cat(
(
torch.bitwise_left_shift(packed_tensor[:, 6:8], 6),
torch.bitwise_or(
torch.bitwise_and(torch.bitwise_left_shift(packed_tensor[:, 6], 4), 64),
torch.bitwise_and(torch.bitwise_left_shift(packed_tensor[:, 7], 5), 128),
).unsqueeze(-1),
),
dim=-1
)
)
return packed_tensor
def pack_uint2(tensor: torch.ByteTensor) -> torch.ByteTensor:
packed_tensor = tensor.contiguous().view(-1, 4)
packed_tensor = torch.bitwise_or(
torch.bitwise_or(packed_tensor[:, 0], torch.bitwise_left_shift(packed_tensor[:, 1], 2)),
torch.bitwise_or(torch.bitwise_left_shift(packed_tensor[:, 2], 4), torch.bitwise_left_shift(packed_tensor[:, 3], 6)),
)
return packed_tensor
def pack_uint1(tensor: torch.Tensor) -> torch.Tensor:
packed_tensor = tensor.contiguous().view(-1, 8)
packed_tensor = torch.bitwise_or(
torch.bitwise_or(
torch.bitwise_or(packed_tensor[:, 0], torch.bitwise_left_shift(packed_tensor[:, 1], 1)),
torch.bitwise_or(torch.bitwise_left_shift(packed_tensor[:, 2], 2), torch.bitwise_left_shift(packed_tensor[:, 3], 3))
),
torch.bitwise_or(
torch.bitwise_or(torch.bitwise_left_shift(packed_tensor[:, 4], 4), torch.bitwise_left_shift(packed_tensor[:, 5], 5)),
torch.bitwise_or(torch.bitwise_left_shift(packed_tensor[:, 6], 6), torch.bitwise_left_shift(packed_tensor[:, 7], 7))
),
)
return packed_tensor
def unpack_uint14(packed_tensor: torch.Tensor, shape: torch.Size) -> torch.Tensor:
result = torch.cat(
(
torch.bitwise_and(packed_tensor[:, :7], 16383),
torch.bitwise_or(
torch.bitwise_or(
torch.bitwise_or(
torch.bitwise_and(torch.bitwise_right_shift(packed_tensor[:, 0], 2), 12288),
torch.bitwise_and(torch.bitwise_right_shift(packed_tensor[:, 1], 4), 3072),
),
torch.bitwise_or(
torch.bitwise_and(torch.bitwise_right_shift(packed_tensor[:, 2], 6), 768),
torch.bitwise_and(torch.bitwise_right_shift(packed_tensor[:, 3], 8), 192),
),
),
torch.bitwise_or(
torch.bitwise_or(
torch.bitwise_and(torch.bitwise_right_shift(packed_tensor[:, 4], 10), 48),
torch.bitwise_and(torch.bitwise_right_shift(packed_tensor[:, 5], 12), 12),
),
torch.bitwise_and(torch.bitwise_right_shift(packed_tensor[:, 6], 14), 3),
),
).unsqueeze(-1)
),
dim=-1
).view(shape)
return result
def unpack_uint12(packed_tensor: torch.Tensor, shape: torch.Size) -> torch.Tensor:
result = torch.cat(
(
torch.bitwise_and(packed_tensor[:, :3], 4095),
torch.bitwise_or(
torch.bitwise_or(
torch.bitwise_and(torch.bitwise_right_shift(packed_tensor[:, 0], 4), 3840),
torch.bitwise_and(torch.bitwise_right_shift(packed_tensor[:, 1], 8), 240),
),
torch.bitwise_and(torch.bitwise_right_shift(packed_tensor[:, 2], 12), 15)
).unsqueeze(-1)
),
dim=-1
).view(shape)
return result
def unpack_uint10(packed_tensor: torch.Tensor, shape: torch.Size) -> torch.Tensor:
result_bitwise_right_shift = torch.bitwise_and(torch.bitwise_right_shift(packed_tensor[:, :3], 10), 63)
result = torch.cat(
(
torch.bitwise_and(packed_tensor[:, :5], 1023),
torch.bitwise_or(
result_bitwise_right_shift[:, :2],
torch.bitwise_and(torch.bitwise_right_shift(packed_tensor[:, 3:5], 4), 960),
),
torch.bitwise_or(
result_bitwise_right_shift[:, 2],
torch.bitwise_or(
torch.bitwise_and(torch.bitwise_right_shift(packed_tensor[:, 3], 6), 768),
torch.bitwise_and(torch.bitwise_right_shift(packed_tensor[:, 4], 8), 192),
),
).unsqueeze(-1),
),
dim=-1
).view(shape)
return result
def unpack_uint7(packed_tensor: torch.ByteTensor, shape: torch.Size) -> torch.ByteTensor:
result = torch.cat(
(
torch.bitwise_and(packed_tensor[:, :7], 127),
torch.bitwise_or(
torch.bitwise_or(
torch.bitwise_or(
torch.bitwise_and(torch.bitwise_right_shift(packed_tensor[:, 0], 1), 64),
torch.bitwise_and(torch.bitwise_right_shift(packed_tensor[:, 1], 2), 32),
),
torch.bitwise_or(
torch.bitwise_and(torch.bitwise_right_shift(packed_tensor[:, 2], 3), 16),
torch.bitwise_and(torch.bitwise_right_shift(packed_tensor[:, 3], 4), 8),
),
),
torch.bitwise_or(
torch.bitwise_or(
torch.bitwise_and(torch.bitwise_right_shift(packed_tensor[:, 4], 5), 4),
torch.bitwise_and(torch.bitwise_right_shift(packed_tensor[:, 5], 6), 2),
),
torch.bitwise_right_shift(packed_tensor[:, 6], 7),
),
).unsqueeze(-1)
),
dim=-1
).view(shape)
return result
def unpack_uint6(packed_tensor: torch.ByteTensor, shape: torch.Size) -> torch.ByteTensor:
result = torch.cat(
(
torch.bitwise_and(packed_tensor[:, :3], 63),
torch.bitwise_or(
torch.bitwise_or(
torch.bitwise_and(torch.bitwise_right_shift(packed_tensor[:, 0], 2), 48),
torch.bitwise_and(torch.bitwise_right_shift(packed_tensor[:, 1], 4), 12),
),
torch.bitwise_right_shift(packed_tensor[:, 2], 6)
).unsqueeze(-1)
),
dim=-1
).view(shape)
return result
def unpack_uint5(packed_tensor: torch.ByteTensor, shape: torch.Size) -> torch.ByteTensor:
result_bitwise_right_shift = torch.bitwise_right_shift(packed_tensor[:, :3], 5)
result = torch.cat(
(
torch.bitwise_and(packed_tensor[:, :5], 31),
torch.bitwise_or(
result_bitwise_right_shift[:, :2],
torch.bitwise_and(torch.bitwise_right_shift(packed_tensor[:, 3:5], 2), 24),
),
torch.bitwise_or(
result_bitwise_right_shift[:, 2],
torch.bitwise_or(
torch.bitwise_and(torch.bitwise_right_shift(packed_tensor[:, 3], 3), 16),
torch.bitwise_and(torch.bitwise_right_shift(packed_tensor[:, 4], 4), 8),
),
).unsqueeze(-1),
),
dim=-1
).view(shape)
return result
def unpack_uint4(packed_tensor: torch.ByteTensor, shape: torch.Size) -> torch.ByteTensor:
result = torch.stack((torch.bitwise_and(packed_tensor, 15), torch.bitwise_right_shift(packed_tensor, 4)), dim=-1).view(shape)
return result
def unpack_uint3(packed_tensor: torch.ByteTensor, shape: torch.Size) -> torch.ByteTensor:
result = torch.bitwise_and(
torch.cat(
(
packed_tensor[:, :3],
torch.bitwise_right_shift(packed_tensor[:, :3], 3),
torch.bitwise_or(
torch.bitwise_right_shift(packed_tensor[:, :2], 6),
torch.bitwise_and(
torch.stack(
(
torch.bitwise_right_shift(packed_tensor[:, 2], 4),
torch.bitwise_right_shift(packed_tensor[:, 2], 5),
),
dim=-1
),
4
),
),
),
dim=-1
),
7
).view(shape)
return result
def unpack_uint2(packed_tensor: torch.ByteTensor, shape: torch.Size) -> torch.ByteTensor:
result = torch.bitwise_and(
torch.stack(
(
packed_tensor,
torch.bitwise_right_shift(packed_tensor, 2),
torch.bitwise_right_shift(packed_tensor, 4),
torch.bitwise_right_shift(packed_tensor, 6),
),
dim=-1
),
3
).view(shape)
return result
def unpack_uint1(packed_tensor: torch.Tensor, shape: torch.Size) -> torch.Tensor:
result = torch.bitwise_and(
torch.stack(
(
packed_tensor,
torch.bitwise_right_shift(packed_tensor, 1),
torch.bitwise_right_shift(packed_tensor, 2),
torch.bitwise_right_shift(packed_tensor, 3),
torch.bitwise_right_shift(packed_tensor, 4),
torch.bitwise_right_shift(packed_tensor, 5),
torch.bitwise_right_shift(packed_tensor, 6),
torch.bitwise_right_shift(packed_tensor, 7),
),
dim=-1
),
1
).view(shape)
return result
packed_int_function_dict = {
"uint14": {"pack": pack_uint14, "unpack": unpack_uint14},
"uint12": {"pack": pack_uint12, "unpack": unpack_uint12},
"uint10": {"pack": pack_uint10, "unpack": unpack_uint10},
"uint7": {"pack": pack_uint7, "unpack": unpack_uint7},
"uint6": {"pack": pack_uint6, "unpack": unpack_uint6},
"uint5": {"pack": pack_uint5, "unpack": unpack_uint5},
"uint4": {"pack": pack_uint4, "unpack": unpack_uint4},
"uint3": {"pack": pack_uint3, "unpack": unpack_uint3},
"uint2": {"pack": pack_uint2, "unpack": unpack_uint2},
"uint1": {"pack": pack_uint1, "unpack": unpack_uint1},
}
packed_int_function_dict["int14"] = packed_int_function_dict["uint14"]
packed_int_function_dict["int12"] = packed_int_function_dict["uint12"]
packed_int_function_dict["int10"] = packed_int_function_dict["uint10"]
packed_int_function_dict["int7"] = packed_int_function_dict["uint7"]
packed_int_function_dict["int6"] = packed_int_function_dict["uint6"]
packed_int_function_dict["int5"] = packed_int_function_dict["uint5"]
packed_int_function_dict["int4"] = packed_int_function_dict["uint4"]
packed_int_function_dict["int3"] = packed_int_function_dict["uint3"]
packed_int_function_dict["int2"] = packed_int_function_dict["uint2"]
packed_int_function_dict["bool"] = packed_int_function_dict["uint1"]
+84
View File
@@ -0,0 +1,84 @@
import torch
from ..common import dtype_dict # noqa: TID252
from .pack import (
pack_uint15,
pack_uint14,
pack_uint13,
pack_uint12,
pack_uint11,
pack_uint10,
pack_uint9,
pack_uint7,
pack_uint6,
pack_uint5,
pack_uint4,
pack_uint3,
pack_uint2,
pack_uint1,
)
from .unpack import (
unpack_uint15,
unpack_uint14,
unpack_uint13,
unpack_uint12,
unpack_uint11,
unpack_uint10,
unpack_uint9,
unpack_uint7,
unpack_uint6,
unpack_uint5,
unpack_uint4,
unpack_uint3,
unpack_uint2,
unpack_uint1,
)
packed_int_function_dict = {
"uint15": {"pack": pack_uint15, "unpack": unpack_uint15},
"uint14": {"pack": pack_uint14, "unpack": unpack_uint14},
"uint13": {"pack": pack_uint13, "unpack": unpack_uint13},
"uint12": {"pack": pack_uint12, "unpack": unpack_uint12},
"uint11": {"pack": pack_uint11, "unpack": unpack_uint11},
"uint10": {"pack": pack_uint10, "unpack": unpack_uint10},
"uint9": {"pack": pack_uint9, "unpack": unpack_uint9},
"uint7": {"pack": pack_uint7, "unpack": unpack_uint7},
"uint6": {"pack": pack_uint6, "unpack": unpack_uint6},
"uint5": {"pack": pack_uint5, "unpack": unpack_uint5},
"uint4": {"pack": pack_uint4, "unpack": unpack_uint4},
"uint3": {"pack": pack_uint3, "unpack": unpack_uint3},
"uint2": {"pack": pack_uint2, "unpack": unpack_uint2},
"uint1": {"pack": pack_uint1, "unpack": unpack_uint1},
}
packed_int_function_dict["int15"] = packed_int_function_dict["uint15"]
packed_int_function_dict["int14"] = packed_int_function_dict["uint14"]
packed_int_function_dict["int13"] = packed_int_function_dict["uint13"]
packed_int_function_dict["int12"] = packed_int_function_dict["uint12"]
packed_int_function_dict["int11"] = packed_int_function_dict["uint11"]
packed_int_function_dict["int10"] = packed_int_function_dict["uint10"]
packed_int_function_dict["int9"] = packed_int_function_dict["uint9"]
packed_int_function_dict["int7"] = packed_int_function_dict["uint7"]
packed_int_function_dict["int6"] = packed_int_function_dict["uint6"]
packed_int_function_dict["int5"] = packed_int_function_dict["uint5"]
packed_int_function_dict["int4"] = packed_int_function_dict["uint4"]
packed_int_function_dict["int3"] = packed_int_function_dict["uint3"]
packed_int_function_dict["int2"] = packed_int_function_dict["uint2"]
packed_int_function_dict["bool"] = packed_int_function_dict["uint1"]
def pack_int(tensor: torch.Tensor, weights_dtype: str) -> torch.Tensor:
if not dtype_dict[weights_dtype]["is_unsigned"]:
tensor = tensor.sub(dtype_dict[weights_dtype]["min"])
return packed_int_function_dict[weights_dtype]["pack"](tensor.to(dtype=dtype_dict[weights_dtype]["storage_dtype"]))
def unpack_int(packed_tensor: torch.Tensor, weights_dtype: str, shape: torch.Size, dtype: torch.dtype = None) -> torch.Tensor:
packed_tensor = packed_int_function_dict[weights_dtype]["unpack"](packed_tensor, shape)
if not dtype_dict[weights_dtype]["is_unsigned"]:
packed_tensor = packed_tensor.to(dtype=dtype_dict[weights_dtype]["torch_dtype"] if dtype is None else dtype).add_(dtype_dict[weights_dtype]["min"])
return packed_tensor
+305
View File
@@ -0,0 +1,305 @@
import torch
def pack_uint15(tensor: torch.Tensor) -> torch.Tensor:
packed_tensor = tensor.contiguous().view(-1, 16)
packed_tensor = torch.bitwise_or(
packed_tensor[:, :15],
torch.bitwise_and(
torch.stack(
(
torch.bitwise_left_shift(packed_tensor[:, 15], 1),
torch.bitwise_left_shift(packed_tensor[:, 15], 2),
torch.bitwise_left_shift(packed_tensor[:, 15], 3),
torch.bitwise_left_shift(packed_tensor[:, 15], 4),
torch.bitwise_left_shift(packed_tensor[:, 15], 5),
torch.bitwise_left_shift(packed_tensor[:, 15], 6),
torch.bitwise_left_shift(packed_tensor[:, 15], 7),
torch.bitwise_left_shift(packed_tensor[:, 15], 8),
torch.bitwise_left_shift(packed_tensor[:, 15], 9),
torch.bitwise_left_shift(packed_tensor[:, 15], 10),
torch.bitwise_left_shift(packed_tensor[:, 15], 11),
torch.bitwise_left_shift(packed_tensor[:, 15], 12),
torch.bitwise_left_shift(packed_tensor[:, 15], 13),
torch.bitwise_left_shift(packed_tensor[:, 15], 14),
torch.bitwise_left_shift(packed_tensor[:, 15], 15),
),
dim=-1
),
32768
),
)
return packed_tensor
def pack_uint14(tensor: torch.Tensor) -> torch.Tensor:
packed_tensor = tensor.contiguous().view(-1, 8)
packed_tensor = torch.bitwise_or(
packed_tensor[:, :7],
torch.bitwise_and(
torch.stack(
(
torch.bitwise_left_shift(packed_tensor[:, 7], 2),
torch.bitwise_left_shift(packed_tensor[:, 7], 4),
torch.bitwise_left_shift(packed_tensor[:, 7], 6),
torch.bitwise_left_shift(packed_tensor[:, 7], 8),
torch.bitwise_left_shift(packed_tensor[:, 7], 10),
torch.bitwise_left_shift(packed_tensor[:, 7], 12),
torch.bitwise_left_shift(packed_tensor[:, 7], 14),
),
dim=-1
),
49152
),
)
return packed_tensor
def pack_uint13(tensor: torch.Tensor) -> torch.Tensor:
packed_tensor = tensor.contiguous().view(-1, 16)
packed_tensor = torch.bitwise_or(
packed_tensor[:, :13],
torch.bitwise_and(
torch.cat(
(
torch.bitwise_left_shift(packed_tensor[:, 13:], 13),
torch.bitwise_left_shift(packed_tensor[:, 13:], 10),
torch.bitwise_left_shift(packed_tensor[:, 13:], 7),
torch.bitwise_left_shift(packed_tensor[:, 13:], 4),
torch.bitwise_or(
torch.bitwise_or(
torch.bitwise_and(torch.bitwise_left_shift(packed_tensor[:, 13], 1), 8192),
torch.bitwise_and(torch.bitwise_left_shift(packed_tensor[:, 14], 2), 16384),
),
torch.bitwise_and(torch.bitwise_left_shift(packed_tensor[:, 15], 3), 32768),
).unsqueeze(-1),
),
dim=-1,
),
57344,
),
)
return packed_tensor
def pack_uint12(tensor: torch.Tensor) -> torch.Tensor:
packed_tensor = tensor.contiguous().view(-1, 4)
packed_tensor = torch.bitwise_or(
packed_tensor[:, :3],
torch.bitwise_and(
torch.stack(
(
torch.bitwise_left_shift(packed_tensor[:, 3], 4),
torch.bitwise_left_shift(packed_tensor[:, 3], 8),
torch.bitwise_left_shift(packed_tensor[:, 3], 12),
),
dim=-1
),
61440
)
)
return packed_tensor
def pack_uint11(tensor: torch.Tensor) -> torch.Tensor:
packed_tensor = tensor.contiguous().view(-1, 16)
packed_tensor = torch.cat(
(
torch.bitwise_or(packed_tensor[:, :8], torch.bitwise_left_shift(packed_tensor[:, 8:], 11)),
torch.bitwise_or(
torch.bitwise_or(
torch.bitwise_and(torch.bitwise_right_shift(packed_tensor[:, 8:11], 5), 63),
torch.bitwise_and(torch.bitwise_left_shift(packed_tensor[:, 11:14], 1), 4032),
),
torch.cat(
(
torch.bitwise_and(torch.bitwise_left_shift(packed_tensor[:, 14:], 7), -4096),
torch.bitwise_or(
torch.bitwise_and(torch.bitwise_left_shift(packed_tensor[:, 14], 3), 12288),
torch.bitwise_and(torch.bitwise_left_shift(packed_tensor[:, 15], 5), -16384),
).unsqueeze(-1),
),
dim=-1,
),
),
),
dim=-1,
)
return packed_tensor
def pack_uint10(tensor: torch.Tensor) -> torch.Tensor:
packed_tensor = tensor.contiguous().view(-1, 8)
packed_tensor = torch.cat(
(
torch.bitwise_or(packed_tensor[:, :3], torch.bitwise_left_shift(packed_tensor[:, 5:8], 10)),
torch.bitwise_or(
packed_tensor[:, 3:5],
torch.bitwise_or(
torch.bitwise_and(torch.bitwise_left_shift(packed_tensor[:, 5:7], 4), 15360),
torch.bitwise_and(
torch.stack(
(
torch.bitwise_left_shift(packed_tensor[:, 7], 6),
torch.bitwise_left_shift(packed_tensor[:, 7], 8),
),
dim=-1,
),
49152
),
),
),
),
dim=-1
)
return packed_tensor
def pack_uint9(tensor: torch.Tensor) -> torch.Tensor:
packed_tensor = tensor.contiguous().view(-1, 16)
packed_tensor = torch.cat(
(
torch.bitwise_or(packed_tensor[:, :8], torch.bitwise_left_shift(packed_tensor[:, 8:], 9)),
torch.bitwise_or(
torch.bitwise_or(
torch.bitwise_or(
torch.bitwise_and(torch.bitwise_right_shift(packed_tensor[:, 8], 7), 3),
torch.bitwise_and(torch.bitwise_right_shift(packed_tensor[:, 9], 5), 12),
),
torch.bitwise_or(
torch.bitwise_and(torch.bitwise_right_shift(packed_tensor[:, 10], 3), 48),
torch.bitwise_and(torch.bitwise_right_shift(packed_tensor[:, 11], 1), 192),
),
),
torch.bitwise_or(
torch.bitwise_or(
torch.bitwise_and(torch.bitwise_left_shift(packed_tensor[:, 12], 1), 768),
torch.bitwise_and(torch.bitwise_left_shift(packed_tensor[:, 13], 3), 3072),
),
torch.bitwise_or(
torch.bitwise_and(torch.bitwise_left_shift(packed_tensor[:, 14], 5), 12288),
torch.bitwise_and(torch.bitwise_left_shift(packed_tensor[:, 15], 7), 49152),
),
),
).unsqueeze(-1),
),
dim=-1,
)
return packed_tensor
def pack_uint7(tensor: torch.ByteTensor) -> torch.ByteTensor:
packed_tensor = tensor.contiguous().view(-1, 8)
packed_tensor = torch.bitwise_or(
packed_tensor[:, :7],
torch.bitwise_and(
torch.stack(
(
torch.bitwise_left_shift(packed_tensor[:, 7], 1),
torch.bitwise_left_shift(packed_tensor[:, 7], 2),
torch.bitwise_left_shift(packed_tensor[:, 7], 3),
torch.bitwise_left_shift(packed_tensor[:, 7], 4),
torch.bitwise_left_shift(packed_tensor[:, 7], 5),
torch.bitwise_left_shift(packed_tensor[:, 7], 6),
torch.bitwise_left_shift(packed_tensor[:, 7], 7),
),
dim=-1
),
128
),
)
return packed_tensor
def pack_uint6(tensor: torch.ByteTensor) -> torch.ByteTensor:
packed_tensor = tensor.contiguous().view(-1, 4)
packed_tensor = torch.bitwise_or(
packed_tensor[:, :3],
torch.bitwise_and(
torch.stack(
(
torch.bitwise_left_shift(packed_tensor[:, 3], 2),
torch.bitwise_left_shift(packed_tensor[:, 3], 4),
torch.bitwise_left_shift(packed_tensor[:, 3], 6),
),
dim=-1
),
192
)
)
return packed_tensor
def pack_uint5(tensor: torch.ByteTensor) -> torch.ByteTensor:
packed_tensor = tensor.contiguous().view(-1, 8)
packed_tensor = torch.cat(
(
torch.bitwise_or(packed_tensor[:, :3], torch.bitwise_left_shift(packed_tensor[:, 5:8], 5)),
torch.bitwise_or(
packed_tensor[:, 3:5],
torch.bitwise_or(
torch.bitwise_and(torch.bitwise_left_shift(packed_tensor[:, 5:7], 2), 96),
torch.bitwise_and(
torch.stack(
(
torch.bitwise_left_shift(packed_tensor[:, 7], 3),
torch.bitwise_left_shift(packed_tensor[:, 7], 4),
),
dim=-1,
),
128,
),
),
),
),
dim=-1
)
return packed_tensor
def pack_uint4(tensor: torch.ByteTensor) -> torch.ByteTensor:
packed_tensor = tensor.contiguous().view(-1, 2)
packed_tensor = torch.bitwise_or(packed_tensor[:, 0], torch.bitwise_left_shift(packed_tensor[:, 1], 4))
return packed_tensor
def pack_uint3(tensor: torch.ByteTensor) -> torch.ByteTensor:
packed_tensor = tensor.contiguous().view(-1, 8)
packed_tensor = torch.bitwise_or(
torch.bitwise_or(packed_tensor[:, :3], torch.bitwise_left_shift(packed_tensor[:, 3:6], 3)),
torch.cat(
(
torch.bitwise_left_shift(packed_tensor[:, 6:8], 6),
torch.bitwise_or(
torch.bitwise_and(torch.bitwise_left_shift(packed_tensor[:, 6], 4), 64),
torch.bitwise_and(torch.bitwise_left_shift(packed_tensor[:, 7], 5), 128),
).unsqueeze(-1),
),
dim=-1
)
)
return packed_tensor
def pack_uint2(tensor: torch.ByteTensor) -> torch.ByteTensor:
packed_tensor = tensor.contiguous().view(-1, 4)
packed_tensor = torch.bitwise_or(
torch.bitwise_or(packed_tensor[:, 0], torch.bitwise_left_shift(packed_tensor[:, 1], 2)),
torch.bitwise_or(torch.bitwise_left_shift(packed_tensor[:, 2], 4), torch.bitwise_left_shift(packed_tensor[:, 3], 6)),
)
return packed_tensor
def pack_uint1(tensor: torch.Tensor) -> torch.Tensor:
packed_tensor = tensor.contiguous().view(-1, 8)
packed_tensor = torch.bitwise_or(
torch.bitwise_or(
torch.bitwise_or(packed_tensor[:, 0], torch.bitwise_left_shift(packed_tensor[:, 1], 1)),
torch.bitwise_or(torch.bitwise_left_shift(packed_tensor[:, 2], 2), torch.bitwise_left_shift(packed_tensor[:, 3], 3))
),
torch.bitwise_or(
torch.bitwise_or(torch.bitwise_left_shift(packed_tensor[:, 4], 4), torch.bitwise_left_shift(packed_tensor[:, 5], 5)),
torch.bitwise_or(torch.bitwise_left_shift(packed_tensor[:, 6], 6), torch.bitwise_left_shift(packed_tensor[:, 7], 7))
),
)
return packed_tensor
+356
View File
@@ -0,0 +1,356 @@
import torch
def unpack_uint15(packed_tensor: torch.Tensor, shape: torch.Size) -> torch.Tensor:
result = torch.cat(
(
torch.bitwise_and(packed_tensor[:, :15], 32767),
torch.bitwise_or(
torch.bitwise_or(
torch.bitwise_or(
torch.bitwise_or(
torch.bitwise_and(torch.bitwise_right_shift(packed_tensor[:, 0], 1), 16384),
torch.bitwise_and(torch.bitwise_right_shift(packed_tensor[:, 1], 2), 8192),
),
torch.bitwise_or(
torch.bitwise_and(torch.bitwise_right_shift(packed_tensor[:, 2], 3), 4096),
torch.bitwise_and(torch.bitwise_right_shift(packed_tensor[:, 3], 4), 2048),
),
),
torch.bitwise_or(
torch.bitwise_or(
torch.bitwise_and(torch.bitwise_right_shift(packed_tensor[:, 4], 5), 1024),
torch.bitwise_and(torch.bitwise_right_shift(packed_tensor[:, 5], 6), 512),
),
torch.bitwise_or(
torch.bitwise_and(torch.bitwise_right_shift(packed_tensor[:, 6], 7), 256),
torch.bitwise_and(torch.bitwise_right_shift(packed_tensor[:, 7], 8), 128),
),
),
),
torch.bitwise_or(
torch.bitwise_or(
torch.bitwise_or(
torch.bitwise_and(torch.bitwise_right_shift(packed_tensor[:, 8], 9), 64),
torch.bitwise_and(torch.bitwise_right_shift(packed_tensor[:, 9], 10), 32),
),
torch.bitwise_or(
torch.bitwise_and(torch.bitwise_right_shift(packed_tensor[:, 10], 11), 16),
torch.bitwise_and(torch.bitwise_right_shift(packed_tensor[:, 11], 12), 8),
),
),
torch.bitwise_or(
torch.bitwise_or(
torch.bitwise_and(torch.bitwise_right_shift(packed_tensor[:, 12], 13), 4),
torch.bitwise_and(torch.bitwise_right_shift(packed_tensor[:, 13], 14), 2),
),
torch.bitwise_and(torch.bitwise_right_shift(packed_tensor[:, 14], 15), 1),
),
)
).unsqueeze(-1)
),
dim=-1
).view(shape)
return result
def unpack_uint14(packed_tensor: torch.Tensor, shape: torch.Size) -> torch.Tensor:
result = torch.cat(
(
torch.bitwise_and(packed_tensor[:, :7], 16383),
torch.bitwise_or(
torch.bitwise_or(
torch.bitwise_or(
torch.bitwise_and(torch.bitwise_right_shift(packed_tensor[:, 0], 2), 12288),
torch.bitwise_and(torch.bitwise_right_shift(packed_tensor[:, 1], 4), 3072),
),
torch.bitwise_or(
torch.bitwise_and(torch.bitwise_right_shift(packed_tensor[:, 2], 6), 768),
torch.bitwise_and(torch.bitwise_right_shift(packed_tensor[:, 3], 8), 192),
),
),
torch.bitwise_or(
torch.bitwise_or(
torch.bitwise_and(torch.bitwise_right_shift(packed_tensor[:, 4], 10), 48),
torch.bitwise_and(torch.bitwise_right_shift(packed_tensor[:, 5], 12), 12),
),
torch.bitwise_and(torch.bitwise_right_shift(packed_tensor[:, 6], 14), 3),
),
).unsqueeze(-1)
),
dim=-1
).view(shape)
return result
def unpack_uint13(packed_tensor: torch.Tensor, shape: torch.Size) -> torch.Tensor:
result = torch.cat(
(
torch.bitwise_and(packed_tensor[:, :13], 8191),
torch.bitwise_or(
torch.bitwise_or(
torch.bitwise_or(
torch.bitwise_and(torch.bitwise_right_shift(packed_tensor[:, :3], 13), 7),
torch.bitwise_and(torch.bitwise_right_shift(packed_tensor[:, 3:6], 10), 56),
),
torch.bitwise_or(
torch.bitwise_and(torch.bitwise_right_shift(packed_tensor[:, 6:9], 7), 448),
torch.bitwise_and(torch.bitwise_right_shift(packed_tensor[:, 9:12], 4), 3584),
),
),
torch.bitwise_and(
torch.stack(
(
torch.bitwise_right_shift(packed_tensor[:, 12], 1),
torch.bitwise_right_shift(packed_tensor[:, 12], 2),
torch.bitwise_right_shift(packed_tensor[:, 12], 3),
),
dim=-1,
),
4096,
)
),
),
dim=-1
).view(shape)
return result
def unpack_uint12(packed_tensor: torch.Tensor, shape: torch.Size) -> torch.Tensor:
result = torch.cat(
(
torch.bitwise_and(packed_tensor[:, :3], 4095),
torch.bitwise_or(
torch.bitwise_or(
torch.bitwise_and(torch.bitwise_right_shift(packed_tensor[:, 0], 4), 3840),
torch.bitwise_and(torch.bitwise_right_shift(packed_tensor[:, 1], 8), 240),
),
torch.bitwise_and(torch.bitwise_right_shift(packed_tensor[:, 2], 12), 15)
).unsqueeze(-1)
),
dim=-1
).view(shape)
return result
def unpack_uint11(packed_tensor: torch.Tensor, shape: torch.Size) -> torch.Tensor:
result = torch.cat(
(
torch.bitwise_and(packed_tensor[:, :8], 2047),
torch.bitwise_or(
torch.bitwise_and(torch.bitwise_right_shift(packed_tensor[:, :8], 11), 31),
torch.bitwise_and(
torch.cat(
(
torch.bitwise_left_shift(packed_tensor[:, 8:], 5),
torch.bitwise_right_shift(packed_tensor[:, 8:], 1),
torch.bitwise_or(
torch.bitwise_and(torch.bitwise_right_shift(packed_tensor[:, 8:10], 7), 480),
torch.bitwise_and(
torch.stack(
(
torch.bitwise_right_shift(packed_tensor[:, 10], 3),
torch.bitwise_right_shift(packed_tensor[:, 10], 5),
),
dim=-1,
),
1536,
),
),
),
dim=-1,
),
2016,
),
),
),
dim=-1
).view(shape)
return result
def unpack_uint10(packed_tensor: torch.Tensor, shape: torch.Size) -> torch.Tensor:
result_bitwise_right_shift = torch.bitwise_and(torch.bitwise_right_shift(packed_tensor[:, :3], 10), 63)
result = torch.cat(
(
torch.bitwise_and(packed_tensor[:, :5], 1023),
torch.bitwise_or(
result_bitwise_right_shift[:, :2],
torch.bitwise_and(torch.bitwise_right_shift(packed_tensor[:, 3:5], 4), 960),
),
torch.bitwise_or(
result_bitwise_right_shift[:, 2],
torch.bitwise_or(
torch.bitwise_and(torch.bitwise_right_shift(packed_tensor[:, 3], 6), 768),
torch.bitwise_and(torch.bitwise_right_shift(packed_tensor[:, 4], 8), 192),
),
).unsqueeze(-1),
),
dim=-1
).view(shape)
return result
def unpack_uint9(packed_tensor: torch.Tensor, shape: torch.Size) -> torch.Tensor:
result = torch.cat(
(
torch.bitwise_and(packed_tensor[:, :8], 511),
torch.bitwise_or(
torch.bitwise_and(torch.bitwise_right_shift(packed_tensor[:, :8], 9), 127),
torch.bitwise_and(
torch.stack(
(
torch.bitwise_left_shift(packed_tensor[:, 8], 7),
torch.bitwise_left_shift(packed_tensor[:, 8], 5),
torch.bitwise_left_shift(packed_tensor[:, 8], 3),
torch.bitwise_left_shift(packed_tensor[:, 8], 1),
torch.bitwise_right_shift(packed_tensor[:, 8], 1),
torch.bitwise_right_shift(packed_tensor[:, 8], 3),
torch.bitwise_right_shift(packed_tensor[:, 8], 5),
torch.bitwise_right_shift(packed_tensor[:, 8], 7),
),
dim=-1,
),
384
)
)
),
dim=-1
).view(shape)
return result
def unpack_uint7(packed_tensor: torch.ByteTensor, shape: torch.Size) -> torch.ByteTensor:
result = torch.cat(
(
torch.bitwise_and(packed_tensor[:, :7], 127),
torch.bitwise_or(
torch.bitwise_or(
torch.bitwise_or(
torch.bitwise_and(torch.bitwise_right_shift(packed_tensor[:, 0], 1), 64),
torch.bitwise_and(torch.bitwise_right_shift(packed_tensor[:, 1], 2), 32),
),
torch.bitwise_or(
torch.bitwise_and(torch.bitwise_right_shift(packed_tensor[:, 2], 3), 16),
torch.bitwise_and(torch.bitwise_right_shift(packed_tensor[:, 3], 4), 8),
),
),
torch.bitwise_or(
torch.bitwise_or(
torch.bitwise_and(torch.bitwise_right_shift(packed_tensor[:, 4], 5), 4),
torch.bitwise_and(torch.bitwise_right_shift(packed_tensor[:, 5], 6), 2),
),
torch.bitwise_right_shift(packed_tensor[:, 6], 7),
),
).unsqueeze(-1)
),
dim=-1
).view(shape)
return result
def unpack_uint6(packed_tensor: torch.ByteTensor, shape: torch.Size) -> torch.ByteTensor:
result = torch.cat(
(
torch.bitwise_and(packed_tensor[:, :3], 63),
torch.bitwise_or(
torch.bitwise_or(
torch.bitwise_and(torch.bitwise_right_shift(packed_tensor[:, 0], 2), 48),
torch.bitwise_and(torch.bitwise_right_shift(packed_tensor[:, 1], 4), 12),
),
torch.bitwise_right_shift(packed_tensor[:, 2], 6)
).unsqueeze(-1)
),
dim=-1
).view(shape)
return result
def unpack_uint5(packed_tensor: torch.ByteTensor, shape: torch.Size) -> torch.ByteTensor:
result_bitwise_right_shift = torch.bitwise_right_shift(packed_tensor[:, :3], 5)
result = torch.cat(
(
torch.bitwise_and(packed_tensor[:, :5], 31),
torch.bitwise_or(
result_bitwise_right_shift[:, :2],
torch.bitwise_and(torch.bitwise_right_shift(packed_tensor[:, 3:5], 2), 24),
),
torch.bitwise_or(
result_bitwise_right_shift[:, 2],
torch.bitwise_or(
torch.bitwise_and(torch.bitwise_right_shift(packed_tensor[:, 3], 3), 16),
torch.bitwise_and(torch.bitwise_right_shift(packed_tensor[:, 4], 4), 8),
),
).unsqueeze(-1),
),
dim=-1
).view(shape)
return result
def unpack_uint4(packed_tensor: torch.ByteTensor, shape: torch.Size) -> torch.ByteTensor:
result = torch.stack((torch.bitwise_and(packed_tensor, 15), torch.bitwise_right_shift(packed_tensor, 4)), dim=-1).view(shape)
return result
def unpack_uint3(packed_tensor: torch.ByteTensor, shape: torch.Size) -> torch.ByteTensor:
result = torch.bitwise_and(
torch.cat(
(
packed_tensor[:, :3],
torch.bitwise_right_shift(packed_tensor[:, :3], 3),
torch.bitwise_or(
torch.bitwise_right_shift(packed_tensor[:, :2], 6),
torch.bitwise_and(
torch.stack(
(
torch.bitwise_right_shift(packed_tensor[:, 2], 4),
torch.bitwise_right_shift(packed_tensor[:, 2], 5),
),
dim=-1
),
4
),
),
),
dim=-1
),
7
).view(shape)
return result
def unpack_uint2(packed_tensor: torch.ByteTensor, shape: torch.Size) -> torch.ByteTensor:
result = torch.bitwise_and(
torch.stack(
(
packed_tensor,
torch.bitwise_right_shift(packed_tensor, 2),
torch.bitwise_right_shift(packed_tensor, 4),
torch.bitwise_right_shift(packed_tensor, 6),
),
dim=-1
),
3
).view(shape)
return result
def unpack_uint1(packed_tensor: torch.Tensor, shape: torch.Size) -> torch.Tensor:
result = torch.bitwise_and(
torch.stack(
(
packed_tensor,
torch.bitwise_right_shift(packed_tensor, 1),
torch.bitwise_right_shift(packed_tensor, 2),
torch.bitwise_right_shift(packed_tensor, 3),
torch.bitwise_right_shift(packed_tensor, 4),
torch.bitwise_right_shift(packed_tensor, 5),
torch.bitwise_right_shift(packed_tensor, 6),
torch.bitwise_right_shift(packed_tensor, 7),
),
dim=-1
),
1
).view(shape)
return result
+22 -9
View File
@@ -44,7 +44,8 @@ def get_scale_symmetric(weight: torch.FloatTensor, reduction_axes: int | list[in
@devices.inference_context()
def quantize_weight(weight: torch.FloatTensor, reduction_axes: int | list[int], weights_dtype: str, dtype: torch.dtype = None, use_stochastic_rounding: bool = False) -> tuple[torch.Tensor, torch.FloatTensor, torch.FloatTensor]:
weight = weight.to(dtype=torch.float32)
if weight.dtype != torch.float64:
weight = weight.to(dtype=torch.float32)
if dtype_dict[weights_dtype]["is_unsigned"]:
scale, zero_point = get_scale_asymmetric(weight, reduction_axes, weights_dtype)
@@ -66,7 +67,8 @@ def quantize_weight(weight: torch.FloatTensor, reduction_axes: int | list[int],
else:
if use_stochastic_rounding:
mantissa_difference = 1 << (23 - dtype_dict[weights_dtype]["mantissa"])
quantized_weight = quantized_weight.view(dtype=torch.int32).add_(torch.randint_like(quantized_weight, low=0, high=mantissa_difference, dtype=torch.int32)).bitwise_and_(-mantissa_difference).view(dtype=torch.float32)
quantized_weight = quantized_weight.to(dtype=torch.float32).view(dtype=torch.int32)
quantized_weight = quantized_weight.add_(torch.randint_like(quantized_weight, low=0, high=mantissa_difference, dtype=torch.int32)).bitwise_and_(-mantissa_difference).view(dtype=torch.float32)
quantized_weight.nan_to_num_()
quantized_weight = quantized_weight.clamp_(dtype_dict[weights_dtype]["min"], dtype_dict[weights_dtype]["max"]).to(dtype_dict[weights_dtype]["torch_dtype"])
return quantized_weight, scale, zero_point
@@ -79,7 +81,8 @@ def apply_svdquant(weight: torch.FloatTensor, rank: int = 32, niter: int = 8, dt
reshape_weight = True
weight_shape = weight.shape
weight = weight.flatten(1,-1)
weight = weight.to(dtype=torch.float32)
if weight.dtype != torch.float64:
weight = weight.to(dtype=torch.float32)
U, S, svd_down = torch.svd_lowrank(weight, q=rank, niter=niter)
svd_up = torch.mul(U, S.unsqueeze(0))
svd_down = svd_down.t_()
@@ -417,7 +420,8 @@ def sdnq_quantize_layer_weight_dynamic(weight, layer_class_name=None, weights_dt
if torch_dtype is None:
torch_dtype = weight.dtype
weights_dtype_order_to_use = weights_dtype_order_fp32 if torch_dtype in {torch.float32, torch.float64} else weights_dtype_order
weight = weight.to(dtype=torch.float32)
if weight.dtype != torch.float64:
weight = weight.to(dtype=torch.float32)
weight_std = weight.std().square_().clamp_(min=1e-8)
if use_svd:
@@ -458,7 +462,7 @@ def sdnq_quantize_layer_weight_dynamic(weight, layer_class_name=None, weights_dt
svd_down = svd_down.t_()
svd_is_transposed = True
quantization_loss = torch.nn.functional.mse_loss(weight, sdnq_dequantizer(quantized_weight, scale, zero_point, svd_up, svd_down, skip_quantized_matmul=sdnq_dequantizer.use_quantized_matmul, dtype=torch.float32, skip_compile=True)).div_(weight_std)
quantization_loss = torch.nn.functional.mse_loss(weight, sdnq_dequantizer(quantized_weight, scale, zero_point, svd_up, svd_down, skip_quantized_matmul=sdnq_dequantizer.use_quantized_matmul, dtype=weight.dtype, skip_compile=True)).div_(weight_std)
if quantization_loss <= dynamic_loss_threshold:
return (quantized_weight, scale, zero_point, svd_up, svd_down, sdnq_dequantizer)
return None
@@ -569,7 +573,7 @@ def apply_sdnq_to_module(model, weights_dtype="int8", quantized_matmul_dtype=Non
if check_param_name_in(param_name, modules_to_not_convert) is not None:
continue
layer_class_name = module.__class__.__name__
if layer_class_name in allowed_types and module.weight.dtype in {torch.float32, torch.float16, torch.bfloat16}:
if layer_class_name in allowed_types and module.weight.dtype in {torch.float64, torch.float32, torch.float16, torch.bfloat16}:
if (layer_class_name in conv_types or layer_class_name in conv_transpose_types) and not quant_conv:
continue
quant_kwargs = {
@@ -811,7 +815,16 @@ class SDNQQuantizer(DiffusersQuantizer, HfQuantizer):
if self.pre_quantized:
layer, tensor_name = get_module_from_name(model, param_name)
if param_value is not None:
return_dtype = param_value.dtype if tensor_name == "weight" else torch.float32 if self.quantization_config.dequantize_fp32 else kwargs.get("dtype", param_value.dtype if self.torch_dtype is None else self.torch_dtype)
if tensor_name == "weight":
return_dtype = param_value.dtype
elif self.quantization_config.dequantize_fp32:
if param_value.dtype != torch.float64 and self.torch_dtype != torch.float64:
return_dtype = torch.float32
else:
return_dtype = torch.float64
else:
return_dtype = kwargs.get("dtype", param_value.dtype if self.torch_dtype is None else self.torch_dtype)
if param_value.dtype == return_dtype and devices.same_device(param_value.device, target_device):
param_value = param_value.clone()
else:
@@ -861,10 +874,10 @@ class SDNQQuantizer(DiffusersQuantizer, HfQuantizer):
}
quant_kwargs = get_quant_kwargs(quant_kwargs, self.quantization_config.modules_quant_config)
if param_value.dtype == torch.float32 and devices.same_device(param_value.device, target_device):
if param_value.dtype in {torch.float32, torch.float64} and devices.same_device(param_value.device, target_device):
param_value = param_value.clone()
else:
param_value = param_value.to(target_device, non_blocking=self.quantization_config.non_blocking).to(dtype=torch.float32)
param_value = param_value.to(target_device, non_blocking=self.quantization_config.non_blocking).to(dtype=torch.float32 if param_value.dtype != torch.float64 else torch.float64)
layer, tensor_name = get_module_from_name(model, param_name)
layer.weight = torch.nn.Parameter(param_value, requires_grad=False)