修改

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 人工合成偏振。

posted @ 2026-05-06 16:02  伟大的船长  阅读(17)  评论(0)    收藏  举报