From 1ab7159cbe115af1510a13c8df4ab3e1f4496c00 Mon Sep 17 00:00:00 2001 From: Heeji Han Date: Sat, 1 Aug 2026 05:06:07 +0000 Subject: [PATCH 1/2] [fix] handle scalar KMAC returns in YOLOX --- compressai_vision/utils/measure_complexity.py | 4 ++-- 1 file changed, 2 insertions(+), 2 deletions(-) diff --git a/compressai_vision/utils/measure_complexity.py b/compressai_vision/utils/measure_complexity.py index 8798623..add6455 100644 --- a/compressai_vision/utils/measure_complexity.py +++ b/compressai_vision/utils/measure_complexity.py @@ -719,7 +719,7 @@ def calc_complexity_nn_part1_yolox(vision_model, img): C, H, W = img.shape[1:] - kmacs, _ = measure_kmacs(partial_model, img) + kmacs = measure_kmacs(partial_model, img) pixels = reduce(operator.mul, [p_size for p_size in img.shape]) return kmacs, pixels @@ -743,7 +743,7 @@ def calc_complexity_nn_part2_yolox(vision_model, dec_features): C, H, W = input_tensor.shape[1:] partial_model = YoloxPart2(vision_model, vision_model.split_id) - kmacs, _ = measure_kmacs(partial_model, input_tensor) + kmacs = measure_kmacs(partial_model, input_tensor) pixels = reduce(operator.mul, input_tensor.shape) From 5b54bd487804b8d7b8882a0c7cebba267b038176 Mon Sep 17 00:00:00 2001 From: Heeji Han Date: Sat, 1 Aug 2026 05:30:06 +0000 Subject: [PATCH 2/2] [fix] preserve CUDA device context for YOLOX inference and MAC computation --- compressai_vision/model_wrappers/yolox.py | 30 ++++++++++++------- compressai_vision/utils/measure_complexity.py | 12 ++++++-- 2 files changed, 29 insertions(+), 13 deletions(-) diff --git a/compressai_vision/model_wrappers/yolox.py b/compressai_vision/model_wrappers/yolox.py index 74ba662..54c2566 100644 --- a/compressai_vision/model_wrappers/yolox.py +++ b/compressai_vision/model_wrappers/yolox.py @@ -28,6 +28,7 @@ # ADVISED OF THE POSSIBILITY OF SUCH DAMAGE. +from contextlib import nullcontext from enum import Enum from pathlib import Path from typing import Dict, List @@ -184,19 +185,26 @@ def input_to_features(self, x, device: str) -> Dict: def features_to_output(self, x: Dict, device: str): """Complete the downstream task from the intermediate deep features""" - self.model = self.model.to(device).eval() + target_device = torch.device(device) + device_context = ( + torch.cuda.device(target_device) + if target_device.type == "cuda" + else nullcontext() + ) - if self.split_id == self.SPLIT_L13: - return self._feature_at_l13_to_output( - x["data"], x["org_input_size"], x["input_size"], device - ) - elif self.split_id == self.SPLIT_L37: - return self._feature_at_l37_to_output( - x["data"], x["org_input_size"], x["input_size"], device - ) - else: - self.logger.error(f"Not supported split points {self.split_id}") + with device_context: + self.model = self.model.to(target_device).eval() + + if self.split_id == self.SPLIT_L13: + return self._feature_at_l13_to_output( + x["data"], x["org_input_size"], x["input_size"], target_device + ) + if self.split_id == self.SPLIT_L37: + return self._feature_at_l37_to_output( + x["data"], x["org_input_size"], x["input_size"], target_device + ) + self.logger.error(f"Not supported split points {self.split_id}") raise NotImplementedError @torch.no_grad() diff --git a/compressai_vision/utils/measure_complexity.py b/compressai_vision/utils/measure_complexity.py index add6455..052b1dd 100644 --- a/compressai_vision/utils/measure_complexity.py +++ b/compressai_vision/utils/measure_complexity.py @@ -719,7 +719,11 @@ def calc_complexity_nn_part1_yolox(vision_model, img): C, H, W = img.shape[1:] - kmacs = measure_kmacs(partial_model, img) + if img.is_cuda: + with torch.cuda.device(img.device): + kmacs = measure_kmacs(partial_model, img) + else: + kmacs = measure_kmacs(partial_model, img) pixels = reduce(operator.mul, [p_size for p_size in img.shape]) return kmacs, pixels @@ -743,7 +747,11 @@ def calc_complexity_nn_part2_yolox(vision_model, dec_features): C, H, W = input_tensor.shape[1:] partial_model = YoloxPart2(vision_model, vision_model.split_id) - kmacs = measure_kmacs(partial_model, input_tensor) + if input_tensor.is_cuda: + with torch.cuda.device(input_tensor.device): + kmacs = measure_kmacs(partial_model, input_tensor) + else: + kmacs = measure_kmacs(partial_model, input_tensor) pixels = reduce(operator.mul, input_tensor.shape)