@@ -134,21 +134,28 @@ def save_history(
134134 save_log (start_time , finish_time , cls_report , cm , log_dir , classes )
135135
136136
137- def train (backbone_ver = "squeezenet1_1" , epoch_num = 40 , iteration = 10 , lr = 0.001 ):
137+ def train (
138+ backbone_ver = "squeezenet1_1" ,
139+ epoch_num = 40 ,
140+ iteration = 10 ,
141+ lr = 0.001 ,
142+ use_wce = True ,
143+ full_finetune = True ,
144+ ):
138145 # device = torch.device("cuda" if torch.cuda.is_available() else "cpu")
139146 tra_acc_list , val_acc_list , loss_list , lr_list = [], [], [], []
140147
141148 # load data
142- ds , classes , num_samples = prepare_data (args . wce )
149+ ds , classes , num_samples = prepare_data (use_wce )
143150 cls_num = len (classes )
144151
145152 # init model
146- model = Net (cls_num , m_ver = backbone_ver , full_finetune = args . fullfinetune )
153+ model = Net (cls_num , m_ver = backbone_ver , full_finetune = full_finetune )
147154 input_size = model ._get_insize ()
148155 traLoader , valLoader , tesLoader = load_data (ds , input_size )
149156
150157 # optimizer and loss
151- criterion = WCE (num_samples ) if args . wce else nn .CrossEntropyLoss ()
158+ criterion = WCE (num_samples ) if use_wce else nn .CrossEntropyLoss ()
152159 optimizer = optim .SGD (model .parameters (), lr , momentum = 0.9 )
153160 scheduler = torch .optim .lr_scheduler .ReduceLROnPlateau (
154161 optimizer ,
@@ -173,9 +180,9 @@ def train(backbone_ver="squeezenet1_1", epoch_num=40, iteration=10, lr=0.001):
173180
174181 # train process
175182 start_time = datetime .now ()
176- log_dir = f"{ LOGS_DIR } /{ args . model } __{ start_time .strftime ('%Y-%m-%d_%H-%M-%S' )} "
183+ log_dir = f"{ LOGS_DIR } /{ backbone_ver } __{ start_time .strftime ('%Y-%m-%d_%H-%M-%S' )} "
177184 create_dir (log_dir )
178- print (f"Start training { args . model } at { start_time } ..." )
185+ print (f"Start training { backbone_ver } at { start_time } ..." )
179186 # loop over the dataset multiple times
180187 for epoch in range (epoch_num ):
181188 epoch_str = f" Epoch { epoch + 1 } /{ epoch_num } "
@@ -244,7 +251,13 @@ def train(backbone_ver="squeezenet1_1", epoch_num=40, iteration=10, lr=0.001):
244251 warnings .filterwarnings ("ignore" )
245252 parser = argparse .ArgumentParser (description = "train" )
246253 parser .add_argument ("--model" , type = str , default = "squeezenet1_1" )
254+ parser .add_argument ("--epoch" , type = int , default = 40 )
247255 parser .add_argument ("--wce" , type = bool , default = True )
248- parser .add_argument ("--fullfinetune" , type = bool , default = False )
256+ parser .add_argument ("--fullfinetune" , type = bool , default = True )
249257 args = parser .parse_args ()
250- train (backbone_ver = args .model , epoch_num = 2 ) # 2 for test
258+ train (
259+ backbone_ver = args .model ,
260+ epoch_num = args .epoch ,
261+ use_wce = args .wce ,
262+ full_finetune = args .fullfinetune ,
263+ )
0 commit comments