google-deepmind / google-deepmind/kfac-jax

TypeError: unhashable type: 'Literal'

Open
#310 1 comment 0 reactions 0 assignees View on GitHub
Dominant language
Python
Stars
331
Forks
34
Avg merge
19h 22m
Merged PRs (30d)
9

Description

### Issue Report: **TypeError with Constant Multiplication in network When Using kfac_jax**

#### **Summary**
When using `kfac_jax` with a network, introducing a scaling factor (e.g., `geo_scale`) for lattice parameters in the computational graph causes a `TypeError` related to the use of `Literal`. This occurs whether `geo_scale` is passed as a parameter or defined as a constant multiplier in the computation.

---

#### **Steps to Reproduce**
Here’s an example illustrating the issue(get_jacobian is part of the network):

1. **Working Example: No Scaling**
The following works without errors:
```python
def get_jacobian(params):
p_cell = params['cell'].ravel() # No scaling applied
return jnp.diag(p_cell)
```

2. **Failing Example 1: Direct Multiplication**
Adding a constant multiplier to `params['cell']` causes a `TypeError`:
```python
def get_jacobian(params):
p_cell = params['cell'].ravel() * 1e-3 # Multiplying with a constant
return jnp.diag(p_cell)
```

**Error Raised:**
```
TypeError: unhashable type: 'Literal'
```

3. **Failing Example 2: Adding `geo_scale` Parameter**
Introducing a `geo_scale` parameter also causes the same `TypeError`:
```python
def get_jacobian(params, geo_scale=1e-3):
p_cell = params['cell'].ravel() * geo_scale
return jnp.diag(p_cell)
```

**Error Raised:**
```
TypeError: unhashable type: 'Literal'
```

---

#### **Questions**
1. Is there a recommended approach for handling constants or scaling factors like `geo_scale` in `kfac_jax` workflows to avoid such issues?

---

Contributor guide

Open the contributing guide

Research direction

Start with the provided get_jacobian reproducer in the issue and compare the working unscaled case with both constant-multiplication cases. Trace how kfac_jax handles the resulting Literal during network processing; done means the scaling examples no longer raise TypeError and the behavior is covered by a regression test.

Written by the indexing model from the issue text.

Assessment

Tech stack
python
Domain
machine-learning
Issue type
Bug
Difficulty
4/5
Estimated time
3-5 days
Activity status
Quiet
Clarity
Mostly clear
Newbie friendliness
42/100

Get new issues in your inbox

A short digest of beginner-friendly GitHub issues.