Skip to content
Open
Show file tree
Hide file tree
Changes from all commits
Commits
File filter

Filter by extension

Filter by extension

Conversations
Failed to load comments.
Loading
Jump to
Jump to file
Failed to load files.
Loading
Diff view
Diff view
30 changes: 19 additions & 11 deletions compressai_vision/model_wrappers/yolox.py
Original file line number Diff line number Diff line change
Expand Up @@ -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
Expand Down Expand Up @@ -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()
Expand Down
12 changes: 10 additions & 2 deletions compressai_vision/utils/measure_complexity.py
Original file line number Diff line number Diff line change
Expand Up @@ -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
Expand All @@ -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)

Expand Down