Skip to content

Commit 2d09181

Browse files
Reversed to no_grad for head_init_scale
1 parent 8630947 commit 2d09181

File tree

1 file changed

+1
-1
lines changed

1 file changed

+1
-1
lines changed

train.py

Lines changed: 1 addition & 1 deletion
Original file line numberDiff line numberDiff line change
@@ -512,7 +512,7 @@ def main():
512512
**args.model_kwargs,
513513
)
514514
if args.head_init_scale is not None:
515-
with torch.inference_mode():
515+
with torch.no_grad():
516516
model.get_classifier().weight.mul_(args.head_init_scale)
517517
model.get_classifier().bias.mul_(args.head_init_scale)
518518
if args.head_init_bias is not None:

0 commit comments

Comments
 (0)