google / google/flax

Common Initializers Does Not Work with Bias

Open
#2,749 6 comments 0 reactions 1 assignee Assigned to @chiamp View on GitHub
Priority: P1 - soon
Dominant language
Jupyter Notebook
Stars
7.3k
Forks
833
Avg merge
5h 11m
Merged PRs (30d)
5

Description

It is expected that bias vector can be initialized not only with zeros but the following code fails.

```python
import jax, jax.numpy as jnp
import flax.linen as nn

layer = nn.Conv(features=32,
kernel_size=(3, 3),
use_bias=True,
bias_init=nn.initializers.lecun_normal())

jax.jit(layer.init)(jax.random.PRNGKey(42), jnp.empty((5, 28, 28, 1)))

# File jax/core.py:1969, in NamedShape.__getitem__(self, idx)
# 1967 try:
# 1968 idx = operator.index(idx)
# -> 1969 return self.__positional[idx]
# 1970 except TypeError:
# 1971 pass
#
# IndexError: tuple index out of range
```

It seems that the issues #1386 is related and maybe #2002 is related as well.

Contributor guide

Open the contributing guide

Assessment

This issue has not been assessed yet.

Get new issues in your inbox

A short digest of beginner-friendly GitHub issues.