Skip to content

Commit ede76f8

Browse files
authored
Update cky.py
1 parent 7013db4 commit ede76f8

File tree

1 file changed

+1
-1
lines changed

1 file changed

+1
-1
lines changed

torch_struct/cky.py

Lines changed: 1 addition & 1 deletion
Original file line numberDiff line numberDiff line change
@@ -23,7 +23,7 @@ def _dp(self, scores, lengths=None, force_grad=False, cache=True):
2323
semiring.convert(roots).requires_grad_(True),
2424
)
2525
if lengths is None:
26-
lengths = torch.LongTensor([N] * batch)
26+
lengths = torch.LongTensor([N] * batch).to(terms.device)
2727

2828
# Charts
2929
beta = [

0 commit comments

Comments
 (0)