diff --git a/modules/sdnq/common.py b/modules/sdnq/common.py index 6e8951505..40e7de200 100644 --- a/modules/sdnq/common.py +++ b/modules/sdnq/common.py @@ -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", ] diff --git a/modules/sdnq/dequantizer.py b/modules/sdnq/dequantizer.py index dae9bdfea..ef65963c7 100644 --- a/modules/sdnq/dequantizer.py +++ b/modules/sdnq/dequantizer.py @@ -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() diff --git a/modules/sdnq/layers/conv/conv_fp16.py b/modules/sdnq/layers/conv/conv_fp16.py index 2f5254eb5..31beb017e 100644 --- a/modules/sdnq/layers/conv/conv_fp16.py +++ b/modules/sdnq/layers/conv/conv_fp16.py @@ -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 diff --git a/modules/sdnq/layers/conv/conv_fp8.py b/modules/sdnq/layers/conv/conv_fp8.py index f1b478891..3595a5366 100644 --- a/modules/sdnq/layers/conv/conv_fp8.py +++ b/modules/sdnq/layers/conv/conv_fp8.py @@ -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) diff --git a/modules/sdnq/layers/conv/conv_fp8_tensorwise.py b/modules/sdnq/layers/conv/conv_fp8_tensorwise.py index ea24bc001..38258bff7 100644 --- a/modules/sdnq/layers/conv/conv_fp8_tensorwise.py +++ b/modules/sdnq/layers/conv/conv_fp8_tensorwise.py @@ -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) diff --git a/modules/sdnq/layers/linear/linear_fp16.py b/modules/sdnq/layers/linear/linear_fp16.py index 3e04d3be6..2999d09cb 100644 --- a/modules/sdnq/layers/linear/linear_fp16.py +++ b/modules/sdnq/layers/linear/linear_fp16.py @@ -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 diff --git a/modules/sdnq/layers/linear/linear_fp8.py b/modules/sdnq/layers/linear/linear_fp8.py index a2105020e..c0b005b75 100644 --- a/modules/sdnq/layers/linear/linear_fp8.py +++ b/modules/sdnq/layers/linear/linear_fp8.py @@ -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]) diff --git a/modules/sdnq/layers/linear/linear_fp8_tensorwise.py b/modules/sdnq/layers/linear/linear_fp8_tensorwise.py index 1d1f894e4..9977fbe7c 100644 --- a/modules/sdnq/layers/linear/linear_fp8_tensorwise.py +++ b/modules/sdnq/layers/linear/linear_fp8_tensorwise.py @@ -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]) diff --git a/modules/sdnq/loader.py b/modules/sdnq/loader.py index 789fd6e96..ff55c0453 100644 --- a/modules/sdnq/loader.py +++ b/modules/sdnq/loader.py @@ -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: diff --git a/modules/sdnq/packed_float.py b/modules/sdnq/packed_float.py index 66b03c873..6e9c7a0f9 100644 --- a/modules/sdnq/packed_float.py +++ b/modules/sdnq/packed_float.py @@ -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"] diff --git a/modules/sdnq/packed_int.py b/modules/sdnq/packed_int.py deleted file mode 100644 index f5e92b53f..000000000 --- a/modules/sdnq/packed_int.py +++ /dev/null @@ -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"] diff --git a/modules/sdnq/packed_int/__init__.py b/modules/sdnq/packed_int/__init__.py new file mode 100644 index 000000000..ddd6d171d --- /dev/null +++ b/modules/sdnq/packed_int/__init__.py @@ -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 diff --git a/modules/sdnq/packed_int/pack.py b/modules/sdnq/packed_int/pack.py new file mode 100644 index 000000000..839398eda --- /dev/null +++ b/modules/sdnq/packed_int/pack.py @@ -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 diff --git a/modules/sdnq/packed_int/unpack.py b/modules/sdnq/packed_int/unpack.py new file mode 100644 index 000000000..ff76482f7 --- /dev/null +++ b/modules/sdnq/packed_int/unpack.py @@ -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 diff --git a/modules/sdnq/quantizer.py b/modules/sdnq/quantizer.py index 61f17f84f..a9a8cd733 100644 --- a/modules/sdnq/quantizer.py +++ b/modules/sdnq/quantizer.py @@ -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)