当前位置:首页 > 智能硬件 > 人工智能AI
[导读] 训练专项网络 还记得我们在开始时丢弃的70%的培训数据吗?结果表明,如果我们想在Kaggle排行榜上获得一个有竞争力的得分,这是一个很糟糕的主意。在70%的数据和挑战的测试集中,我们的模

训练专项网络

还记得我们在开始时丢弃的70%的培训数据吗?结果表明,如果我们想在Kaggle排行榜上获得一个有竞争力的得分,这是一个很糟糕的主意。在70%的数据和挑战的测试集中,我们的模型还有相当多特征没有看到。

因此,改变之前只训练单个模型的方式,让我们训练几个专项网络,每个专项网络预测一组不同的目标值。我们将训练一个只预测left_eye_center和right_eye_center的模型,一个仅用于nose_TIp等等;总的来说,我们将有六个模型。这将允许我们使用完整的训练数据集,并希望获得整体更有竞争力的分数。

六个专项网络都将使用完全相同的网络架构(一种简单的方法,不一定是最好的)。因为训练必须比以前花费更长的时间,所以让我们考虑一个策略,以便我们不必等待max_epochs完成,即使验证错误停止提高很多。这被称为早期停止,我们将写另一个on_epoch_finished回调来处理。这里的实现:
class EarlyStopping(object):
def __init__(self, paTIence=100):
self.paTIence = paTIence
self.best_valid = np.inf
self.best_valid_epoch = 0
self.best_weights = None

def __call__(self, nn, train_history):
current_valid = train_history[-1]['valid_loss']
current_epoch = train_history[-1]['epoch']
if current_valid < self.best_valid:
self.best_valid = current_valid
self.best_valid_epoch = current_epoch
self.best_weights = nn.get_all_params_values()
elif self.best_valid_epoch + self.patience < current_epoch:
print("Early stopping.")
print("Best valid loss was {:.6f} at epoch {}.".format(
self.best_valid, self.best_valid_epoch))
nn.load_params_from(self.best_weights)
raise StopIteration()

可以看到,在call函数里面有两个分支:第一个是现在的验证错误比我们之前看到的要好,第二个是最好的验证错误所在的迭代次数和当前迭代次数的距离已经超过了我们的耐心。在第一个分支里,我们存下网络的权重:
self.best_weights = nn.get_all_params_values()

第二个分支里,我们将网络的权重设置成最优的验证错误时存下的值,然后发出一个StopIteration,告诉NeuralNet我们想要停止训练。
nn.load_params_from(self.best_weights)
raise StopIteration()

让我们在net的定义中更新on_epoch_finished处理程序的列表,并添加EarlyStopping:
net8 = NeuralNet(
# ...
on_epoch_finished=[
AdjustVariable('update_learning_rate', start=0.03, stop=0.0001),
AdjustVariable('update_momentum', start=0.9, stop=0.999),
EarlyStopping(patience=200),
],
# ...
)

到目前为止一切顺利,但是如何定义这些专项网络进行相应的预测呢?让我们做一个列表:
SPECIALIST_SETTINGS = [
dict(
columns=(
'left_eye_center_x', 'left_eye_center_y',
'right_eye_center_x', 'right_eye_center_y',
),
flip_indices=((0, 2), (1, 3)),
),

dict(
columns=(
'nose_tip_x', 'nose_tip_y',
),
flip_indices=(),
),

dict(
columns=(
'mouth_left_corner_x', 'mouth_left_corner_y',
'mouth_right_corner_x', 'mouth_right_corner_y',
'mouth_center_top_lip_x', 'mouth_center_top_lip_y',
),
flip_indices=((0, 2), (1, 3)),
),

dict(
columns=(
'mouth_center_bottom_lip_x',
'mouth_center_bottom_lip_y',
),
flip_indices=(),
),

dict(
columns=(
'left_eye_inner_corner_x', 'left_eye_inner_corner_y',
'right_eye_inner_corner_x', 'right_eye_inner_corner_y',
'left_eye_outer_corner_x', 'left_eye_outer_corner_y',
'right_eye_outer_corner_x', 'right_eye_outer_corner_y',
),
flip_indices=((0, 2), (1, 3), (4, 6), (5, 7)),
),

dict(
columns=(
'left_eyebrow_inner_end_x', 'left_eyebrow_inner_end_y',
'right_eyebrow_inner_end_x', 'right_eyebrow_inner_end_y',
'left_eyebrow_outer_end_x', 'left_eyebrow_outer_end_y',
'right_eyebrow_outer_end_x', 'right_eyebrow_outer_end_y',
),
flip_indices=((0, 2), (1, 3), (4, 6), (5, 7)),
),
]

