@@ -30,7 +30,7 @@ def assign_labels(
3030 n_neurons = spikes .size (2 )
3131
3232 if rates is None :
33- rates = torch .zeros (n_neurons , n_labels )
33+ rates = torch .zeros (( n_neurons , n_labels ), device = spikes . device )
3434
3535 # Sum over time dimension (spike ordering doesn't matter).
3636 spikes = spikes .sum (1 )
@@ -112,7 +112,7 @@ def all_activity(
112112 # Sum over time dimension (spike ordering doesn't matter).
113113 spikes = spikes .sum (1 )
114114
115- rates = torch .zeros (n_samples , n_labels )
115+ rates = torch .zeros (( n_samples , n_labels ), device = spikes . device )
116116 for i in range (n_labels ):
117117 # Count the number of neurons with this label assignment.
118118 n_assigns = torch .sum (assignments == i ).float ()
@@ -153,7 +153,7 @@ def proportion_weighting(
153153 # Sum over time dimension (spike ordering doesn't matter).
154154 spikes = spikes .sum (1 )
155155
156- rates = torch .zeros (n_samples , n_labels )
156+ rates = torch .zeros (( n_samples , n_labels ), device = spikes . device )
157157 for i in range (n_labels ):
158158 # Count the number of neurons with this label assignment.
159159 n_assigns = torch .sum (assignments == i ).float ()
@@ -191,7 +191,7 @@ def ngram(
191191 """
192192 predictions = []
193193 for activity in spikes :
194- score = torch .zeros (n_labels )
194+ score = torch .zeros (n_labels , device = spikes . device )
195195
196196 # Aggregate all of the firing neurons' indices
197197 fire_order = []
0 commit comments