Fix gemlite import (#2553)
This commit is contained in:
@@ -2,8 +2,14 @@
|
||||
Common utilities for torchao.
|
||||
"""
|
||||
|
||||
import logging
|
||||
import os
|
||||
import pwd
|
||||
|
||||
import torch
|
||||
|
||||
logger = logging.getLogger(__name__)
|
||||
|
||||
|
||||
def apply_torchao_config_to_model(
|
||||
model: torch.nn.Module, torchao_config: str, filter_fn=None
|
||||
@@ -50,27 +56,17 @@ def apply_torchao_config_to_model(
|
||||
elif "gemlite" in torchao_config:
|
||||
# gemlite-<packing_bitwidth>-<bit_width>-<group_size> or
|
||||
# gemlite-<bit_width>-<group_size> (packing_bitwidth defaults to 32)
|
||||
import os
|
||||
import pwd
|
||||
|
||||
import gemlite
|
||||
from gemlite.core import GemLiteLinearTriton, set_autotune
|
||||
|
||||
try:
|
||||
from torchao.quantization import gemlite_uintx_weight_only
|
||||
except:
|
||||
print(
|
||||
f"import `gemlite_uintx_weight_only` failed, please use torchao nightly to use gemlite quantization"
|
||||
)
|
||||
return model
|
||||
from gemlite.core import GemLiteLinearTriton
|
||||
from torchao.quantization import gemlite_uintx_weight_only
|
||||
|
||||
_quant_args = torchao_config.split("-")
|
||||
bit_width = int(_quant_args[-2])
|
||||
group_size = None if _quant_args[-1] == "None" else int(_quant_args[-1])
|
||||
|
||||
try:
|
||||
packing_bitwidth = int(_quant_args[-3])
|
||||
except:
|
||||
# if only 2 inputs found, use default value
|
||||
except (ValueError, IndexError):
|
||||
# if only 2 inputs found or conversion fails, use default value
|
||||
packing_bitwidth = 32
|
||||
|
||||
quantize_(
|
||||
|
||||
Reference in New Issue
Block a user