Skip to content

Commit 46204af

Browse files
committed
Apply isort and black
1 parent f4b03e9 commit 46204af

5 files changed

Lines changed: 11 additions & 12 deletions

File tree

onecomp/quantized_model_loader.py

Lines changed: 1 addition & 1 deletion
Original file line numberDiff line numberDiff line change
@@ -23,8 +23,8 @@
2323
from .quantizer.gptq.config import resolve_gptq_layer_group_size, resolve_gptq_layer_wbits
2424
from .quantizer.gptq.gptq_layer import GPTQLinear
2525
from .quantizer.onebit.onebit_layer import OneBitLinear
26-
from .utils.dtype import needs_bfloat16
2726
from .utils.device import get_default_device
27+
from .utils.dtype import needs_bfloat16
2828
from .utils.quant_config import get_quant_param
2929

3030
logger = getLogger(__name__)

onecomp/quantizer/autobit/activation_stats.py

Lines changed: 1 addition & 1 deletion
Original file line numberDiff line numberDiff line change
@@ -16,7 +16,7 @@
1616
get_blocks_and_inputs,
1717
move_kwargs_to_device,
1818
)
19-
from onecomp.utils.device import get_default_device, empty_cache
19+
from onecomp.utils.device import empty_cache, get_default_device
2020

2121

2222
def _find_head_modules(model, blocks):

onecomp/quantizer/gptq/_gptq.py

Lines changed: 1 addition & 1 deletion
Original file line numberDiff line numberDiff line change
@@ -20,8 +20,8 @@
2020
from transformers import Conv1D
2121

2222
from onecomp.quantizer._quantizer import QuantizationResult, Quantizer
23-
from onecomp.utils.quant_config import get_quant_param
2423
from onecomp.utils.device import empty_cache
24+
from onecomp.utils.quant_config import get_quant_param
2525

2626

2727
@dataclass

onecomp/runner.py

Lines changed: 5 additions & 8 deletions
Original file line numberDiff line numberDiff line change
@@ -330,13 +330,8 @@ def check(self):
330330
device = self.model_config.device
331331
if is_mps_device(device):
332332
if self.multi_gpu:
333-
raise ValueError(
334-
"multi_gpu is not supported on MPS device."
335-
)
336-
all_quantizers = (
337-
self.quantizers if self.quantizers is not None
338-
else [self.quantizer]
339-
)
333+
raise ValueError("multi_gpu is not supported on MPS device.")
334+
all_quantizers = self.quantizers if self.quantizers is not None else [self.quantizer]
340335
for i, q in enumerate(all_quantizers):
341336
label = f"quantizers[{i}]" if self.quantizers else "quantizer"
342337
if isinstance(q, AutoBitQuantizer):
@@ -583,7 +578,9 @@ def auto_run(
583578
enable_fused_groups=True,
584579
)
585580
qep_config = QEPConfig(device=device)
586-
runner = cls(model_config=model_config, quantizer=quantizer, qep=qep, qep_config=qep_config)
581+
runner = cls(
582+
model_config=model_config, quantizer=quantizer, qep=qep, qep_config=qep_config
583+
)
587584
runner.run()
588585

589586
if evaluate:

onecomp/utils/perplexity.py

Lines changed: 3 additions & 1 deletion
Original file line numberDiff line numberDiff line change
@@ -125,7 +125,9 @@ def calculate_perplexity(
125125
if stride is None:
126126
stride = max_length
127127
seq_len = encodings.input_ids.size(1)
128-
use_cpu_accum = device.type == "mps" if isinstance(device, torch.device) else str(device).startswith("mps")
128+
use_cpu_accum = (
129+
device.type == "mps" if isinstance(device, torch.device) else str(device).startswith("mps")
130+
)
129131
accum_device = torch.device("cpu") if use_cpu_accum else device
130132
nll_sum = torch.tensor(0.0, dtype=torch.float64, device=accum_device)
131133
n_tokens = 0

0 commit comments

Comments
 (0)