Skip to content

Commit e8c1cf0

Browse files
BowenBaofacebook-github-bot
authored andcommitted
Remove unnecessary model attribute assignment on 'freqs_cis' (#1766)
Summary: Fixes #1767, more details there. Upstream PR link meta-llama/llama#349 Pull Request resolved: #1766 Reviewed By: msaroufim Differential Revision: D47556245 Pulled By: xuzhao9 fbshipit-source-id: 5ad541ff3e38a36b6515b1b7135b3da92109c9f5
1 parent 411e388 commit e8c1cf0

File tree

1 file changed

+3
-2
lines changed

1 file changed

+3
-2
lines changed

torchbenchmark/models/llama/model.py

+3-2
Original file line numberDiff line numberDiff line change
@@ -224,8 +224,9 @@ def forward(self, tokens: torch.Tensor, start_pos: int):
224224

225225
h = self.tok_embeddings(tokens)
226226

227-
self.freqs_cis = self.freqs_cis.to(h.device)
228-
freqs_cis = self.freqs_cis[start_pos : start_pos + seqlen]
227+
# Reference: https://github.com/facebookresearch/llama/pull/349
228+
freqs_cis = self.freqs_cis.to(h.device)
229+
freqs_cis = freqs_cis[start_pos : start_pos + seqlen]
229230

230231
mask = None
231232

0 commit comments

Comments
 (0)