Skip to content

Commit 833a91d

Browse files
committed
supporting pytorch 1.1.0
1 parent f6ba20b commit 833a91d

2 files changed

Lines changed: 9 additions & 9 deletions

File tree

README.md

Lines changed: 1 addition & 1 deletion
Original file line numberDiff line numberDiff line change
@@ -41,7 +41,7 @@ Figure 2: CondenseNets with Fully Dense Connectivity and Increasing Growth Rate.
4141
### Dependencies
4242

4343
- [Python3](https://www.python.org/downloads/)
44-
- [PyTorch(0.1.12+)](http://pytorch.org)
44+
- [PyTorch(1.1.0)](http://pytorch.org)
4545
- [ImageNet](https://www.image-net.org/challenges/LSVRC/2012/)
4646

4747
### Train

main.py

Lines changed: 8 additions & 8 deletions
Original file line numberDiff line numberDiff line change
@@ -295,7 +295,7 @@ def train(train_loader, model, criterion, optimizer, epoch):
295295
### Measure data loading time
296296
data_time.update(time.time() - end)
297297

298-
target = target.cuda(async=True)
298+
target = target.cuda(non_blocking=True)
299299
input_var = torch.autograd.Variable(input)
300300
target_var = torch.autograd.Variable(target)
301301

@@ -312,9 +312,9 @@ def train(train_loader, model, criterion, optimizer, epoch):
312312

313313
### Measure accuracy and record loss
314314
prec1, prec5 = accuracy(output.data, target, topk=(1, 5))
315-
losses.update(loss.data[0], input.size(0))
316-
top1.update(prec1[0], input.size(0))
317-
top5.update(prec5[0], input.size(0))
315+
losses.update(loss.item(), input.size(0))
316+
top1.update(prec1.item(), input.size(0))
317+
top5.update(prec5.item(), input.size(0))
318318

319319
### Compute gradient and do SGD step
320320
optimizer.zero_grad()
@@ -349,7 +349,7 @@ def validate(val_loader, model, criterion):
349349

350350
end = time.time()
351351
for i, (input, target) in enumerate(val_loader):
352-
target = target.cuda(async=True)
352+
target = target.cuda(non_blocking=True)
353353
input_var = torch.autograd.Variable(input, volatile=True)
354354
target_var = torch.autograd.Variable(target, volatile=True)
355355

@@ -359,9 +359,9 @@ def validate(val_loader, model, criterion):
359359

360360
### Measure accuracy and record loss
361361
prec1, prec5 = accuracy(output.data, target, topk=(1, 5))
362-
losses.update(loss.data[0], input.size(0))
363-
top1.update(prec1[0], input.size(0))
364-
top5.update(prec5[0], input.size(0))
362+
losses.update(loss.data.item(), input.size(0))
363+
top1.update(prec1.item(), input.size(0))
364+
top5.update(prec5.item(), input.size(0))
365365

366366
### Measure elapsed time
367367
batch_time.update(time.time() - end)

0 commit comments

Comments
 (0)