FluxML/Flux.jl

Shape-propagating Chain

已关闭

#703 创建于 2019年3月26日

 (8 条评论) (8 个反应) (0 位负责人)Julia (619 个派生)batch import
discussionenhancementhelp wanted

仓库指标

星标
 (4,725 个星标)
PR 合并指标
 (平均合并 3天 6小时) (30 天内合并 6 个 PR)

描述

It'd be nice to be able to write something like

model = @Chain(
  Input(28^2),
  Dense(32, relu),
  Dense(10),
  softmax)

It's a relatively minor convenience but it does avoid some redundancy when specifying chains, which is tedious to correct and easy to get wrong when trying different layer sizes.

Here's roughly how I imagine this working. The @Chain would expand to something like

shape = nothing
layer1, shape = fromshape(Input, shape, 10)
layer2, shape = fromshape(Dense, shape, 32, relu)
...
Chain(layer1, layer2, ...)

fromshape can then forward to an appropriate constructor or error for non-supported layers. Hopefully this strikes the right balance of simplicity/generality and we don't end up having to turn it into a full shape inference system.

贡献者指南