| |
| |
|
|
|
|
| from __future__ import annotations |
|
|
|
|
| def _scan_partition_sizes(eqn) -> tuple[int, int, int]: |
|
|
| consts, carry, xs = eqn.params["ft_in"].unpack() |
| return len(consts), len(carry), len(xs) |
|
|
|
|
| def _extend_scan_flat_trees( |
| params: dict, |
| *, |
| extra_xs: int, |
| extra_ys: int, |
| ) -> None: |
|
|
| from jax._src import flattree as _ft |
|
|
| consts, carry, xs = params["ft_in"].unpack() |
| carry_out, ys = params["ft_out"].unpack() |
| if extra_xs: |
| xs = _ft.pack((xs, _ft.nones(extra_xs))) |
| if extra_ys: |
| ys = _ft.pack((ys, _ft.nones(extra_ys))) |
| params["ft_in"] = _ft.pack((consts, carry, xs)) |
| params["ft_out"] = _ft.pack((carry_out, ys)) |
|
|
|
|
| def _patch() -> None: |
|
|
| from jax._src import source_info_util as jex_source_info_util |
| from kfac_jax._src import tag_graph_matcher as tgm |
|
|
| if getattr(tgm.eval_jaxpr_eqn, "__hamiltonzero_patched__", False): |
| return |
|
|
| def eval_jaxpr_eqn(eqn, in_values): |
| bind_params = eqn.primitive.get_bind_params(eqn.params) |
| user_context = jex_source_info_util.user_context |
| with user_context(eqn.source_info.traceback): |
| output = eqn.primitive.bind(*in_values, **bind_params) |
| return [output] if not isinstance(output, list) else output |
|
|
| eval_jaxpr_eqn.__hamiltonzero_patched__ = True |
| tgm.eval_jaxpr_eqn = eval_jaxpr_eqn |
|
|
|
|
| def _patch_allow_multiple_registrations() -> None: |
|
|
| import threading |
|
|
| from kfac_jax._src import tag_graph_matcher as tgm |
|
|
| if getattr(tgm.auto_register_tags, "__hamiltonzero_allow_multi__", False): |
| return |
|
|
| _orig_auto = tgm.auto_register_tags |
| _orig_check = tgm.TaggedFunction.check_multiple_registrations |
| _state = threading.local() |
|
|
| def auto_register_tags( |
| func, |
| func_args, |
| *, |
| allow_multiple_registrations: bool = False, |
| **kwargs, |
| ): |
| prev = getattr(_state, "allow", False) |
| _state.allow = bool(allow_multiple_registrations) |
| try: |
| return _orig_auto(func, func_args, **kwargs) |
| finally: |
| _state.allow = prev |
|
|
| def check_multiple_registrations(self): |
| if getattr(_state, "allow", False): |
| return |
| return _orig_check(self) |
|
|
| auto_register_tags.__hamiltonzero_allow_multi__ = True |
| tgm.auto_register_tags = auto_register_tags |
| tgm.TaggedFunction.check_multiple_registrations = check_multiple_registrations |
|
|
|
|
| def _patch_orphan_registration_in_sub_graphs() -> None: |
|
|
| from kfac_jax._src import tag_graph_matcher as tgm |
|
|
| if getattr(tgm._auto_register_tags, "__hamiltonzero_orphan_subgraph__", False): |
| return |
|
|
| _orig = tgm._auto_register_tags |
|
|
| def _patched(graph, *args, register_orphans=True, **kwargs): |
|
|
| return _orig(graph, *args, register_orphans=True, **kwargs) |
|
|
| _patched.__hamiltonzero_orphan_subgraph__ = True |
| tgm._auto_register_tags = _patched |
|
|
|
|
| def _patch_manual_tag_outputs_that_are_graph_inputs() -> None: |
|
|
| from kfac_jax._src import tag_graph_matcher as tgm |
| import jax.extend as jex |
|
|
| graph_cls = tgm.JaxprGraph |
| if getattr( |
| graph_cls.sub_graph_eqns, |
| "__hamiltonzero_graph_input_tag_output__", |
| False, |
| ): |
| return |
| original = graph_cls.sub_graph_eqns |
|
|
| def sub_graph_eqns(self, root_vars, leaf_vars): |
| kept = [] |
| for value in leaf_vars: |
| if ( |
| isinstance(value, jex.core.Literal) |
| or value in self.params_vars |
| or value in self.var_to_creation_op |
| ): |
| kept.append(value) |
| elif value in self.jaxpr.invars: |
| continue |
| else: |
| raise KeyError(value) |
| return original(self, root_vars, tuple(kept)) |
|
|
| sub_graph_eqns.__hamiltonzero_graph_input_tag_output__ = True |
| graph_cls.sub_graph_eqns = sub_graph_eqns |
|
|
|
|
| def _patch_hoist_layer_tags_from_scan() -> None: |
|
|
| from kfac_jax._src import tag_graph_matcher as tgm |
| from kfac_jax._src import layers_and_loss_tags as tags |
| import jax |
| from jax._src import core as _jcore |
| from kfac_jax._src.tag_graph_matcher import ( |
| ClosedJaxpr, |
| HIGHER_ORDER_NAMES, |
| to_closed_jaxpr, |
| to_jaxpr_or_closed_jaxpr, |
| ) |
|
|
| from jax.extend.core import gensym, new_jaxpr_eqn |
|
|
| if getattr(tgm.clean_layer_tags_jaxpr, "__hamiltonzero_hoist_tags__", False): |
| return |
|
|
| _orig_clean_layer = tgm.clean_layer_tags_jaxpr |
|
|
| def _make_transpose_swap01_eqn(scan_outvar, make_var_func): |
|
|
| ndim = len(scan_outvar.aval.shape) |
| if ndim < 2: |
| return None, scan_outvar |
| from jax._src.lax import lax as _jlax |
|
|
| permutation = (1, 0) + tuple(range(2, ndim)) |
| new_shape = tuple(scan_outvar.aval.shape[p] for p in permutation) |
| new_aval = _jcore.ShapedArray(new_shape, scan_outvar.aval.dtype) |
| new_outvar = make_var_func(new_aval) |
| eqn = new_jaxpr_eqn( |
| invars=[scan_outvar], |
| outvars=[new_outvar], |
| primitive=_jlax.transpose_p, |
| params={"permutation": permutation}, |
| effects=frozenset(), |
| ) |
| return eqn, new_outvar |
|
|
| def _hoist_tags_recursive(closed_jaxpr, make_var_func): |
|
|
| new_eqns = [] |
| hoisted_tags = [] |
|
|
| for eqn in closed_jaxpr.jaxpr.eqns: |
| if eqn.primitive.name not in HIGHER_ORDER_NAMES: |
| new_eqns.append(eqn) |
| continue |
| if eqn.primitive.name == "cond": |
| new_eqns.append(eqn) |
| continue |
| if eqn.primitive.name == "while": |
| body_jaxpr = eqn.params["body_jaxpr"] |
| key = "body_jaxpr" |
| supports_extension = False |
| elif eqn.primitive.name == "scan": |
| body_jaxpr = eqn.params["jaxpr"] |
| key = "jaxpr" |
| supports_extension = True |
| elif eqn.primitive.name == "pjit": |
| body_jaxpr = eqn.params["jaxpr"] |
| key = "jaxpr" |
| supports_extension = False |
| elif eqn.primitive.name in ("xla_call", "xla_pmap"): |
| body_jaxpr = eqn.params["call_jaxpr"] |
| key = "call_jaxpr" |
| supports_extension = False |
| else: |
| new_eqns.append(eqn) |
| continue |
|
|
| body_closed = to_closed_jaxpr(body_jaxpr) |
| new_body_closed, nested_hoisted = _hoist_tags_recursive( |
| body_closed, make_var_func |
| ) |
|
|
| body_invars = body_jaxpr.jaxpr.invars |
| body_eqns_no_tags = [] |
| tag_var_map = {} |
|
|
| new_body_captures: list = [] |
| new_body_capture_id_to_idx: dict[int, int] = {} |
|
|
| output_var_to_aux_xs: dict[int, tuple] = {} |
|
|
| scan_length = eqn.params["length"] if eqn.primitive.name == "scan" else None |
|
|
| deferred_specs: list = [] |
|
|
| import jax.extend as _jex_chain |
|
|
| def _resolve_tag_chain(w): |
|
|
| while not isinstance(w, _jex_chain.core.Literal) and w in tag_var_map: |
| w = tag_var_map[w] |
| return w |
|
|
| for body_eqn in new_body_closed.jaxpr.eqns: |
| if not isinstance(body_eqn.primitive, tags.LayerTag): |
| body_eqns_no_tags.append(body_eqn) |
| continue |
|
|
| meta = body_eqn.params["meta"] |
| for ind1, ind2 in enumerate(meta.outputs_index): |
| tag_var_map[body_eqn.outvars[ind1]] = body_eqn.invars[ind2] |
|
|
| params_index_set = set(meta.params_index) |
| partial_invars: list = [None] * len(body_eqn.invars) |
| deferred: list = [] |
| hoistable = True |
|
|
| ( |
| _scan_num_consts, |
| _scan_num_carry, |
| _, |
| ) = _scan_partition_sizes(eqn) |
| _xs_threshold = _scan_num_consts + _scan_num_carry |
|
|
| tag_has_xs_iterated_params = False |
| output_idx_set = set(meta.outputs_index) |
|
|
| for _arg_idx_pre, _v_pre in enumerate(body_eqn.invars): |
| if _arg_idx_pre not in params_index_set: |
| continue |
| _v_pre_resolved = _resolve_tag_chain(_v_pre) |
| if _v_pre_resolved in body_invars: |
| _idx = body_invars.index(_v_pre_resolved) |
| if eqn.primitive.name == "scan" and _idx >= _xs_threshold: |
| tag_has_xs_iterated_params = True |
| break |
|
|
| use_aux_xs_for_outputs = supports_extension and scan_length is not None |
|
|
| use_accumulating_aux_base = False |
| if ( |
| use_aux_xs_for_outputs |
| and not tag_has_xs_iterated_params |
| and getattr(meta, "variant", None) == "dense" |
| and len(meta.inputs_index) == 1 |
| and len(meta.outputs_index) == 1 |
| and len(meta.params_index) >= 1 |
| ): |
| _const_indices = [] |
| for _candidate_idx in ( |
| *meta.inputs_index, |
| *meta.params_index, |
| ): |
| _candidate = _resolve_tag_chain(body_eqn.invars[_candidate_idx]) |
| if _candidate not in body_invars: |
| _const_indices = [] |
| break |
| _candidate_body_idx = body_invars.index(_candidate) |
| if _candidate_body_idx >= _scan_num_consts: |
| _const_indices = [] |
| break |
| _const_indices.append(_candidate_body_idx) |
| _output_candidate = _resolve_tag_chain( |
| body_eqn.invars[meta.outputs_index[0]] |
| ) |
| if ( |
| len(_const_indices) |
| == len(meta.inputs_index) + len(meta.params_index) |
| and _output_candidate not in body_invars |
| ): |
| _outer_input = eqn.invars[_const_indices[0]] |
| use_accumulating_aux_base = ( |
| _outer_input.aval.shape[:-1] |
| == _output_candidate.aval.shape[:-1] |
| ) |
|
|
| for arg_idx, v in enumerate(body_eqn.invars): |
| v_resolved = _resolve_tag_chain(v) |
| if v_resolved in body_invars: |
| idx = body_invars.index(v_resolved) |
| is_xs_iterated = ( |
| eqn.primitive.name == "scan" and idx >= _xs_threshold |
| ) |
| is_param = arg_idx in params_index_set |
| if is_xs_iterated and is_param: |
| tag_has_xs_iterated_params = True |
|
|
| _slot_is_scan_carry = ( |
| eqn.primitive.name == "scan" |
| and idx >= _scan_num_consts |
| and idx < _xs_threshold |
| ) |
| if ( |
| (tag_has_xs_iterated_params or _slot_is_scan_carry) |
| and not is_param |
| and not is_xs_iterated |
| and supports_extension |
| and scan_length is not None |
| ): |
| cap_id = id(v_resolved) |
| if cap_id not in new_body_capture_id_to_idx: |
| new_body_capture_id_to_idx[cap_id] = len( |
| new_body_captures, |
| ) |
| new_body_captures.append(v_resolved) |
| deferred.append( |
| (arg_idx, new_body_capture_id_to_idx[cap_id]), |
| ) |
| continue |
|
|
| partial_invars[arg_idx] = eqn.invars[idx] |
| elif arg_idx in params_index_set: |
| hoistable = False |
| break |
| elif ( |
| arg_idx in output_idx_set |
| and use_aux_xs_for_outputs |
| and supports_extension |
| and scan_length is not None |
| ): |
| body_v_id = id(v_resolved) |
| if body_v_id not in output_var_to_aux_xs: |
| aux_body_invar = make_var_func(v_resolved.aval) |
| aux_outer_aval = _jcore.ShapedArray( |
| (scan_length, *v_resolved.aval.shape), |
| v_resolved.aval.dtype, |
| ) |
| aux_outer_var = make_var_func(aux_outer_aval) |
| aux_tag_var = ( |
| make_var_func(v_resolved.aval) |
| if use_accumulating_aux_base |
| else aux_outer_var |
| ) |
| output_var_to_aux_xs[body_v_id] = ( |
| v_resolved, |
| aux_body_invar, |
| aux_outer_var, |
| aux_tag_var, |
| ) |
| ( |
| _, |
| _, |
| aux_outer_var, |
| aux_tag_var, |
| ) = output_var_to_aux_xs[body_v_id] |
| if ( |
| aux_tag_var is not aux_outer_var |
| ) != use_accumulating_aux_base: |
| raise ValueError( |
| "Conflicting scan accumulation contracts for " |
| "the same hoisted layer output." |
| ) |
| partial_invars[arg_idx] = aux_tag_var |
| elif supports_extension: |
| cap_id = id(v_resolved) |
| if cap_id not in new_body_capture_id_to_idx: |
| new_body_capture_id_to_idx[cap_id] = len( |
| new_body_captures, |
| ) |
| new_body_captures.append(v_resolved) |
| deferred.append( |
| (arg_idx, new_body_capture_id_to_idx[cap_id]), |
| ) |
| else: |
| import jax.extend as _jex |
| import numpy as _np |
|
|
| zero_val = _np.zeros( |
| v_resolved.aval.shape, |
| dtype=v_resolved.aval.dtype, |
| ) |
| partial_invars[arg_idx] = _jex.core.Literal( |
| zero_val, |
| v_resolved.aval, |
| ) |
|
|
| if not hoistable: |
| body_eqns_no_tags.append(body_eqn) |
| continue |
|
|
| deferred_specs.append( |
| ( |
| body_eqn, |
| partial_invars, |
| deferred, |
| tag_has_xs_iterated_params, |
| use_aux_xs_for_outputs, |
| ) |
| ) |
|
|
| import jax.extend as _jex |
|
|
| def _remap_invars(eqns): |
|
|
| out = [] |
| for e in eqns: |
| new_invars = [ |
| _resolve_tag_chain(w) |
| if not isinstance(w, _jex.core.Literal) |
| else w |
| for w in e.invars |
| ] |
| out.append(e.replace(invars=new_invars)) |
| return out |
|
|
| body_eqns_no_tags = _remap_invars(body_eqns_no_tags) |
| new_body_outvars = [ |
| _resolve_tag_chain(v) if not isinstance(v, _jex.core.Literal) else v |
| for v in new_body_closed.jaxpr.outvars |
| ] |
|
|
| if output_var_to_aux_xs: |
| from jax._src.lax import lax as _jlax |
|
|
| aug_for_id: dict[int, tuple] = {} |
| for body_v_id, ( |
| body_v, |
| aux_body_invar, |
| _, |
| _, |
| ) in output_var_to_aux_xs.items(): |
| aug_var = make_var_func(body_v.aval) |
| aug_for_id[body_v_id] = (aug_var, aux_body_invar) |
|
|
| seen_ids: set[int] = set() |
|
|
| def _retarget_to_aug(w): |
| if isinstance(w, _jex.core.Literal): |
| return w |
| wid = id(w) |
| if wid in aug_for_id and wid in seen_ids: |
| return aug_for_id[wid][0] |
| return w |
|
|
| augmented_eqns = [] |
| for body_eqn_clean in body_eqns_no_tags: |
| augmented_eqns.append( |
| body_eqn_clean.replace( |
| invars=[_retarget_to_aug(w) for w in body_eqn_clean.invars] |
| ) |
| ) |
| for o in body_eqn_clean.outvars: |
| oid = id(o) |
| if oid in aug_for_id and oid not in seen_ids: |
| aug_var, aux_body_invar = aug_for_id[oid] |
| aug_eqn = new_jaxpr_eqn( |
| invars=[o, aux_body_invar], |
| outvars=[aug_var], |
| primitive=_jlax.add_p, |
| params={}, |
| effects=frozenset(), |
| ) |
| augmented_eqns.append(aug_eqn) |
| seen_ids.add(oid) |
|
|
| assert seen_ids == set(aug_for_id.keys()), ( |
| "aux-xs injection: some output Vars not encountered as " |
| "body-eqn outvars" |
| ) |
| body_eqns_no_tags = augmented_eqns |
| new_body_outvars = [ |
| _retarget_to_aug(v) if not isinstance(v, _jex.core.Literal) else v |
| for v in new_body_outvars |
| ] |
|
|
| new_body_outvars = list(new_body_outvars) + list(new_body_captures) |
|
|
| new_body_invars_list = list(new_body_closed.jaxpr.invars) + [ |
| aux_body_invar |
| for _, aux_body_invar, _, _ in output_var_to_aux_xs.values() |
| ] |
|
|
| new_body_jaxpr = new_body_closed.jaxpr.replace( |
| eqns=body_eqns_no_tags, |
| outvars=new_body_outvars, |
| invars=new_body_invars_list, |
| ) |
| new_body_closed_clean = ClosedJaxpr( |
| new_body_jaxpr, |
| new_body_closed.consts, |
| ) |
|
|
| params_dict = dict(**eqn.params) |
| params_dict[key] = to_jaxpr_or_closed_jaxpr( |
| new_body_closed_clean, |
| body_jaxpr, |
| ) |
| if eqn.primitive.name == "scan": |
| _extend_scan_flat_trees( |
| params_dict, |
| extra_xs=len(output_var_to_aux_xs), |
| extra_ys=len(new_body_captures), |
| ) |
|
|
| if output_var_to_aux_xs: |
| from jax._src.lax import lax as _jlax |
| import jax.extend as _jex |
| import numpy as _np |
|
|
| for ( |
| body_v, |
| _aux_body_invar, |
| aux_outer_var, |
| aux_tag_var, |
| ) in output_var_to_aux_xs.values(): |
| zero_scalar_aval = _jcore.ShapedArray( |
| (), |
| body_v.aval.dtype, |
| ) |
| zero_scalar_literal = _jex.core.Literal( |
| _np.array(0.0, dtype=body_v.aval.dtype), |
| zero_scalar_aval, |
| ) |
| first_bcast_outvar = aux_tag_var |
| first_bcast_shape = aux_tag_var.aval.shape |
| bcast_eqn = new_jaxpr_eqn( |
| invars=[zero_scalar_literal], |
| outvars=[first_bcast_outvar], |
| primitive=_jlax.broadcast_in_dim_p, |
| params={ |
| "shape": first_bcast_shape, |
| "broadcast_dimensions": (), |
| "sharding": None, |
| }, |
| effects=frozenset(), |
| ) |
| new_eqns.append(bcast_eqn) |
| if aux_tag_var is not aux_outer_var: |
| expand_eqn = new_jaxpr_eqn( |
| invars=[aux_tag_var], |
| outvars=[aux_outer_var], |
| primitive=_jlax.broadcast_in_dim_p, |
| params={ |
| "shape": ( |
| scan_length, |
| *body_v.aval.shape, |
| ), |
| "broadcast_dimensions": tuple( |
| range(1, body_v.aval.ndim + 1) |
| ), |
| "sharding": None, |
| }, |
| effects=frozenset(), |
| ) |
| new_eqns.append(expand_eqn) |
|
|
| new_capture_outvars = [] |
| for cap_v in new_body_captures: |
| cap_aval = _jcore.ShapedArray( |
| (scan_length, *cap_v.aval.shape), |
| cap_v.aval.dtype, |
| ) |
| new_capture_outvars.append(make_var_func(cap_aval)) |
|
|
| new_scan_invars = list(eqn.invars) + [ |
| aux_outer_var |
| for _, _, aux_outer_var, _ in output_var_to_aux_xs.values() |
| ] |
| new_eqn = eqn.replace( |
| params=params_dict, |
| invars=new_scan_invars, |
| outvars=list(eqn.outvars) + list(new_capture_outvars), |
| ) |
| new_eqns.append(new_eqn) |
|
|
| aux_xs_capture_ids: set[int] = set() |
| for _be, _pi, _df, _has_xs, _use_aux in deferred_specs: |
| if _use_aux: |
| for _aidx, _cidx in _df: |
| aux_xs_capture_ids.add(_cidx) |
| transposed_outvars: list = [] |
| for _cap_idx, cap_outvar in enumerate(new_capture_outvars): |
| if _cap_idx in aux_xs_capture_ids: |
| transposed_outvars.append(cap_outvar) |
| continue |
| t_eqn, t_outvar = _make_transpose_swap01_eqn( |
| cap_outvar, |
| make_var_func, |
| ) |
| if t_eqn is not None: |
| new_eqns.append(t_eqn) |
| transposed_outvars.append(t_outvar) |
|
|
| for ( |
| body_eqn, |
| partial_invars, |
| deferred, |
| has_xs_iter_params, |
| _use_aux_xs, |
| ) in deferred_specs: |
| final_invars = list(partial_invars) |
| for arg_idx, cap_idx in deferred: |
| final_invars[arg_idx] = transposed_outvars[cap_idx] |
| new_outvars = [make_var_func(v.aval) for v in body_eqn.outvars] |
|
|
| hoisted_params = body_eqn.params |
| if has_xs_iter_params: |
| import dataclasses as _dc |
|
|
| orig_meta = body_eqn.params["meta"] |
|
|
| _v = orig_meta.variant or "" |
| new_variant = None |
| if _v == "scale_and_shift": |
| new_variant = "stacked_scale_and_shift" |
| elif _v == "structural_repeated_dense": |
| new_variant = "structural_stacked_repeated_dense" |
| elif _v == "structural_scale_and_shift": |
| new_variant = "structural_stacked_scale_and_shift" |
| if new_variant is not None: |
| new_meta = _dc.replace(orig_meta, variant=new_variant) |
| hoisted_params = {**body_eqn.params, "meta": new_meta} |
| hoisted_tags.append( |
| new_jaxpr_eqn( |
| invars=final_invars, |
| outvars=new_outvars, |
| primitive=body_eqn.primitive, |
| params=hoisted_params, |
| effects=body_eqn.effects, |
| ) |
| ) |
|
|
| for nh_eqn in nested_hoisted: |
| outer_remapped = [] |
| ok = True |
| for v in nh_eqn.invars: |
| if v in body_jaxpr.jaxpr.invars: |
| idx = body_jaxpr.jaxpr.invars.index(v) |
| outer_remapped.append(eqn.invars[idx]) |
| else: |
| ok = False |
| break |
| if ok: |
| new_outvars2 = [make_var_func(v.aval) for v in nh_eqn.outvars] |
| hoisted_tags.append( |
| new_jaxpr_eqn( |
| invars=outer_remapped, |
| outvars=new_outvars2, |
| primitive=nh_eqn.primitive, |
| params=nh_eqn.params, |
| effects=nh_eqn.effects, |
| ) |
| ) |
|
|
| new_closed = ClosedJaxpr( |
| closed_jaxpr.jaxpr.replace(eqns=new_eqns), |
| closed_jaxpr.consts, |
| ) |
| return new_closed, hoisted_tags |
|
|
| def clean_layer_tags_jaxpr_patched(jaxpr, only_remove_auto_tags=False): |
|
|
| closed = to_closed_jaxpr(jaxpr) |
| make_var_func = gensym() |
| closed, hoisted = _hoist_tags_recursive(closed, make_var_func) |
|
|
| seen_param_keys = set() |
| deduped = [] |
| for h in hoisted: |
| meta = h.params["meta"] |
| key = tuple(id(h.invars[i]) for i in meta.params_index) |
| if key in seen_param_keys: |
| continue |
| seen_param_keys.add(key) |
| deduped.append(h) |
| hoisted = deduped |
|
|
| if hoisted: |
| new_eqns = list(closed.jaxpr.eqns) + list(hoisted) |
| closed = ClosedJaxpr( |
| closed.jaxpr.replace(eqns=new_eqns), |
| closed.consts, |
| ) |
| return _orig_clean_layer( |
| to_jaxpr_or_closed_jaxpr(closed, jaxpr), |
| only_remove_auto_tags=only_remove_auto_tags, |
| ) |
|
|
| clean_layer_tags_jaxpr_patched.__hamiltonzero_hoist_tags__ = True |
| tgm.clean_layer_tags_jaxpr = clean_layer_tags_jaxpr_patched |
|
|
|
|
| def _patch_kfactor_identity_init() -> None: |
|
|
| import jax.numpy as _jnp |
|
|
| from kfac_jax._src.curvature_blocks import ( |
| kronecker_factored as _kf, |
| ) |
| from kfac_jax._src import utils as _kfac_utils |
|
|
| if getattr(_kf.KroneckerFactored._init, "__hamiltonzero_kfactor_identity__", False): |
| return |
|
|
| _orig_kf_init = _kf.KroneckerFactored._init |
|
|
| def _patched_kf_init( |
| self, rng, exact_powers_to_cache, approx_powers_to_cache, cache_eigenvalues |
| ): |
|
|
| cache = {} |
| factors = [] |
| for i, d in enumerate(self.array_shape): |
| eye = _jnp.eye(d, dtype=self.dtype) * _jnp.asarray(1.0, dtype=self.dtype) |
|
|
| wma = _kfac_utils.WeightedMovingAverage( |
| value=eye, |
| weight=_jnp.asarray(1.0, dtype=self.dtype), |
| ) |
| factors.append(wma) |
| if cache_eigenvalues or exact_powers_to_cache: |
| cache[f"{i}_factor_eigenvalues"] = _jnp.ones((d,), dtype=self.dtype) |
| if exact_powers_to_cache: |
| cache[f"{i}_factor_eigen_vectors"] = _jnp.eye(d, dtype=self.dtype) |
| for power in approx_powers_to_cache: |
| if power != -1: |
| raise NotImplementedError( |
| f"Approximations for power {power} not implemented." |
| ) |
| if str(power) not in cache: |
| cache[str(power)] = {} |
| cache[str(power)][f"{i}_factor"] = _jnp.eye(d, dtype=self.dtype) |
| return _kf.KroneckerFactored.State( |
| cache=cache, |
| factors=tuple(factors), |
| ) |
|
|
| _patched_kf_init.__hamiltonzero_kfactor_identity__ = True |
| _kf.KroneckerFactored._init = _patched_kf_init |
|
|
| if getattr( |
| _kf.RepeatedDenseKroneckerFactored._init, |
| "__hamiltonzero_avg_repeats_one__", |
| False, |
| ): |
| return |
|
|
| def _patched_rd_init( |
| self, rng, exact_powers_to_cache, approx_powers_to_cache, cache_eigenvalues |
| ): |
| super_state = _kf.KroneckerFactored._init( |
| self, |
| rng, |
| exact_powers_to_cache, |
| approx_powers_to_cache, |
| cache_eigenvalues, |
| ) |
| avg = _kfac_utils.WeightedMovingAverage( |
| value=_jnp.asarray(1.0, dtype=self.dtype), |
| weight=_jnp.asarray(1.0, dtype=self.dtype), |
| ) |
| return _kf.RepeatedDenseKroneckerFactored.State( |
| average_repeats=avg, |
| **super_state.__dict__, |
| ) |
|
|
| _patched_rd_init.__hamiltonzero_avg_repeats_one__ = True |
| _kf.RepeatedDenseKroneckerFactored._init = _patched_rd_init |
|
|
|
|
| def _patch_pi_adjusted_kronecker_factors_floor() -> None: |
|
|
| import jax.numpy as _jnp |
| from kfac_jax._src.utils import math as _kfac_math |
|
|
| if getattr( |
| _kfac_math.pi_adjusted_kronecker_factors, |
| "__hamiltonzero_kron_floor__", |
| False, |
| ): |
| return |
|
|
| _orig = _kfac_math.pi_adjusted_kronecker_factors |
| EPS_FLOOR = 1e-6 |
| EPS_REL = 1e-4 |
|
|
| def _shift_from_avg_diag(avg_diag, scale): |
| eps_abs = _jnp.asarray(EPS_FLOOR, dtype=avg_diag.dtype) |
| eps_rel = _jnp.asarray(EPS_REL, dtype=avg_diag.dtype) |
| floor = _jnp.maximum(eps_abs, eps_rel * scale) |
| return _jnp.maximum(floor, floor - avg_diag) |
|
|
| def _floor_factor(f): |
| if f.ndim == 0 or f.size == 1: |
| return f + _shift_from_avg_diag(f, _jnp.abs(f)) |
| if f.ndim == 1: |
| avg_diag = _jnp.mean(f) |
| scale = _jnp.max(_jnp.abs(f)) |
| return f + _shift_from_avg_diag(avg_diag, scale) |
| if f.ndim == 2: |
| d = f.shape[-1] |
| diag = _jnp.diagonal(f) |
| avg_diag = _jnp.sum(diag) / d |
| scale = _jnp.max(diag) |
| shift = _shift_from_avg_diag(avg_diag, scale) |
| return f + shift * _jnp.eye(d, dtype=f.dtype) |
|
|
| if f.ndim >= 3 and f.shape[-1] == f.shape[-2]: |
| d = f.shape[-1] |
| eye = _jnp.eye(d, dtype=f.dtype) |
| for _ in range(f.ndim - 2): |
| eye = eye[None, ...] |
| diag = _jnp.diagonal(f, axis1=-2, axis2=-1) |
| avg_diag = _jnp.mean(diag, axis=-1) |
| scale = _jnp.max(diag, axis=-1) |
| shift = _shift_from_avg_diag(avg_diag, scale) |
| return f + shift[..., None, None] * eye |
| return f |
|
|
| def patched(*factors, damping): |
| floored = tuple(_floor_factor(f) for f in factors) |
| return _orig(*floored, damping=damping) |
|
|
| patched.__hamiltonzero_kron_floor__ = True |
| _kfac_math.pi_adjusted_kronecker_factors = patched |
|
|
| from kfac_jax._src import utils as _kfac_utils_pkg |
|
|
| if hasattr(_kfac_utils_pkg, "pi_adjusted_kronecker_factors"): |
| _kfac_utils_pkg.pi_adjusted_kronecker_factors = patched |
|
|
|
|
| def _patch_nested_scan_parent_walk() -> None: |
|
|
| from kfac_jax._src import tag_graph_matcher as tgm |
|
|
| _TagLocation = tgm.TagLocation |
| if getattr(_TagLocation, "__hamiltonzero_nested_parent_walk__", False): |
| return |
|
|
| def _invars_of(eqn): |
| nm = eqn.primitive.name |
| if nm in ("scan", "pjit"): |
| return eqn.params["jaxpr"].jaxpr.invars |
| if nm == "while": |
| return eqn.params["body_jaxpr"].jaxpr.invars |
| if nm in ("xla_call", "xla_pmap"): |
| return eqn.params["call_jaxpr"].invars |
| raise NotImplementedError(f"higher-order primitive {nm!r}") |
|
|
| def _walk(param_vars, eqns_in_order): |
| for eqn, _ in eqns_in_order: |
| invars = _invars_of(eqn) |
| p_indexes = [invars.index(p) for p in param_vars] |
| param_vars = tuple(eqn.invars[pi] for pi in p_indexes) |
| return param_vars |
|
|
| def _top_level_parameters(self): |
| pv = self.bottom_level_parameters |
| return _walk(pv, list(self.parent_equations)) |
|
|
| def _full_name_ordered(self, eqns_in_order): |
| param_vars = self.bottom_level_parameters |
| parts = [] |
| for eqn, n in eqns_in_order: |
| nm = eqn.primitive.name |
| invars = _invars_of(eqn) |
| p_indexes = [invars.index(p) for p in param_vars] |
| piece = f"{nm}_{n}/" |
| if nm == "scan": |
| num_consts, _, _ = _scan_partition_sizes(eqn) |
| checks = [pi < num_consts for pi in p_indexes] |
| if not (all(checks) or all(not ci for ci in checks)): |
| raise ValueError( |
| "Parameters inside scan of the same tag are not both " |
| "carry or const." |
| ) |
| piece = piece + ("const/" if all(checks) else "carry/") |
| parts.append(piece) |
| param_vars = [eqn.invars[pi] for pi in p_indexes] |
|
|
| prefix = "".join(reversed(parts)) |
| return prefix + self.base_name |
|
|
| def _full_name(self): |
| return _full_name_ordered(self, list(self.parent_equations)) |
|
|
| _TagLocation.top_level_parameters = property(_top_level_parameters) |
| _TagLocation.full_name = property(_full_name) |
| _TagLocation.__hamiltonzero_nested_parent_walk__ = True |
|
|
|
|
| _patch() |
| _patch_allow_multiple_registrations() |
| _patch_orphan_registration_in_sub_graphs() |
| _patch_manual_tag_outputs_that_are_graph_inputs() |
| _patch_hoist_layer_tags_from_scan() |
| _patch_kfactor_identity_init() |
| _patch_pi_adjusted_kronecker_factors_floor() |
| _patch_nested_scan_parent_walk() |
|
|
|
|
| __all__: list[str] = [] |
|
|