from __future__ import print_function
import errno
import os
import numpy as np
# from PIL import Image
import torch
import [Link] as nn
EPS = 1e-7
def assert_eq(real, expected):
assert real == expected, '%s (true) vs %s (expected)' % (real, expected)
def assert_array_eq(real, expected):
assert ([Link](real-expected) < EPS).all(), \
'%s (true) vs %s (expected)' % (real, expected)
def load_folder(folder, suffix):
imgs = []
for f in sorted([Link](folder)):
if [Link](suffix):
[Link]([Link](folder, f))
return imgs
# def load_imageid(folder):
# images = load_folder(folder, 'jpg')
# img_ids = set()
# for img in images:
# img_id = int([Link]('/')[-1].split('.')[0].split('_')[-1])
# img_ids.add(img_id)
# return img_ids
# def pil_loader(path):
# with open(path, 'rb') as f:
# with [Link](f) as img:
# return [Link]('RGB')
def weights_init(m):
"""custom weights initialization."""
cname = m.__class__
if cname == [Link] or cname == nn.Conv2d or cname == nn.ConvTranspose2d:
[Link].normal_(0.0, 0.02)
elif cname == nn.BatchNorm2d:
[Link].normal_(1.0, 0.02)
[Link].fill_(0)
else:
print('%s is not initialized.' % cname)
def init_net(net, net_file):
if net_file:
net.load_state_dict([Link](net_file))
else:
[Link](weights_init)
def create_dir(path):
if not [Link](path):
try:
[Link](path)
except OSError as exc:
if [Link] != [Link]:
raise
class Logger(object):
def __init__(self, output_name):
dirname = [Link](output_name)
if not [Link](dirname):
[Link](dirname)
self.log_file = open(output_name, 'w')
[Link] = {}
def append(self, key, val):
vals = [Link](key, [])
[Link](val)
def log(self, extra_msg=''):
msgs = [extra_msg]
for key, vals in [Link]():
[Link]('%s %.6f' % (key, [Link](vals)))
msg = '\n'.join(msgs)
self.log_file.write(msg + '\n')
self.log_file.flush()
[Link] = {}
return msg
def write(self, msg):
self.log_file.write(msg + '\n')
self.log_file.flush()
print(msg)