turns-00052.parquet:48552
5740f26b9541aac98d6b7912
turn 7/7gpt-4o-2024-08-06EnglishChina1827 words
degenerate_repetitionAbsentFinal dense release
USER
import argparse
import os,glob,sys,yaml,cv2
import time,datetime
import numpy as np
import torch
import torch.nn as nn
import torch.nn.parallel
import torch.backends.cudnn as cudnn
import torch.optim as optim
import torch.nn.functional as F
import torch.utils.data
from thop import profile
INVALID_DISPARITY_F32 = 0.-300.
INVALID_DISPARITY_U16 = 32768
from util import save_checkpoint,AverageMeter,save_float_to_colormap,normerWrong,normer,normerImageNet,resizer,constantNorm
# from lightStereo.lightstereo_refine import LightStereo
from lightStereo.lightstereo import LightStereo
def dispFloatToUint16(dispFOrig, ratio=1.):
'''
disparity range from -300 to -0 is currently not suppoted for read/write to png file
'''
dispF = dispFOrig.copy()*ratio
dispInvMsk=(dispF<=INVALID_DISPARITY_F32*ratio)
dispValMsk=(dispF>INVALID_DISPARITY_F32*ratio)
dispF[dispValMsk] *= 128.
dispF[dispInvMsk] = INVALID_DISPARITY_U16
dispU16 = dispF.astype(np.uint16)
return dispU16
def convert_relu_to_relu6(model):
for child_name, child in model.named_children():
if isinstance(child, (nn.ReLU, nn.LeakyReLU)):
#print(child_name)
setattr(model, child_name, nn.ReLU6(inplace=True))
else:
convert_relu_to_relu6(child)
# gaussian blur
def gaussBlurCore(img, ksize=(3,3), sigma=0.35):
kw = int(ksize[0])
kh = int(ksize[1])
blurImg = cv2.GaussianBlur(img, (kw,kh), sigma)
return blurImg
def gaussBlur(left, right):
left, right = left.copy(), right.copy()
left, right = left, right
left = gaussBlurCore(left)
left,right = left, right
return left, right
def colorTransfer(Isrc, Iref):
Is = Isrc.copy()
Ir = Iref.copy()
#plt.imshow(Ir[...,::-1])
#plt.show()
#plt.imshow(Is[...,::-1])
#plt.show()
#BGR->LAB
valid = ((Iref[:,:,0]>0) & (Iref[:,:,1]>0) & (Iref[:,:,2]>0))
LabIs = cv2.cvtColor(Is,cv2.COLOR_BGR2LAB)
LabIr = cv2.cvtColor(Ir,cv2.COLOR_BGR2LAB)
#mean、std
Is_means = [0, 0, 0]
Ir_means = [0, 0, 0]
Is_stdevs = [0, 0, 0]
Ir_stdevs = [0, 0, 0]
LabIs = LabIs.astype(np.float32)/255.
LabIr = LabIr.astype(np.float32)/255.
for i in range(3):
Is_means[i] += LabIs[:,:,i].mean()
Ir_means[i] += LabIr[:,:,i][valid].mean()
Is_stdevs[i] += (LabIs[:,:,i].std() + 1e-9)
Ir_stdevs[i] += (LabIr[:,:,i][valid].std() + 1e-9)
thresh = [a/b for a,b in zip(Ir_stdevs,Is_stdevs)]
LabIt = thresh*(LabIs -Is_means) + Ir_means
# [0-255]
LabIt = (LabIt*255.)
LabIt *= (LabIt>0)
LabIt = (LabIt * (LabIt<=255) + 255 * (LabIt>255)).astype(np.uint8)
#show
It = cv2.cvtColor(LabIt,cv2.COLOR_LAB2BGR)
#plt.imshow(It)
#plt.show()
return It
def getYamlCfg(fn='models/configModelTest.yaml'):
f = open(fn, 'r', encoding='utf-8')
cfg = f.read()
content = yaml.load(cfg, Loader=yaml.FullLoader)
return content
def initModels(content, modelFolder='models', args=None):
numOfModels = len(list(content.keys()))
#importedFolder= __import__(modelFolder)
singleModel = {}
for idx in range(numOfModels):
modelName = list(content.keys())[idx]
if content[modelName]['used4Train'] is not True:
continue
print('=> Importing model-{}'.format(modelName))
importedFolder = __import__(content[modelName]['modelFolder'])
moduleImported = getattr(importedFolder, content[modelName]['moduleName'])
classImported = getattr(moduleImported, content[modelName]['className'])
singleModel['classInst'] = classImported(batchNorm=content[modelName]['batchNorm'], pred_disp=content[modelName]['pred_disp'], args=args)
break
return singleModel['classInst']
def initModelsBug(content, modelFolder='models'):
numOfModels = len(list(content.keys()))
importedFolder= __import__(modelFolder)
singleModel = {}
for idx in range(numOfModels):
modelName = list(content.keys())[idx]
if content[modelName]['used4Train'] is not True:
continue
print('=> Importing model-{}'.format(modelName))
moduleImported = getattr(importedFolder, content[modelName]['moduleName'])
classImported = getattr(moduleImported, content[modelName]['className'])
singleModel['classInst'] = classImported(batchNorm=content[modelName]['batchNorm'], pred_disp=content[modelName]['pred_disp'])
break
return singleModel['classInst']
def loadPretrain(model, pretrained, strict=False):
oldCkpt = torch.load(pretrained)
if pretrained.find('.tar')>=0:
oldCkpt = oldCkpt['state_dict']
newCkpt = {}
for k in oldCkpt.keys():
newCkpt[k.replace('module.','')] = oldCkpt[k]
model.load_state_dict(newCkpt, strict=strict)
return model
# Training settings
parser = argparse.ArgumentParser(description='Stereo training...')
parser.add_argument('--data_parent_path', type=str, default='/home/notebook/data/group/xizhong_xxz/', help="location to save models")
parser.add_argument('--data_config_path', type=str, default='./configs/configDatasetMultiscale.yaml', help="location to save models")
parser.add_argument('--model_config_path', type=str, default='./configs/configModel.yaml', help="location to save models")
parser.add_argument('--loss_config_path', type=str, default='./configs/configLoss.yaml', help="location to save models")
parser.add_argument('--save_path', type=str, default='./rectify', help="location to save models")
parser.add_argument('--use_gru_refine', type=int, default=0, help="if use gru refine for disp output")
parser.add_argument('--useBackboneV2', type=int, default=0, help='1 use mobilenet')
parser.add_argument('--att_relu6', type=int, default=0, help='1 use relu6 for att module')
parser.add_argument('--convex_up', type=int, default=0, help='1 use convex upsample')
parser.add_argument('--transAtt', type=int, default=0, help='1 use self att')
parser.add_argument('--agg_blocks', default=[8,16,32], metavar='N', nargs='*', help='set the block number for each stage')
parser.add_argument('--pretrained_Backbone', type=int, default=1, help='1 use pretrained backbone to extractor feature')
parser.add_argument('--sample_tail', type=int, default=0, help="if add a tail refinement for output quantize")
parser.add_argument('--use_mono', type=int, default=0, help="if use mono optimize stereo")
parser.add_argument('--useResNet', type=int, default=0, help='use resnet as backbone')
parser.add_argument('--validateOnly', type=int, default=0, help='do validation only. Default=0')
parser.add_argument('--warmup', type=int, default=0, help='do warmup training or not, Default=0')
parser.add_argument('--abnormAug', type=int, default=0, help='do abnormal augmenation such as thin plate transform or diagonal epipolar transform')
parser.add_argument('--eraseAug', type=int, default=0, help='do erasing augmenation')
parser.add_argument('--occAug', type=int, default=0, help='do occlusion augmentation in right image branch')
parser.add_argument('--defocusAug', type=int, default=0, help='do blur augmentation in all image branches')
parser.add_argument('--gradLoss', type=int, default=0, help='use gradient loss for supervision')
parser.add_argument('--normLoss', type=int, default=0, help='use normal loss for supervision')
parser.add_argument('--smtLoss', type=int, default=0, help='use smoothness loss for self-supervision')
parser.add_argument('--use_dn', type=int, default=0, help='0 for bn, 1 for dn')
parser.add_argument('--mobile', type=int, default=0, help='0 for bn, 1 for mobile')
parser.add_argument('--use_relu6', type=int, default=0, help='0 for relu, 1 for relu6')
parser.add_argument('--use_refine_relu6', type=int, default=0, help='0 for relu, 1 for relu6')
parser.add_argument('--backbone_relu6_only', type=int, default=0, help='0 for relu, 1 for relu6')
parser.add_argument('--noNormInput', type=int, default=0, help='normalize input to 0.0 ~ 1.0')
parser.add_argument('--rgbOrder', type=int, default=0, help='0 for bgr, 1 for rgb')
parser.add_argument('--normGT', type=int, default=0, help='normalize GT to 0.0 ~ 1.0')
parser.add_argument('--ada_sigma', type=int, default=0, help='0 for fixed sigma, 1 for adaptive sigma')
parser.add_argument('--depth_wise', type=int, default=0, help='depth wise conv for decoder')
parser.add_argument('--ada_cost', type=int, default=1, help='0 for fixed cost, 1 for adaptive cost')
parser.add_argument('--simple_refine', type=int, default=0, help='1 for simple refiner, 0 for warp refiner')
parser.add_argument('--rightDown2x', type=int, default=0, help='1 for right image plane down 2x, 0 for no downscale')
parser.add_argument('--clipGrad', type=int, default=0, help='use gradient clipping for training')
parser.add_argument('--clip', type=float, default=2.1, help='clipping range')
parser.add_argument('--skyFill', type=int, default=1, help='fill sky area with disparity 0.0 with semantic label')
parser.add_argument('--use_refine', type=int, default=0, help='use refinement blocks')
parser.add_argument('--refine_stages', type=int, default=2, help='use refinement blocks')
parser.add_argument('--warp_feature', type=int, default=0, help='warp feature or original image')
parser.add_argument('--wRec', type=float, default=0.0, help='gradient loss weight')
parser.add_argument('--proj_ch', type=int, default=3, help='channels to reconsrtuct')
parser.add_argument('--pre_scale', type=float, default=0.666667, help='pre-scale')
parser.add_argument('--cat2Fuse', type=int, default=0, help='costs cat to fuse')
parser.add_argument('--trainSceneFlowOnly', type=int, default=0, help='train scene flow dataset only')
parser.add_argument('--wGrad', type=float, default=0.5, help='gradient loss weight')
parser.add_argument('--wSmt', type=float, default=0.5, help='smooth loss weight')
parser.add_argument('--inconsistAug', type=int, default=0, help='inconsistancy augmentation for LR pairs')
parser.add_argument('--imageNetInputNorm', type=int, default=0, help='inputs normalized like in imagenet')
parser.add_argument('--constantNorm', type=int, default=0, help='inputs normalized to [-1, 1]')
parser.add_argument('--useTransformer', type=int, default=0, help='use transformer blocks')
parser.add_argument('--correctCorr', type=int, default=0, help='use correct correlation')
parser.add_argument('--copyPaste', type=int, default=1, help='use copy paste tech.')
parser.add_argument('--zeroBlack', type=int, default=0, help='black area filled with zero when normalizing.')
parser.add_argument('--zeroInvalidGt', type=int, default=0, help='black area filled with zero gt.')
parser.add_argument('--maxdisp', type=int, default=120, help="max disparity range")
#parser.add_argument('--maxdisp', type=int, default=120, help="max disparity range")
parser.add_argument('--shift', type=int, default=0, help='random shift of left image. Default=0')
#shift value must be zero for cost volume aggregation (soft argmin) based network
parser.add_argument('--crop_height', type=int, default=640, help="crop height")
parser.add_argument('--crop_width', type=int, default=768, help="crop width")
parser.add_argument('--resume', type=str, default='', help="resume from saved model")
parser.add_argument('--start_epoch', type=int, default=0, help='the epoch to resume training')
parser.add_argument('--pretrained', dest='pretrained', default=None, help='path to pre-trained model')
parser.add_argument('--batch_size', type=int, default=8, help='training batch size')
parser.add_argument('--val_batch_size', type=int, default=8, help='testing batch size')
parser.add_argument('--workers', type=int, default=32, help='number of threads')
parser.add_argument('--nEpochs', type=int, default=64, help='number of epochs to train for')
parser.add_argument('--solver', default='adam',choices=['adam','sgd'], help='solver algorithms')
parser.add_argument('--lr', type=float, default=0.0008, help='Learning Rate. Default=0.001')
parser.add_argument('--lr_decay', type=float, default=0.5, help='Learning Rate descending rete. Default=0.5')
'''
Attention: milestones should be ignored(keep unchanged), since it's adapted to nEpochs in the latter main function
'''
#####parser.add_argument('--milestones', default=[14,28,42,56], metavar='N', nargs='*', help='epochs at which learning rate is divided by 2')
parser.add_argument('--milestones', default=[10,20,30,40,50,60], metavar='N', nargs='*', help='epochs at which learning rate is divided by 2')
#parser.add_argument('--milestones', default=[13,23,33,43,53,63], metavar='N', nargs='*', help='epochs at which learning rate is divided by 2')
#parser.add_argument('--weight_decay', '--wd', default=4e-4, type=float, metavar='W', help='weight decay')
parser.add_argument('--weight_decay', '--wd', default=1e-4, type=float, metavar='W', help='weight decay')
parser.add_argument('--bias_decay', default=0, type=float, metavar='B', help='bias decay')
parser.add_argument('--momentum', default=0.9, type=float, metavar='M', help='momentum for sgd, alpha parameter for adam')
parser.add_argument('--beta', default=0.999, type=float, metavar='M', help='beta parameter for adam')
'''
ascending order means use loss.sum()/batchSize, which is map level average
descending order means use loss.mean(), which is pixel level average
'''
#parser.add_argument('--multiscale_weights', '-w', default=[0.005,0.01,0.02,0.08,0.32], type=float, nargs=5, help='training weight for each scale, from highest resolution (flow2) to lowest (flow6)', metavar=('W2', 'W3', 'W4', 'W5', 'W6'))
#parser.add_argument('--multiscale_weights', '-w', default=[0.32,0.08,0.02,0.01,0.005], type=float, nargs=5, help='training weight for each scale, from highest resolution (flow2) to lowest (flow6)', metavar=('W2', 'W3', 'W4', 'W5', 'W6'))
parser.add_argument('--multiscale_weights', '-w', default=[1./3, 2./3, 1., 1., 1.], type=float, nargs=5, help='training weight for each scale, from highest resolution (flow2) to lowest (flow6)', metavar=('W2', 'W3', 'W4', 'W5', 'W6'))
parser.add_argument('--seed', type=int, default=2022, help='random seed to use. Default=123')
parser.add_argument('--print_freq', '-p', default=50, type=int, metavar='N', help='print frequency')
parser.add_argument('--save_freq', '-s', default=5, type=int, metavar='N', help='save checkpoint frequency')
device = torch.device("cuda" if torch.cuda.is_available() else "cpu")
def main():
args = parser.parse_args()
device = torch.device("cuda" if torch.cuda.is_available() else "cpu")
print('=> get test data pairs')
# args.data_path="/home/notebook/data/group/liulibo/stereoProject/rectify/bars/"
# args.data_path="/home/notebook/data/group/liulibo/MiDaS/langan_refine_data/"
args.data_path="/home/notebook/data/group/Ruby/stereo-from-mono/dirty_results/"
args.flip_LR=0
args.useMinMax=0
lists =[args.data_path+file for file in os.listdir(args.data_path) if(("left" in file or "master" in file) and "disp" not in file)]
total = len(lists)
print(total)
#pause
print('=> get models from configuration file')
content = getYamlCfg(fn='configs/configModelTest.yaml')
model = LightStereo(parser)
model = model.to(device).eval()
# model = loadPretrain(model, "./checkpoint/refiner_finetune_20240918/model_best.pth.tar", strict=False)
model = loadPretrain(model, "./checkpoint/gru/best.pth.tar", strict=False)
print('=> infer')
count=0
# for idx in range(total):
for file in os.listdir(args.data_path):
if(("left" in file or "master" in file)):
if "Disp" in file:
continue
try:
srcLFn = args.data_path+file
srcRFn = srcLFn.replace('left','right').replace('master', 'slave')
print(srcLFn)
print(srcRFn)
lImg = cv2.imread(srcLFn,1)
rImg = cv2.imread(srcRFn,1)
h,w,c=lImg.shape
hh,ww,cc=rImg.shape
lMask = ((lImg[:,:,0]==0) & (lImg[:,:,1]==0) & (lImg[:,:,2]==0)).astype(np.uint8)
except:
continue
if args.flip_LR>0:
lImg = lImg[:,::-1,:]
rImg = rImg[:,::-1,:]
lMask = lMask[:,::-1]
rszH,rszW = lImg.shape[0],lImg.shape[1]
lImg,rImg = resizer(lImg,rImg,args)
# if args.rgbOrder > 0:
# lImg= lImg[:,:,::-1].copy()
# rImg= rImg[:,:,::-1].copy()
# if args.imageNetInputNorm <= 0:
# imgPair = normer(lImg, rImg)
# else:
# if args.constantNorm:
# imgPair = constantNorm(lImg, rImg)
# else:
imgPair = normerImageNet(lImg, rImg)
imgPairTs = torch.from_numpy(imgPair).to(device).float().unsqueeze(0)
with torch.no_grad():
output= model(imgPairTs)
if isinstance(output, list):
output = output[-1].squeeze().cpu().numpy()
else:
output = output.squeeze().cpu().numpy()
if args.flip_LR>0:
output = output[:,::-1]
lMask = lMask[:,::-1]
outH,outW = output.shape[0],output.shape[1]
scaleDisp = rszW * 1. / outW
output = cv2.resize(output, (rszW,rszH), interpolation=cv2.INTER_LINEAR) * scaleDisp
#print(output.min())
# dispU16 = dispFloatToUint16(output, ratio=3)
output_path = '/home/notebook/data/group/Ruby/stereoProject/test_results'
cv2.imwrite(output_path +'/'+os.path.basename(srcLFn).replace('left','pred').replace('master','pred').replace('.jpg','.png'), output*10)
# cv2.imwrite(args.data_path+'/'+os.path.basename(srcLFn).replace('left','pred').replace('master','pred').replace('.jpg','.png'), output*10)
# print(args.save_path)
# cv2.imwrite(args.data_path+'/'+os.path.basename(srcLFn).replace('left','gru_leftDisp_gray').replace('master','gru_leftDisp_gray').replace('.jpg','.png'), dispU16)
# color = save_float_to_colormap(output, saveFn=args.data_path+'/'+os.path.basename(srcLFn).replace('left','gru_leftDisp').replace('master','gru_leftDisp').replace('.jpg','.png'), inputLimMin=0, inputLimMax=120, useMinMaxFirst=(args.useMinMax>0), saveU8=False)
# cv2.imwrite(args.data_path+'/'+os.path.basename(srcLFn).replace('left','leftDisp_gray').replace('master','leftDisp_gray').replace('.jpg','.png'), dispU16)
# color = save_float_to_colormap(output, saveFn=args.data_path+'/'+os.path.basename(srcLFn).replace('left','leftDisp').replace('master','leftDisp').replace('.jpg','.png'), inputLimMin=0, inputLimMax=120, useMinMaxFirst=(args.useMinMax>0), saveU8=False)
if __name__ == '__main__':
main()
把这个代码改成只处理文件名5000到5537的文件ASSISTANT
为了让脚本只处理文件名在5000到5537之间的文件,我们需要在处理文件时加入过滤条件。假设这些文件名中间部分是数字(并取决于这个假设来过滤)。这里是修改后的代码:
```python
import argparse
import os
import cv2
import numpy as np
import torch
import torch.nn as nn
import torch.nn.functional as F
from util import normerImageNet
from lightStereo.lightstereo import LightStereo
INVALID_DISPARITY_F32 = 0.-300.
INVALID_DISPARITY_U16 = 32768
def dispFloatToUint16(dispFOrig, ratio=1.):
dispF = dispFOrig.copy() * ratio
dispInvMsk = (dispF <= INVALID_DISPARITY_F32 * ratio)
dispValMsk = (dispF > INVALID_DISPARITY_F32 * ratio)
dispF[dispValMsk] *= 128.
dispF[dispInvMsk] = INVALID_DISPARITY_U16
dispU16 = dispF.astype(np.uint16)
return dispU16
def loadPretrain(model, pretrained, strict=False):
oldCkpt = torch.load(pretrained)
if pretrained.find('.tar') >= 0:
oldCkpt = oldCkpt['state_dict']
newCkpt = {k.replace('module.', ''): v for k, v in oldCkpt.items()}
model.load_state_dict(newCkpt, strict=strict)
return model
def main():
parser = argparse.ArgumentParser(description='Stereo training...')
# Add relevant arguments
parser.add_argument('--data_parent_path', type=str, default='/home/notebook/data/group/xizhong_xxz/', help="location to save models")
args = parser.parse_args()
device = torch.device("cuda" if torch.cuda.is_available() else "cpu")
args.data_path = "/home/notebook/data/group/Ruby/stereo-from-mono/dirty_results/"
args.flip_LR = 0
args.useMinMax = 0
print('=> get models from configuration file')
model = LightStereo(parser)
model = model.to(device).eval()
model = loadPretrain(model, "./checkpoint/gru/best.pth.tar", strict=False)
print('=> infer')
for file in os.listdir(args.data_path):
# Assume filenames are like '<number>_left.jpg' or '<number>_master.jpg'
try:
# Extract the numeric part
file_number = int(file.split('_')[0])
if 5000 <= file_number <= 5537 and ("left" in file or "master" in file):
if "Disp" in file:
continue
srcLFn = os.path.join(args.data_path, file)
srcRFn = srcLFn.replace('left', 'right').replace('master', 'slave')
print(srcLFn)
print(srcRFn)
lImg = cv2.imread(srcLFn, 1)
rImg = cv2.imread(srcRFn, 1)
lMask = ((lImg[:,:,0] == 0) & (lImg[:,:,1] == 0) & (lImg[:,:,2] == 0)).astype(np.uint8)
if args.flip_LR > 0:
lImg = lImg[:, ::-1, :]
rImg = rImg[:, ::-1, :]
lMask = lMask[:, ::-1]
lImg, rImg = resizer(lImg, rImg, args)
imgPair = normerImageNet(lImg, rImg)
imgPairTs = torch.from_numpy(imgPair).to(device).float().unsqueeze(0)
with torch.no_grad():
output = model(imgPairTs)
if isinstance(output, list):
output = output[-1].squeeze().cpu().numpy()
else:
output = output.squeeze().cpu().numpy()
if args.flip_LR > 0:
output = output[:, ::-1]
rszH, rszW = lImg.shape[:2]
outH, outW = output.shape
scaleDisp = rszW * 1. / outW
output = cv2.resize(output, (rszW, rszH), interpolation=cv2.INTER_LINEAR) * scaleDisp
output_path = '/home/notebook/data/group/Ruby/stereoProject/test_results'
cv2.imwrite(os.path.join(output_path, os.path.basename(srcLFn).replace('left','pred').replace('master','pred').replace('.jpg','.png')), output * 10)
except Exception as e:
print(f"Error processing file {file}: {e}")
continue
if __name__ == '__main__':
main()
```
### 关键调整:
- **文件过滤**: 仅处理 5000 到 5537 范围内包含数字的文件。例如,数字提取代码行:`file_number = int(file.split('_')[0])` 假设文件名是以数字开头,紧跟下划线和描述性文字(如 left 或 master)。
- **异常处理**: 引入 `try-except` 块,捕获并报告在处理文件时可能出现的错误。
- **主要处理逻辑**: 代码包含的关键图像处理和模型推断功能保持不变。
请确保文件名格式允许 `file.split('_')[0]` 能提取到整数字符串,否则可能需要修改字符串解析逻辑来适应具体文件名模式。