Commit 1d8511f8 authored by VictorSanh's avatar VictorSanh
Browse files

FIX small bugs in `run_classifier_pytorch.py`

parent 936eb4c3
......@@ -512,7 +512,7 @@ def main():
train_dataloader = DataLoader(train_data, sampler=train_sampler, batch_size=args.train_batch_size)
model.train()
for epoch in range(args.num_train_epochs):
for epoch in range(int(args.num_train_epochs)):
for input_ids, input_mask, segment_ids, label_ids in train_dataloader:
input_ids = input_ids.to(device)
input_mask = input_mask.float().to(device)
......
Markdown is supported
0% or .
You are about to add 0 people to the discussion. Proceed with caution.
Finish editing this message first!
Please register or to comment