Skip to content

Commit 0f71774

Browse files
Fix missing compute energy in energy() call
1 parent 68079b9 commit 0f71774

1 file changed

Lines changed: 6 additions & 2 deletions

File tree

accelforge/mapper/FFM/mappings.py

Lines changed: 6 additions & 2 deletions
Original file line numberDiff line numberDiff line change
@@ -1,3 +1,4 @@
1+
from accelforge.frontend.renames import TensorName
12
from collections import defaultdict
23
from accelforge.frontend import arch
34
from accelforge.frontend.spec import Spec
@@ -430,7 +431,10 @@ def energy(
430431
result = {}
431432
for einsum in self.einsum_names:
432433
einsum_accessed = energy.access(einsum, col_idx=0)
433-
for tensor in self.spec.workload.einsums[einsum].tensor_names:
434+
# None for computes
435+
for tensor in list[TensorName](
436+
self.spec.workload.einsums[einsum].tensor_names
437+
) + ["None"]:
434438
tensor_accessed = einsum_accessed.access(tensor, col_idx=1)
435439
for col in tensor_accessed._get_keys_of_length(2):
436440
component, action = col.split("<SEP>")
@@ -448,7 +452,7 @@ def energy(
448452
if not keep_indices:
449453
v = sum(result.values())
450454
if value_if_one_mapping and len(self.data) == 1:
451-
return v.iloc[0]
455+
return v.iloc[0] if isinstance(v, pd.Series) else v
452456
return v
453457

454458
new_result = defaultdict(float)

0 commit comments

Comments
 (0)