修改
1. 数据集:basicsr/data/paired_image_dataset.py
新增/改造 Dataset_PolaRGBImage,用于读取 HuggingFace PolaRGB 官方结构。
原来旧逻辑是读:
lq, gt, t, s
现在读:
lq, gt, p0, p45, p90, p135
核心初始化:
class Dataset_PolaRGBImage(data.Dataset):
"""PolaRGB dataset with RGB/LQ, four polarization images, and GT."""
def __init__(self, opt):
super(Dataset_PolaRGBImage, self).__init__()
self.opt = opt
self.file_client = None
self.io_backend_opt = opt['io_backend']
self.mean = opt['mean'] if 'mean' in opt else None
self.std = opt['std'] if 'std' in opt else None
self.gt_folder = opt['dataroot_gt']
self.lq_folder = opt['dataroot_lq']
self.p0_folder = opt.get('dataroot_p0')
self.p45_folder = opt.get('dataroot_p45')
self.p90_folder = opt.get('dataroot_p90')
self.p135_folder = opt.get('dataroot_p135')
if opt.get('pola_rgb_layout') == 'official':
self.paths = self._official_paths_from_root()
else:
self.paths = self._paths_from_folders()
这里 pola_rgb_layout: official 是新增配置开关。因为 PolaRGB 不是 rgb/0/45/90/135 平行目录,而是:
train/easy/input/00/0000_rgb.png
train/easy/input/00/0000_000.png
train/easy/input/00/0000_045.png
train/easy/input/00/0000_090.png
train/easy/input/00/0000_135.png
train/easy/gt/00/0000_rgb.png
官方结构扫描代码:
def _official_paths_from_root(self):
split_root = self.lq_folder
if self.opt['phase'] == 'train':
subsets = self.opt.get('subsets', ['easy', 'hard'])
input_roots = [osp.join(split_root, subset, 'input') for subset in subsets]
gt_roots = [osp.join(split_root, subset, 'gt') for subset in subsets]
else:
input_roots = [osp.join(split_root, 'input')]
gt_roots = [osp.join(split_root, 'gt')]
训练集读取 easy + hard;验证集读取 test/input 和 test/gt。
匹配 RGB、四张偏振、GT:
polar_suffixes = {
'p0': '_000',
'p45': '_045',
'p90': '_090',
'p135': '_135'
}
for input_root, gt_root in zip(input_roots, gt_roots):
rgb_paths = sorted([
path for path in scandir(input_root, recursive=True)
if path.endswith('_rgb.png')
])
for rgb_rel in rgb_paths:
scene = osp.dirname(rgb_rel)
basename, ext = osp.splitext(osp.basename(rgb_rel))
sample_id = basename[:-4]
例如:
basename = 0000_rgb
sample_id = 0000
GT 匹配:
gt_candidates = [
osp.join(gt_root, scene, f'{sample_id}_rgb{ext}'),
osp.join(gt_root, scene, f'0000_rgb{ext}'),
osp.join(gt_root, scene, f'000_rgb{ext}')
]
gt_path = next((path for path in gt_candidates if osp.isfile(path)), gt_candidates[0])
这里是为了兼容你数据里的 0000_rgb.png,也兼容可能存在的 000_rgb.png。
构造一条样本:
item = {
'lq_path': osp.join(input_root, rgb_rel),
'gt_path': gt_path
}
for key, suffix in polar_suffixes.items():
item[f'{key}_path'] = osp.join(
input_root, scene, f'{sample_id}{suffix}{ext}')
也就是自动找:
0000_rgb.png
0000_000.png
0000_045.png
0000_090.png
0000_135.png
为了避免某些样本缺文件导致 DataLoader 中途崩,构建阶段检查完整性:
if all(osp.isfile(path) for path in item.values()):
paths.append(item)
else:
skipped += 1
读取图像:
img_lq = self._read_img(paths['lq_path'], 'lq')
img_gt = self._read_img(paths['gt_path'], 'gt')
img_p0 = self._read_img(paths['p0_path'], 'p0')
img_p45 = self._read_img(paths['p45_path'], 'p45')
img_p90 = self._read_img(paths['p90_path'], 'p90')
img_p135 = self._read_img(paths['p135_path'], 'p135')
训练增强同步作用在 6 张图上:
imgs = [img_lq, img_gt, img_p0, img_p45, img_p90, img_p135]
if self.opt['phase'] == 'train':
gt_size = self.opt['gt_size']
imgs = self._padding(imgs, gt_size)
imgs = self._random_crop(imgs, gt_size, scale, paths['gt_path'])
if self.geometric_augs:
imgs = random_augmentation(*imgs)
返回字段:
return {
'lq': img_lq,
'gt': img_gt,
'p0': img_p0,
'p45': img_p45,
'p90': img_p90,
'p135': img_p135,
'lq_path': paths['lq_path'],
'gt_path': paths['gt_path'],
'p0_path': paths['p0_path'],
'p45_path': paths['p45_path'],
'p90_path': paths['p90_path'],
'p135_path': paths['p135_path']
}
2. 网络:basicsr/models/archs/PolarFMamba_arch.py
原来网络输入是:
forward(x, t, s)
内部调用:
generate_polarized_images(x, t, s)
现在删掉偏振生成,改成直接输入真实偏振:
forward(x, p0, p45, p90, p135)
偏振分解模块:
class PolarFrequencyDecomposition(nn.Module):
def __init__(self, d_model, polar_in_channels=12):
super().__init__()
self.polar_compression = nn.Sequential(
nn.Conv2d(polar_in_channels, 3, 1),
nn.GELU()
)
为什么是 12:
p0 3 channels
p45 3 channels
p90 3 channels
p135 3 channels
总共 12 channels
forward:
def forward(self, *polar_imgs):
polar_stack = torch.cat(polar_imgs, dim=1)
polar_compressed = self.polar_compression(polar_stack)
主网络:
class PolarFMamba(nn.Module):
def __init__(self, d_model=32, depth=4, polar_in_channels=12):
super().__init__()
self.input_proj = nn.Conv2d(3, d_model, 3, padding=1)
self.polar_decomp = PolarFrequencyDecomposition(d_model, polar_in_channels)
新的 forward:
def forward(self, x, p0, p45, p90, p135):
x_feat = self.input_proj(x)
polar_high, polar_low = self.polar_decomp(p0, p45, p90, p135)
for block in self.blocks:
fused_feat, _, _ = block['dual_freq'](x_feat, x_feat)
x_feat = block['freq_filter'](fused_feat, polar_high, polar_low)
x_feat = x_feat + fused_feat
return self.reconstruction(x_feat)
核心变化就是:
polar_high, polar_low = self.polar_decomp(p0, p45, p90, p135)
不再有:
generate_polarized_images(...)
3. 模型训练封装:basicsr/models/image_restoration_model.py
新增 PolaRGB 数据喂入:
def feed_train_polargbdata(self, data):
self.lq = data['lq'].to(self.device)
self.p0 = data['p0'].to(self.device)
self.p45 = data['p45'].to(self.device)
self.p90 = data['p90'].to(self.device)
self.p135 = data['p135'].to(self.device)
if 'gt' in data:
self.gt = data['gt'].to(self.device)
PolaRGB 的 mixup 要同步混所有输入,否则 RGB 和偏振会错位:
def mixup_multi(self, target, *inputs):
lam = self.dist.rsample((1, 1)).item()
r_index = torch.randperm(target.size(0)).to(self.device)
target = lam * target + (1 - lam) * target[r_index, :]
inputs = [
lam * input_ + (1 - lam) * input_[r_index, :]
for input_ in inputs
]
return target, *inputs
在 PolaRGB feed 中使用:
if self.mixing_flag:
use_identity = self.mixing_augmentation.use_identity
apply_mixup = random.randint(0, 1) == 0 if use_identity else True
if apply_mixup:
self.gt, self.lq, self.p0, self.p45, self.p90, self.p135 = (
self.mixing_augmentation.mixup_multi(
self.gt, self.lq, self.p0, self.p45, self.p90, self.p135))
新增优化函数:
def optimize_parameters_polargb(self, current_iter):
self.optimizer_g.zero_grad()
with autocast(enabled=self.use_amp):
preds = self.net_g(self.lq, self.p0, self.p45, self.p90, self.p135)
if not isinstance(preds, list):
preds = [preds]
self.output = preds[-1]
loss_dict = OrderedDict()
l_pix = 0.
for pred in preds:
l_pix += self.cri_pix(pred, self.gt)
loss_dict['l_pix'] = l_pix
self.amp_scaler.scale(l_pix).backward()
self.amp_scaler.unscale_(self.optimizer_g)
if self.opt['train']['use_grad_clip']:
torch.nn.utils.clip_grad_norm_(self.net_g.parameters(), 0.01)
self.amp_scaler.step(self.optimizer_g)
self.amp_scaler.update()
self.log_dict = self.reduce_loss_dict(loss_dict)
验证推理也支持五输入:
def nonpad_test(self, img=None, t=None, s=None, p0=None, p45=None, p90=None, p135=None):
img = self.lq if img is None else img
p0 = getattr(self, 'p0', None) if p0 is None else p0
p45 = getattr(self, 'p45', None) if p45 is None else p45
p90 = getattr(self, 'p90', None) if p90 is None else p90
p135 = getattr(self, 'p135', None) if p135 is None else p135
model = getattr(self, 'net_g_ema', self.net_g)
model.eval()
with torch.no_grad():
if all(v is not None for v in [p0, p45, p90, p135]):
pred = model(img, p0, p45, p90, p135)
elif t is not None and s is not None:
pred = model(img, t, s)
else:
pred = model(img)
self.output = pred[-1] if isinstance(pred, list) else pred
4. 训练入口:basicsr/PolarFMamba_train.py
训练循环里原来取:
lq = train_data['lq']
gt = train_data['gt']
t = train_data['t']
s = train_data['s']
现在改成:
lq = train_data['lq']
gt = train_data['gt']
p0 = train_data['p0']
p45 = train_data['p45']
p90 = train_data['p90']
p135 = train_data['p135']
mini batch 抽样同步处理:
if mini_batch_size < batch_size:
indices = random.sample(range(0, batch_size), k=mini_batch_size)
lq = lq[indices]
gt = gt[indices]
p0 = p0[indices]
p45 = p45[indices]
p90 = p90[indices]
p135 = p135[indices]
progressive crop 同步裁剪:
if mini_gt_size < gt_size:
x0 = int((gt_size - mini_gt_size) * random.random())
y0 = int((gt_size - mini_gt_size) * random.random())
x1 = x0 + mini_gt_size
y1 = y0 + mini_gt_size
lq = lq[:, :, x0:x1, y0:y1]
gt = gt[:, :, x0 * scale:x1 * scale, y0 * scale:y1 * scale]
p0 = p0[:, :, x0:x1, y0:y1]
p45 = p45[:, :, x0:x1, y0:y1]
p90 = p90[:, :, x0:x1, y0:y1]
p135 = p135[:, :, x0:x1, y0:y1]
调用新的训练接口:
model.feed_train_polargbdata({
'lq': lq,
'gt': gt,
'p0': p0,
'p45': p45,
'p90': p90,
'p135': p135
})
model.optimize_parameters_polargb(current_iter)
另外删掉了顶部类似这种硬编码:
os.environ["CUDA_VISIBLE_DEVICES"] = "6,7"
否则四卡 DDP 会被脚本内部覆盖,导致 Duplicate GPU detected。
5. 配置文件:Options/PolarFMamba_cityscapes_v5.yml
数据集类型改成:
type: Dataset_PolaRGBImage
pola_rgb_layout: official
训练集:
datasets:
train:
name: TrainSet
type: Dataset_PolaRGBImage
pola_rgb_layout: official
dataroot_gt: /home/student_account/student27/PolaRGB/train
dataroot_lq: /home/student_account/student27/PolaRGB/train
subsets: [easy, hard]
验证集:
val:
name: ValSet
type: Dataset_PolaRGBImage
pola_rgb_layout: official
dataroot_gt: /home/student_account/student27/PolaRGB/test
dataroot_lq: /home/student_account/student27/PolaRGB/test
网络:
network_g:
type: PolarFMamba
d_model: 32
polar_in_channels: 12
四卡:
num_gpu: 4
后来为显存稳定,额外建了 bs24 配置:
batch_size_per_gpu: 24
mini_batch_sizes: [24]
现在训练日志里的:
Batch size per gpu: 24
World size: 4
Batch_Size to 96
说明最终实际训练 batch 是:
24 * 4 = 96
6. 启动命令
必须带 PYTHONPATH,否则会 import 到旧的:
/home/student/student22/RetinexMamba
正确命令:
cd /home/student_account/student27/PolarMamba
env PYTHONPATH=/home/student_account/student27/PolarMamba \
PYTORCH_CUDA_ALLOC_CONF=max_split_size_mb:128 \
CUDA_VISIBLE_DEVICES=1,2,3,4 torchrun \
--nproc_per_node=4 \
--master_port=4335 \
basicsr/PolarFMamba_train.py \
--opt Options/PolarFMamba_cityscapes_v5_bs24.yml \
--launcher pytorch
总流程变化
修改前:
RGB + t + s
-> generate_polarized_images()
-> 生成模拟偏振图
-> PolarFMamba
修改后:
PolaRGB 官方数据集
-> RGB + 000 + 045 + 090 + 135
-> 直接送入 PolarFMamba
-> 不再生成偏振图
核心收益是:训练使用真实偏振观测,而不是通过 t/s 人工合成偏振。

浙公网安备 33010602011771号