@@ -100,9 +100,7 @@ def forward(self, tensor):
100100 sin_inp_y = torch .einsum ("i,j->ij" , pos_y , self .inv_freq )
101101 emb_x = get_emb (sin_inp_x ).unsqueeze (1 )
102102 emb_y = get_emb (sin_inp_y )
103- emb = torch .zeros ((x , y , self .channels * 2 ), device = tensor .device ).type (
104- tensor .type ()
105- )
103+ emb = torch .zeros ((x , y , self .channels * 2 ), device = tensor .device ).type (tensor .type ())
106104 emb [:, :, : self .channels ] = emb_x
107105 emb [:, :, self .channels : 2 * self .channels ] = emb_y
108106
@@ -165,9 +163,7 @@ def forward(self, tensor):
165163 emb_x = get_emb (sin_inp_x ).unsqueeze (1 ).unsqueeze (1 )
166164 emb_y = get_emb (sin_inp_y ).unsqueeze (1 )
167165 emb_z = get_emb (sin_inp_z )
168- emb = torch .zeros ((x , y , z , self .channels * 3 ), device = tensor .device ).type (
169- tensor .type ()
170- )
166+ emb = torch .zeros ((x , y , z , self .channels * 3 ), device = tensor .device ).type (tensor .type ())
171167 emb [:, :, :, : self .channels ] = emb_x
172168 emb [:, :, :, self .channels : 2 * self .channels ] = emb_y
173169 emb [:, :, :, 2 * self .channels :] = emb_z
0 commit comments