I have some trouble understanding why the computation time spikes significantly in my program.
I am trying to obtain some indices and I want to time which one is faster. The following code is part of a class:
def _obtain_indices(self):
"""
Set the indices of the object for all
indices that exceed a threshold.
Parameters
----------
tensor: array like
Tensor of which to select elements
Returns
-------
None
"""
start = datetime.now()
abs_tens = torch.abs(self.residual_grad)
end = datetime.now()
t_abs = (end - start)/timedelta(milliseconds=1)
# Func gauss
start = datetime.now()
self.indices = self.compress(abs_tens, self.K)
end = datetime.now()
t_gauss = end - start
t_gauss = t_gauss/timedelta(milliseconds=1)
# Func top_k
start = datetime.now()
self.indices = torch.topk(abs_tens, int(self.K))[1]
end = datetime.now()
t_topk = end - start
t_topk = t_topk/timedelta(milliseconds=1)
wandb.log({'t_abs' : t_abs, 't_gauss': t_gauss, 'top-k': t_topk})
Now there's some trouble here. wandb.log() logs the values.
I have some example output graphs:

This is with the original code above. When I switch the functions (TopK first, Gausssecond) I get the following graph:
What is happening? I was thinking maybe something with heap allocation, but I'm not sure and feel unqualified to make statements about it. Can somebody point me in the right direction?


