Summary
flax.traverse_util.flatten_dict, when called with a separator, produces flattened keys that unflatten_dict cannot faithfully restore for two classes of input. In one case the result is a silently corrupted round-trip; in the other it is a low-level TypeError from deep inside the call. Both are reachable with ordinary dictionaries, and both are surprising given that the flatten_dict docstring directs the reader to unflatten_dict "on how to restore the nested dictionary structure."
The proposal is to fail fast with a clear, actionable error when a separator cannot represent the given keys unambiguously, rather than corrupt the data or raise an opaque error.
Case 1: a key that contains the separator round-trips to a different structure
from flax.traverse_util import flatten_dict, unflatten_dict
xs = {"a/b": 1, "c": {"d": 2}}
flat = flatten_dict(xs, sep="/")
back = unflatten_dict(flat, sep="/")
print(flat) # {'a/b': 1, 'c/d': 2}
print(back) # {'a': {'b': 1}, 'c': {'d': 2}}
print(back == xs) # False
The key "a/b" and the nested path ("c", "d") both flatten to a slash-joined string, so unflatten_dict splits "a/b" into a two-level path and the original structure is lost. No error is raised.
Case 2: a non-string key under a separator raises an opaque error
from flax.traverse_util import flatten_dict
flatten_dict({1: {2: 3}}, sep="/")
# TypeError: sequence item 0: expected str instance, int found
The failure surfaces from sep.join(path) inside _flatten, so the message describes an internal join rather than the actual problem, which is that a separator cannot be applied to non-string keys.
Why this is worth addressing
flatten_dict and unflatten_dict are commonly used to move parameter and state trees to and from a flat, string-keyed representation, for example when interfacing with checkpoint formats or logging. Silent structural corruption of such a tree is difficult to detect downstream, and the non-string-key error gives no indication of its cause. Integer keys and keys containing punctuation both occur in real parameter dictionaries.
Proposed change
When sep is not None, validate each key before joining and raise a single, descriptive error for the two conditions that make the flattened key ambiguous or ill-formed:
- a key that is not a string, and
- a string key that contains
sep.
For example:
ValueError: flatten_dict with sep='/' requires string keys that do not
contain the separator; got key 'a/b' at path ('a/b',)
This turns two silent or opaque failure modes into one clear error, and leaves every input that currently round-trips correctly unchanged. The tuple-key mode (sep=None) is unaffected, since it already round-trips these inputs faithfully.
I am happy to prepare the pull request, including regression tests for both cases and a short note in the flatten_dict docstring, if this direction is agreeable.
Environment
- flax 0.12.8
- jax 0.11.0
- numpy 2.5.2
- Python 3.12
Summary
flax.traverse_util.flatten_dict, when called with a separator, produces flattened keys thatunflatten_dictcannot faithfully restore for two classes of input. In one case the result is a silently corrupted round-trip; in the other it is a low-levelTypeErrorfrom deep inside the call. Both are reachable with ordinary dictionaries, and both are surprising given that theflatten_dictdocstring directs the reader tounflatten_dict"on how to restore the nested dictionary structure."The proposal is to fail fast with a clear, actionable error when a separator cannot represent the given keys unambiguously, rather than corrupt the data or raise an opaque error.
Case 1: a key that contains the separator round-trips to a different structure
The key
"a/b"and the nested path("c", "d")both flatten to a slash-joined string, sounflatten_dictsplits"a/b"into a two-level path and the original structure is lost. No error is raised.Case 2: a non-string key under a separator raises an opaque error
The failure surfaces from
sep.join(path)inside_flatten, so the message describes an internaljoinrather than the actual problem, which is that a separator cannot be applied to non-string keys.Why this is worth addressing
flatten_dictandunflatten_dictare commonly used to move parameter and state trees to and from a flat, string-keyed representation, for example when interfacing with checkpoint formats or logging. Silent structural corruption of such a tree is difficult to detect downstream, and the non-string-key error gives no indication of its cause. Integer keys and keys containing punctuation both occur in real parameter dictionaries.Proposed change
When
sepis notNone, validate each key before joining and raise a single, descriptive error for the two conditions that make the flattened key ambiguous or ill-formed:sep.For example:
This turns two silent or opaque failure modes into one clear error, and leaves every input that currently round-trips correctly unchanged. The tuple-key mode (
sep=None) is unaffected, since it already round-trips these inputs faithfully.I am happy to prepare the pull request, including regression tests for both cases and a short note in the
flatten_dictdocstring, if this direction is agreeable.Environment