FluxML/Flux.jl

Weights shape not validated against kernel, channels

Open

#2,506 opened on Oct 25, 2024

View on GitHub
 (5 comments) (1 reaction) (0 assignees)Julia (619 forks)batch import
good first issuehelp wanted

Repository metrics

Stars
 (4,725 stars)
PR merge metrics
 (Avg merge 3d 6h) (6 merged PRs in 30d)

Description

weights = Flux.kaiming_normal()(3, 3, 1)
Conv((3, 3), 1 => 1; pad = (1, 1), init = (_...) -> weights)
# Conv((3,), 3 => 1, pad=1)  # 10 parameters

weights = Flux.kaiming_normal()(3, 3, 1, 1)
Conv((3, 3), 1 => 1; pad = (1, 1), init = (_...) -> weights)
# Conv((3, 3), 1 => 1, pad=1)  # 10 parameters

I wanted to strictly specify the weight init for testing, but encountered this odd result. I think there should be validation to ensure that the weight shape matches the kernel size and input channels, and error if there is a mismatch.

Contributor guide