File size: 10,621 Bytes
3e02ab8 | 1 2 3 4 5 6 7 8 9 10 11 12 13 14 15 16 17 18 19 20 21 22 23 24 25 26 27 28 29 30 31 32 33 34 35 36 37 38 39 40 41 42 43 44 45 46 47 48 49 50 51 52 53 54 55 56 57 58 59 60 61 62 63 64 65 66 67 68 69 70 71 72 73 74 75 76 77 78 79 80 81 82 83 84 85 86 87 88 89 90 91 92 93 94 95 96 97 98 99 100 101 102 103 104 105 106 107 108 109 110 111 112 113 114 115 116 117 118 119 120 121 122 123 124 125 126 127 128 129 130 131 132 133 134 135 136 137 138 139 140 141 142 143 144 145 146 147 148 149 150 151 152 153 154 155 156 157 158 159 160 161 162 163 164 165 166 167 168 169 170 171 172 173 174 175 176 177 178 179 180 181 182 183 184 185 186 187 188 189 190 191 192 193 194 195 196 197 198 199 200 201 202 203 204 205 206 207 208 209 210 211 212 213 214 215 216 217 218 219 220 221 222 223 224 225 226 227 228 229 230 231 232 233 234 235 236 237 | # This file is a part of the `nequip` package. Please see LICENSE and README at the root for information on using it.
import torch
from onescience.datapipes.materials.nequip import AtomicDataDict
from .graph_model import GraphModel
from ._graph_mixin import GraphModuleMixin
from onescience.utils.nequip.internal.dtype import (
test_model_output_similarity_by_dtype,
_pt2_compile_error_message,
)
from onescience.utils.nequip.internal.fx import nequip_make_fx
from onescience.utils.nequip.internal.dtype import dtype_to_name
from typing import Dict, Sequence, List, Optional, Any, Final
from torch.func import functional_call
def _list_to_dict(
keys: Sequence[str], args: List[torch.Tensor]
) -> Dict[str, torch.Tensor]:
return {key: arg for key, arg in zip(keys, args)}
def _list_from_dict(
keys: Sequence[str], data: Dict[str, torch.Tensor]
) -> List[torch.Tensor]:
return [data[key] for key in keys]
class ListInputOutputWrapper(torch.nn.Module):
"""
Wraps a ``torch.nn.Module`` that takes and returns ``Dict[str, torch.Tensor]`` to have it take and return ``Sequence[torch.Tensor]`` for specified input and output fields.
"""
def __init__(
self,
model: torch.nn.Module,
input_keys: Sequence[str],
output_keys: Sequence[str],
):
super().__init__()
self.model = model
self.input_keys = list(input_keys)
self.output_keys = list(output_keys)
def forward(self, *args: torch.Tensor) -> List[torch.Tensor]:
inputs = _list_to_dict(self.input_keys, args)
outputs = self.model(inputs)
return _list_from_dict(self.output_keys, outputs)
class DictInputOutputWrapper(torch.nn.Module):
"""
Wraps a model that takes and returns ``Sequence[torch.Tensor]`` to have it take and return ``Dict[str, torch.Tensor]`` for specified input and output fields (i.e. the opposite of ``ListInputOutputWrapper``).
"""
def __init__(self, model, input_keys: List[str], output_keys: List[str]):
super().__init__()
self.model = model
self.input_keys = input_keys
self.output_keys = output_keys
def forward(self, data: AtomicDataDict.Type) -> AtomicDataDict.Type:
inputs = _list_from_dict(self.input_keys, data)
with torch.inference_mode():
outputs = self.model(inputs)
return _list_to_dict(self.output_keys, outputs)
class ListInputOutputStateDictWrapper(ListInputOutputWrapper):
"""Like ``ListInputOutputWrapper``, but also updates the model with state dict entries before each ``forward`` using ``functional_call``."""
def __init__(
self,
model: torch.nn.Module,
input_keys: Sequence[str],
output_keys: Sequence[str],
state_dict_keys: Sequence[str],
):
super().__init__(model, input_keys, output_keys)
self.state_dict_keys = state_dict_keys
def forward(self, *args: torch.Tensor) -> List[torch.Tensor]:
# won't check that `args` is of the correct length
input_dict = _list_to_dict(self.input_keys, args[: len(self.input_keys)])
state_dict = _list_to_dict(self.state_dict_keys, args[len(self.input_keys) :])
# use functional_call to avoid in-place modification
output_dict = functional_call(self.model, state_dict, args=(input_dict,))
return _list_from_dict(self.output_keys, output_dict)
class CompileGraphModel(GraphModel):
"""Wrapper that uses ``torch.compile`` to optimize the wrapped module while allowing it to be trained.
The cache is keyed by input signature (input keys only).
For each input signature, the eager model is run to determine the output keys, and then a compiled model is created for that input/output combination.
The compiled model and output keys are stored together in the cache.
"""
is_compile_graph_model: Final[bool] = True
# ^ to identify `GraphModel` types from `nequip-package`d models (see https://pytorch.org/docs/stable/package.html#torch-package-sharp-edges)
def __init__(
self,
model: GraphModuleMixin,
model_config: Optional[Dict[str, str]] = None,
model_input_fields: Dict[str, Any] = {},
) -> None:
super().__init__(model, model_config, model_input_fields)
# cache for multiple compiled variants based on input key signatures
# cache structure: {input_signature: (compiled_model, output_fields)}
# NOTE: the cache dict is wrapped in a tuple so that it's not registered and saved in the state dict -- this is necessary to enable `GraphModel` to load `CompileGraphModel` state dicts
# see https://discuss.pytorch.org/t/saving-nn-module-to-parent-nn-module-without-registering-paremeters/132082/6
self._compiled_cache = ({},)
# weights and buffers should be done lazily because model modification can happen after instantiation
# such that parameters and buffers may change between class instantiation and the lazy compilation in the `forward`
self.weight_names = None
self.buffer_names = None
def _get_input_signature(self, data: AtomicDataDict.Type) -> tuple:
"""Compute a hashable signature for the input keys.
The unique set of input keys determines a unique set of output keys when run through the model,
so we only need the input keys for the cache lookup signature.
Uses intersection of data keys and GraphModel inputs, which assumes:
- correctness of irreps registration system
- this particular batch contains all necessary inputs for this variant
"""
input_keys = tuple(sorted(data.keys() & self.model_input_fields))
return input_keys
def forward(self, data: AtomicDataDict.Type) -> AtomicDataDict.Type:
# short-circuit if one of the batch dims is 1 (0 would be an error)
# this is related to the 0/1 specialization problem
# see https://docs.google.com/document/d/16VPOa3d-Liikf48teAOmxLc92rgvJdfosIy-yoT38Io/edit?fbclid=IwAR3HNwmmexcitV0pbZm_x1a4ykdXZ9th_eJWK-3hBtVgKnrkmemz6Pm5jRQ&tab=t.0#heading=h.ez923tomjvyk
# we just need something that doesn't have a batch dim of 1 to `make_fx` or else it'll shape specialize
# the models compiled for more batch_size > 1 data cannot be used for batch_size=1 data
# (under specific cases related to the `PerTypeScaleShift` module)
# for now we just make sure to always use the eager model when the data has any batch dims of 1
if (
AtomicDataDict.num_nodes(data) < 2
or AtomicDataDict.num_frames(data) < 2
or AtomicDataDict.num_edges(data) < 2
):
# use parent class's forward
return super().forward(data)
# === get or compile variant for this input signature ===
# compilation happens lazily when we encounter a new combination of input keys
input_signature = self._get_input_signature(data)
cache = self._compiled_cache[0]
if input_signature not in cache:
# get weight names and buffers (only once on first compilation)
if self.weight_names is None:
self.weight_names = [n for n, _ in self.model.named_parameters()]
self.buffer_names = [n for n, _ in self.model.named_buffers()]
# == get input fields for this variant ==
input_fields = list(input_signature)
# == run eager model to determine output fields ==
eager_output = super().forward(data.copy())
output_fields = tuple(sorted(eager_output.keys()))
del eager_output
# == preprocess model and make_fx ==
model_to_trace = ListInputOutputStateDictWrapper(
model=self.model,
input_keys=input_fields,
output_keys=output_fields,
state_dict_keys=self.weight_names + self.buffer_names,
)
weights, buffers = self._get_weights_buffers()
fx_model = nequip_make_fx(
model=model_to_trace,
data=data,
fields=input_fields,
extra_inputs=weights + buffers,
)
del weights, buffers
# == compile exported program ==
# see https://pytorch.org/tutorials/intermediate/torch_export_tutorial.html#running-the-exported-program
# TODO: compile options
compiled_model = torch.compile(
fx_model,
dynamic=True,
fullgraph=False,
)
# store in cache: (compiled_model, output_fields)
cache[input_signature] = (compiled_model, output_fields)
# run original model and compiled model with data to sanity check
def compiled_forward_for_test(data_test):
return self._compiled_forward(
data_test, compiled_model, input_fields, output_fields
)
# only test output fields that are present in data (i.e. labels are present)
test_fields = sorted(set(output_fields) & data.keys())
test_model_output_similarity_by_dtype(
compiled_forward_for_test,
self.model,
{k: data[k] for k in input_fields},
dtype_to_name(self.model_dtype),
fields=test_fields,
error_message=_pt2_compile_error_message,
)
# === run compiled model for this variant ===
compiled_model, output_fields = cache[input_signature]
out_dict = self._compiled_forward(
data, compiled_model, input_signature, output_fields
)
to_return = data.copy()
to_return.update(out_dict)
return to_return
def _compiled_forward(self, data, compiled_model, input_fields, output_fields):
# run compiled model with data
weights, buffers = self._get_weights_buffers()
data_list = _list_from_dict(input_fields, data)
out_list = compiled_model(*(data_list + weights + buffers))
out_dict = _list_to_dict(output_fields, out_list)
return out_dict
def _get_weights_buffers(self):
# get weights and buffers from trainable model
weight_dict = dict(self.model.named_parameters())
weights = [weight_dict[name] for name in self.weight_names]
buffer_dict = dict(self.model.named_buffers())
buffers = [buffer_dict[name] for name in self.buffer_names]
return weights, buffers
|