mirror of
https://github.com/igv/FSRCNN-TensorFlow.git
synced 2026-10-06 17:04:12 +08:00
README and is_train -> train
This commit is contained in:
@@ -16,7 +16,7 @@ Can specify epochs, learning rate, data directory, etc:
|
||||
`python main.py --epochs 10 --learning_rate 0.0001 --data_dir Train`
|
||||
<br>
|
||||
<br>
|
||||
For testing: `python main.py --is_train False`
|
||||
For testing: `python main.py --train False`
|
||||
|
||||
To use FSCRNN-s instead of FSCRNN: `python main.py --fast True`
|
||||
|
||||
|
||||
@@ -18,7 +18,7 @@ flags.DEFINE_integer("stride", 4, "The size of stride to apply to input image [4
|
||||
flags.DEFINE_string("checkpoint_dir", "checkpoint", "Name of checkpoint directory [checkpoint]")
|
||||
flags.DEFINE_string("output_dir", "result", "Name of test output directory [result]")
|
||||
flags.DEFINE_string("data_dir", "FastTrain", "Name of data directory to train on [FastTrain]")
|
||||
flags.DEFINE_boolean("is_train", True, "True for training, False for testing [True]")
|
||||
flags.DEFINE_boolean("train", True, "True for training, false for testing [True]")
|
||||
flags.DEFINE_integer("threads", 1, "Number of processes to pre-process data with [1]")
|
||||
flags.DEFINE_boolean("params", False, "Save weight and bias parameters [False]")
|
||||
|
||||
|
||||
@@ -23,7 +23,7 @@ class FSRCNN(object):
|
||||
def __init__(self, sess, config):
|
||||
self.sess = sess
|
||||
self.fast = config.fast
|
||||
self.is_train = config.is_train
|
||||
self.train = config.train
|
||||
self.c_dim = config.c_dim
|
||||
self.is_grayscale = (self.c_dim == 1)
|
||||
self.epoch = config.epoch
|
||||
@@ -38,7 +38,7 @@ class FSRCNN(object):
|
||||
# Different image/label sub-sizes for different scaling factors x2, x3, x4
|
||||
scale_factors = [[10, 20], [11, 21], [6, 24]]
|
||||
self.image_size, self.label_size = scale_factors[self.scale - 2]
|
||||
if not self.is_train:
|
||||
if not self.train:
|
||||
self.stride = [10, 7, 6][self.scale - 2]
|
||||
|
||||
# Different model layer counts/filter sizes for FSRCNN vs FSRCNN-s (fast)
|
||||
@@ -102,7 +102,7 @@ class FSRCNN(object):
|
||||
|
||||
if self.params:
|
||||
save_params(self.sess, self.weights, self.biases)
|
||||
elif self.is_train:
|
||||
elif self.train:
|
||||
self.train()
|
||||
else:
|
||||
self.test()
|
||||
|
||||
@@ -65,7 +65,7 @@ def prepare_data(sess, dataset):
|
||||
|
||||
For train dataset, output data would be ['.../t1.bmp', '.../t2.bmp', ..., '.../t99.bmp']
|
||||
"""
|
||||
if FLAGS.is_train:
|
||||
if FLAGS.train:
|
||||
filenames = os.listdir(dataset)
|
||||
data_dir = os.path.join(os.getcwd(), dataset)
|
||||
else:
|
||||
@@ -77,9 +77,9 @@ def prepare_data(sess, dataset):
|
||||
def make_data(sess, checkpoint_dir, data, label):
|
||||
"""
|
||||
Make input data as h5 file format
|
||||
Depending on 'is_train' (flag value), savepath would be changed.
|
||||
Depending on 'train' (flag value), savepath would be changed.
|
||||
"""
|
||||
if FLAGS.is_train:
|
||||
if FLAGS.train:
|
||||
savepath = os.path.join(os.getcwd(), '{}/train.h5'.format(checkpoint_dir))
|
||||
else:
|
||||
savepath = os.path.join(os.getcwd(), '{}/test.h5'.format(checkpoint_dir))
|
||||
|
||||
Reference in New Issue
Block a user