lhallee commited on
Commit
a6afadf
·
verified ·
1 Parent(s): 3e05dca

Give this artifact's bridge its own config and model subclasses (FastPLMs#50)

Browse files

Artifacts share one FastPLMs runtime per process. Re-exported shared classes let one artifact's model class answer for another with the same config class, for example ESMFold2-300 built as ESMFold2Model after ESMFold2-Fast loaded. Only modeling_fastplms.py changes; the runtime bundle and weights are untouched. https://github.com/Synthyra/FastPLMs/issues/50

Files changed (1) hide show
  1. modeling_fastplms.py +27 -9
modeling_fastplms.py CHANGED
@@ -180,12 +180,30 @@ def _install_runtime():
180
  return package
181
 
182
  _install_runtime()
183
- _module_182 = _import_without_bytecode("fastplms.models.esmfold.modeling_fast_esmfold")
184
- FastEsmFoldConfig = _module_182.FastEsmFoldConfig
185
- FastEsmFoldConfig.__module__ = __name__
186
- FastEsmForProteinFolding = _module_182.FastEsmForProteinFolding
187
- FastEsmForProteinFolding.__module__ = __name__
188
- FastEsmForSequenceClassification = _module_182.FastEsmForSequenceClassification
189
- FastEsmForSequenceClassification.__module__ = __name__
190
- FastEsmForTokenClassification = _module_182.FastEsmForTokenClassification
191
- FastEsmForTokenClassification.__module__ = __name__
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
180
  return package
181
 
182
  _install_runtime()
183
+
184
+ def _artifact_class(base, config_class=None):
185
+ """Subclass a shared runtime class so this artifact registers its own classes.
186
+
187
+ The base constructor is kept explicitly because Transformers turns every config
188
+ subclass into a dataclass, which would otherwise generate a new constructor.
189
+ """
190
+ namespace = {
191
+ "__module__": __name__,
192
+ "__qualname__": base.__name__,
193
+ "__doc__": base.__doc__,
194
+ "__init__": base.__init__,
195
+ }
196
+ base_config_class = getattr(base, "config_class", None)
197
+ if (
198
+ config_class is not None
199
+ and isinstance(base_config_class, type)
200
+ and issubclass(config_class, base_config_class)
201
+ ):
202
+ namespace["config_class"] = config_class
203
+ return type(base.__name__, (base,), namespace)
204
+
205
+ _module_204 = _import_without_bytecode("fastplms.models.esmfold.modeling_fast_esmfold")
206
+ FastEsmFoldConfig = _artifact_class(_module_204.FastEsmFoldConfig)
207
+ FastEsmForProteinFolding = _artifact_class(_module_204.FastEsmForProteinFolding, config_class=FastEsmFoldConfig)
208
+ FastEsmForSequenceClassification = _artifact_class(_module_204.FastEsmForSequenceClassification, config_class=FastEsmFoldConfig)
209
+ FastEsmForTokenClassification = _artifact_class(_module_204.FastEsmForTokenClassification, config_class=FastEsmFoldConfig)