如何使用image-super-resolution仓库加载与保存模型权重?
ISR模型权重的保存与加载方法
一、权重自动保存
你的代码中已配置log_dirs={'logs': './logs', 'weights': './weights'},Trainer会自动完成权重保存:
- 训练过程中每个epoch结束后,会保存最新权重到
./weights目录 - 你设置的
monitored_metrics={'val_PSNR_Y': 'max'}会触发自动保存验证集PSNR_Y指标最优的权重文件
若需自定义保存策略,可添加ModelCheckpoint回调:
from keras.callbacks import ModelCheckpoint # 自定义生成器权重保存回调 checkpoint_gen = ModelCheckpoint( './weights/custom_generator_best.h5', monitor='val_PSNR_Y', save_best_only=True, save_weights_only=True, mode='max', verbose=1 ) # 在训练时传入回调 trainer.train( epochs=80, steps_per_epoch=100, batch_size=16, monitored_metrics={'val_PSNR_Y': 'max'}, callbacks=[checkpoint_gen] )
二、加载预训练权重
1. 初始化Trainer时加载
直接在Trainer初始化参数中指定权重文件路径:
trainer = Trainer( generator=rrdn, discriminator=discr, feature_extractor=f_ext, lr_train_dir='lrtrain', hr_train_dir='hrtrain', lr_valid_dir='lrval', hr_valid_dir='hrval', loss_weights=loss_weights, learning_rate=learning_rate, flatness=flatness, dataname='image_dataset', log_dirs=log_dirs, weights_generator='./weights/generator_best_val_PSNR_Y.h5', # 替换为你的权重路径 weights_discriminator='./weights/discriminator_best_val_PSNR_Y.h5', # 判别器权重可选 n_validation=40, )
2. 手动加载权重
若已完成模型初始化,可直接调用模型的load_weights方法:
# 加载生成器权重 rrdn.load_weights('./weights/generator_best_val_PSNR_Y.h5') # 加载判别器权重 discr.load_weights('./weights/discriminator_best_val_PSNR_Y.h5')
3. 加载官方预训练权重
初始化模型时直接指定官方预训练权重标识即可:
rrdn = RRDN( arch_params={'C':4, 'D':3, 'G':64, 'G0':64, 'T':10, 'x':scale}, patch_size=lr_train_patch_size, weights='psnr-small' # 可选标识:'psnr-large'、'gan'等 )
内容的提问来源于stack exchange,提问作者Entretoize
相关产品推荐
相关产品推荐