我们很早前就讨论过在数据扩充中flip_indices的重要性。在数据介绍部分,我们的load_data()函数也接受一个可选参数,来抽取某些列。我们将在用专项网络预测结果的fit_specialists()中使用这些特性:
from collections import OrderedDict
from sklearn.base import clone

def fit_specialists():
specialists = OrderedDict()

for setting in SPECIALIST_SETTINGS:
cols = setting['columns']
X, y = load2d(cols=cols)

本站声明: 本文章由作者或相关机构授权发布,目的在于传递更多信息,并不代表本站赞同其观点,本站亦不保证或承诺内容真实性等。需要转载请联系该专栏作者,如若文章内容侵犯您的权益,请及时联系本站删除。
换一批
延伸阅读

国际独立第三方检测、检验和认证机构德国莱茵TUV大中华区为全球化IoT开发平台服务商杭州涂鸦信息技术有限公司的智能网关产品Smart Wired Gateway颁发了Matter 1.0认证证书。Matter连接标准通过...

关键字: 智能网关 TE IP 网络技术

北京2022年7月1日 /美通社/ -- 随着数字经济的蓬勃发展和"东数西算"工程全面启动,算力已成为新的生产力。计算场景的多元化、泛在化需要更高效的连接,云计算和一体化大数据中心的新型算力网络体系将...

关键字: 网络技术 IC NI SMART

厦门2022年6月30日 /美通社/ -- 随着WLAN技术的发展,室内场景更倾向于依赖无线通信技术,2021年互联网约有50%的数据流量采用WiFi接入(来源:思...

关键字: Wi-Fi 局域网络 砷化镓 网络技术

5G从2019年年度开始进入商用,从现在已经有接近两年半的时间,为什么还有不少人都没有在使用5G?说起这个话题,每次网友都是七嘴八舌。

关键字: 5G 4G 网络技术

如今的这个时代,已经很难有针对普通人的发财和改变命运的机会!创业你现在也需要累积资本。要么是有一定的高学历顶尖级人才的资本,要么是累积有一定的行业资源资本,要么是你有一定的资金资本(还不一定行)!

关键字: 5G 4G 网络技术

自全球5G开启后,通信企业迎来了最好的发展时期。来自GSMA移动经济报告中数据显示,截至2021年底,全球移动用户达53亿,预计到2025年,全球移动用户将达57亿,其中,到2022年,全球5G连接总数将达到10亿。

关键字: 5G 4G 网络技术

(全球TMT2022年3月3日讯)日前,包括爱立信在内的合作伙伴与中国移动携手多家运营商,在2022 MWC期间发布《5G-Advanced网络技术演进 -- 面向万物智联新时代2.0》白皮书。对5G产业的进展与未来趋...

关键字: Advance 网络技术 爱立信

近日美媒发表一篇题为《中国5G远超美国》的文章称,美国威瑞森电信公司和电话电报公司研发的新5G网络,比此前的4G网络还要慢很多,这些号称全世界最快和最可靠的5G服务,正在误导美国民众。

关键字: 5G 4G 网络技术

1月28日电(记者盖博铭、张漫子)记者28日从北京市科委、中关村管委会获悉,全球首台套5G+8K全业务转播车已开进北京冬奥会,将在冬奥会期间为全球观众带来5G+8K超高清视频体验,全方位展示中国超高清视频产业的能力和水平...

关键字: 5G 8K 网络技术

5G即将实现全面商用,可以预见的是,5G将作为核心底层基础设施渗透到各行业中。5G时代,通信行业产生的电力消耗也可想而知,有相关预测指出,到2025年,通信行业将消耗全球20%的电力。其中,大约80%的能耗来自广泛分布的...

关键字: 5G 3G 网络技术
关闭
关闭