@@ -255,17 +255,17 @@ def __init__(self, num_tags: int = None):
255255 self .b_start = None # Start boundary energy
256256 self .b_end = None # End boundary energy
257257
258- def build (self , num_tags : int ):
258+ def build (self , num_tags : int , device = None , dtype = None ):
259259 """Initialize layer weights."""
260260 self ._num_tags = num_tags
261261
262262 # Transition matrix (energy between tag pairs)
263- self .U = nn .Parameter (torch .empty (num_tags , num_tags ))
263+ self .U = nn .Parameter (torch .empty (num_tags , num_tags , device = device , dtype = dtype ))
264264 nn .init .xavier_uniform_ (self .U )
265265
266266 # Boundary energies
267- self .b_start = nn .Parameter (torch .zeros (num_tags ))
268- self .b_end = nn .Parameter (torch .zeros (num_tags ))
267+ self .b_start = nn .Parameter (torch .zeros (num_tags , device = device , dtype = dtype ))
268+ self .b_end = nn .Parameter (torch .zeros (num_tags , device = device , dtype = dtype ))
269269
270270 self ._built = True
271271
@@ -288,7 +288,7 @@ def forward(
288288 During inference: Best tag sequence [batch_size, seq_len]
289289 """
290290 if not self ._built :
291- self .build (emissions .size (- 1 ))
291+ self .build (emissions .size (- 1 ), device = emissions . device , dtype = emissions . dtype )
292292
293293 if tags is not None :
294294 # Training: compute loss
0 commit comments