Skip to content
Merged
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
Original file line number Diff line number Diff line change
Expand Up @@ -142,7 +142,12 @@ def __call__(
self.bitstream_name,
file_prefix,
)
self.update_time_elapsed("encode", (time_measure() - start))
elapsed_enc_time = time_measure() - start
total_enc_runtime = sum(enc_time_by_module.values())
total_enc_time = (
total_enc_runtime if total_enc_runtime != 0 else elapsed_enc_time
)
self.update_time_elapsed("encode", (total_enc_time))
if self.is_mac_calculation:
self.acc_kmac_and_pixels_info(
"feature_reduction", enc_complexity[0], enc_complexity[1]
Expand Down Expand Up @@ -183,7 +188,12 @@ def __call__(
dec_features, dec_time_by_module, dec_complexity = self._decompress(
codec, res["bitstream"], self.codec_output_dir, file_prefix
)
self.update_time_elapsed("decode", (time_measure() - start))
elapsed_dec_time = time_measure() - start
total_dec_runtime = sum(dec_time_by_module.values())
total_dec_time = (
total_dec_runtime if total_dec_runtime != 0 else elapsed_dec_time
)
self.update_time_elapsed("decode", (total_dec_time))
if self.is_mac_calculation:
self.acc_kmac_and_pixels_info(
"feature_restoration", dec_complexity[0], dec_complexity[1]
Expand Down
Original file line number Diff line number Diff line change
Expand Up @@ -247,7 +247,12 @@ def __call__(
res, enc_time_by_module, enc_complexity = self._compress(
codec, features, self.codec_output_dir, self.bitstream_name, ""
)
self.update_time_elapsed("encode", (time_measure() - start))
elapsed_enc_time = time_measure() - start
total_enc_runtime = sum(enc_time_by_module.values())
total_enc_time = (
total_enc_runtime if total_enc_runtime != 0 else elapsed_enc_time
)
self.update_time_elapsed("encode", (total_enc_time))
self.add_time_details("encode", enc_time_by_module)
if self.is_mac_calculation:
self.add_kmac_and_pixels_info(
Expand Down Expand Up @@ -286,7 +291,12 @@ def __call__(
dec_features, dec_time_by_module, dec_complexity = self._decompress(
codec, res["bitstream"], self.codec_output_dir, ""
)
self.update_time_elapsed("decode", (time_measure() - start))
elapsed_dec_time = time_measure() - start
total_dec_runtime = sum(dec_time_by_module.values())
total_dec_time = (
total_dec_runtime if total_dec_runtime != 0 else elapsed_dec_time
)
self.update_time_elapsed("decode", (total_dec_time))
self.add_time_details("decode", dec_time_by_module)
if self.is_mac_calculation:
self.add_kmac_and_pixels_info(
Expand Down
29 changes: 29 additions & 0 deletions compressai_vision/utils/measure_complexity.py
Original file line number Diff line number Diff line change
Expand Up @@ -203,6 +203,24 @@ def _flops_to_kmacs(total_flops: float) -> float:
return float(total_flops) / 1e3


# (mem-fix) KMACs are deterministic w.r.t. (module, input shapes), so trace each
# shape only once. Per-image FlopCountAnalysis (torch.jit tracing) retains feature-map
# tensors and steadily grows CPU RAM (OOM); caching avoids that and speeds up
# repeated measurement.
_KMACS_CACHE: dict = {}


def _shape_key(x):
"""Build a hashable shape key from inputs (possibly nested)."""
if torch.is_tensor(x):
return ("t", tuple(x.shape))
if isinstance(x, (list, tuple)):
return ("s", tuple(_shape_key(v) for v in x))
if isinstance(x, dict):
return ("d", tuple((k, _shape_key(v)) for k, v in x.items()))
return ("o", repr(x))


def measure_kmacs(module: nn.Module, inputs, tag: str = None) -> float:
"""
Measure KMACs for a given module using fvcore.
Expand Down Expand Up @@ -249,6 +267,15 @@ def _cast(x):
elif not isinstance(inputs, tuple):
inputs = (inputs,) # safe fallback

# (mem-fix) Shape-based cache lookup: skip re-tracing for the same (module, input shapes).
cache_key = None
try:
cache_key = (module.__class__.__qualname__, id(p), _shape_key(inputs))
except Exception:
cache_key = None
if cache_key is not None and cache_key in _KMACS_CACHE:
return _KMACS_CACHE[cache_key]

with torch.no_grad():
flops = FlopCountAnalysis(module, inputs)
flops.set_op_handle(
Expand Down Expand Up @@ -280,6 +307,8 @@ def _cast(x):
kmacs = _flops_to_kmacs(total_flops)
name = tag or module.__class__.__name__
# print(f"[INFO] {name}: KMACs = {kmacs}")
if cache_key is not None:
_KMACS_CACHE[cache_key] = kmacs
return kmacs


Expand Down