-
Notifications
You must be signed in to change notification settings - Fork 654
feat: add the Uni-Mol v1 backbone and its self-supervised pretraining #6019
New issue
Have a question about this project? Sign up for a free GitHub account to open an issue and contact its maintainers and the community.
By clicking “Sign up for GitHub”, you agree to our terms of service and privacy statement. We’ll occasionally send you account related emails.
Already on GitHub? Sign in to your account
Open
iProzd
wants to merge
34
commits into
deepmodeling:master
Choose a base branch
from
iProzd:0912_unimol_core
base: master
Could not load branches
Branch not found: {{ refName }}
Loading
Could not load tags
Nothing to show
Loading
Are you sure you want to change the base?
Some commits from the old base branch may be removed from the timeline,
and old review comments may become outdated.
Open
Changes from all commits
Commits
Show all changes
34 commits
Select commit
Hold shift + click to select a range
b465834
feat(dpmodel): add Uni-Mol v1 encoder blocks and an exact GELU
iProzd a6a196f
feat(dpmodel): add the Uni-Mol v1 data-side transforms
iProzd 4591035
feat: register the exact GELU in every backend activation table
iProzd 0fecbd6
feat(dpmodel): add the Uni-Mol v1 pretraining heads and loss
iProzd c6b85ef
feat(dpmodel): add the unimol descriptor
iProzd df72506
feat: import released Uni-Mol v1 checkpoints
iProzd bc53505
feat(dpmodel): add the unimol_pretrain fitting
iProzd d488e55
feat: register the unimol descriptor, fitting and loss in argcheck
iProzd dbd8ab3
feat(dpmodel): register the unimol_pretrain model
iProzd 24f46d4
feat(pt-expt): wire the unimol model into the PyTorch-Exportable backend
iProzd e6fb4d7
test: parity tests for the Uni-Mol v1 port
iProzd b77b533
feat(dpmodel): apply Uni-Mol's dropout during training
iProzd d98f85a
feat: convert Uni-Mol pretraining data into a deepmd dataset
iProzd 8042bec
feat: make Uni-Mol pretraining reachable from a training run
iProzd e69c6af
fix(dpmodel): place every constructed array on the input's device
iProzd 04d917d
fix: refuse a periodic cell at the model boundary
iProzd 63251e2
feat: let an objective declare the data transform it needs
iProzd 18649c0
docs: describe how a Uni-Mol run is configured, and validate the example
iProzd 5d4043d
fix: make `dp --pt-expt train` work for the Uni-Mol objective
iProzd 9aa9930
test: train from a configuration, the way the feature is used
iProzd f333398
fix: say what the locality guard actually rules out
iProzd d0c192d
feat: expose the Adam epsilon
iProzd 6ae1a1e
fix(dpmodel): train the token embedding and the distance basis
iProzd 7fbcd2e
fix: corruption varies per epoch, and cropping moves to the converter
iProzd fe88321
fix: keep the embedding and the distance basis in the gradient
iProzd 8cc8e69
fix: act on the review of this pull request
iProzd 7f486c1
fix: keep the objective finite when nothing is corrupted
iProzd ee98665
fix: let the corruption survive a decoder worker
iProzd e82b575
fix: honour a precision the backbone does not share
iProzd aacaa7e
feat: match the published LMDB layout, and train in single precision
iProzd 3b48dd1
fix: act on the review of this pull request
iProzd 92ef3ef
fix: act on the second review of this pull request
iProzd 4f43a33
test: cover the two review paths that had no test
iProzd a309c88
fix: give `xp_erf` a TensorFlow branch
iProzd File filter
Filter by extension
Conversations
Failed to load comments.
Loading
Jump to
Jump to file
Failed to load files.
Loading
Diff view
Diff view
There are no files selected for viewing
This file contains hidden or bidirectional Unicode text that may be interpreted or compiled differently than what appears below. To review, open the file in an editor that reveals hidden Unicode characters.
Learn more about bidirectional Unicode characters
| Original file line number | Diff line number | Diff line change |
|---|---|---|
|
|
@@ -53,6 +53,7 @@ | |
| "tanh", | ||
| "gelu", | ||
| "gelu_tf", | ||
| "gelu_erf", | ||
| "silu", | ||
| "silut", | ||
| "none", | ||
|
|
||
This file contains hidden or bidirectional Unicode text that may be interpreted or compiled differently than what appears below. To review, open the file in an editor that reveals hidden Unicode characters.
Learn more about bidirectional Unicode characters
This file contains hidden or bidirectional Unicode text that may be interpreted or compiled differently than what appears below. To review, open the file in an editor that reveals hidden Unicode characters.
Learn more about bidirectional Unicode characters
This file contains hidden or bidirectional Unicode text that may be interpreted or compiled differently than what appears below. To review, open the file in an editor that reveals hidden Unicode characters.
Learn more about bidirectional Unicode characters
| Original file line number | Diff line number | Diff line change |
|---|---|---|
| @@ -0,0 +1,106 @@ | ||
| # SPDX-License-Identifier: LGPL-3.0-or-later | ||
| """Atomic model for Uni-Mol v1 self-supervised pretraining.""" | ||
|
|
||
| from typing import ( | ||
| Any, | ||
| ) | ||
|
|
||
| from deepmd.dpmodel.array_api import ( | ||
| Array, | ||
| ) | ||
| from deepmd.dpmodel.descriptor.unimol import ( | ||
| DescrptUniMol, | ||
| ) | ||
| from deepmd.dpmodel.fitting.unimol_pretrain import ( | ||
| UniMolPretrainFitting, | ||
| ) | ||
|
|
||
| from .dp_atomic_model import ( | ||
| DPAtomicModel, | ||
| ) | ||
|
|
||
|
|
||
| class DPUniMolAtomicModel(DPAtomicModel): | ||
| r"""Uni-Mol pretraining, wired at token resolution. | ||
|
|
||
| The standard path hands a descriptor's five-tuple to a fitting, which | ||
| cannot carry the two virtual tokens, the pair channel or the norm | ||
| regularisers that these heads read. This model therefore overrides one | ||
| method to route the backbone's token-resolution output straight into the | ||
| heads. Nothing else about the atomic model changes. | ||
| """ | ||
|
|
||
| def __init__( | ||
| self, descriptor: Any, fitting: Any, type_map: list[str], **kwargs: Any | ||
| ) -> None: | ||
| if not isinstance(descriptor, DescrptUniMol): | ||
| raise TypeError( | ||
| "DPUniMolAtomicModel needs the unimol descriptor, which is the only " | ||
| "one producing a Uni-Mol token sequence" | ||
| ) | ||
| if not isinstance(fitting, UniMolPretrainFitting): | ||
| raise TypeError("DPUniMolAtomicModel needs the unimol_pretrain fitting") | ||
| # The objective compares against distances whose virtual tokens sit at | ||
| # the origin, which is where upstream puts them and, once the transform | ||
| # has centred a frame, where the clean centroid is. Under "centroid" the | ||
| # descriptor instead places them at the centroid of the coordinates it | ||
| # is handed -- the corrupted ones -- so the two virtual columns of every | ||
| # corrupted row would be regressed against a label for a different | ||
| # position, by about the size of the noise. Centring costs nothing here | ||
| # because the transform always centres, so the only effect would be that | ||
| # silent mismatch. | ||
| if getattr(descriptor, "virtual_token_position", "origin") != "origin": | ||
| raise ValueError( | ||
| "unimol pretraining needs virtual_token_position='origin': the " | ||
| "distance target places the virtual tokens at the origin, and " | ||
| f"this descriptor places them at the " | ||
| f"{descriptor.virtual_token_position}, so the two virtual " | ||
| "columns would train against the wrong label. The corruption " | ||
| "centres every frame, so 'origin' is the centroid anyway" | ||
| ) | ||
| super().__init__(descriptor, fitting, type_map, **kwargs) | ||
|
|
||
| def forward_atomic( | ||
| self, | ||
| extended_coord: Array, | ||
| extended_atype: Array, | ||
| nlist: Array, | ||
| mapping: Array | None = None, | ||
| fparam: Array | None = None, | ||
| aparam: Array | None = None, | ||
| comm_dict: dict | None = None, | ||
| charge_spin: Array | None = None, | ||
| ) -> dict[str, Array]: | ||
| """Run the backbone and its heads at token resolution. | ||
|
|
||
| Parameters | ||
| ---------- | ||
| extended_coord | ||
| nf x (nall x 3) coordinates; the descriptor rejects any frame that | ||
| carries periodic images. | ||
| extended_atype | ||
| nf x nall element types, already clamped to be non-negative. | ||
| nlist | ||
| nf x nloc x nnei neighbour list, which is how real atoms are told | ||
| apart from padding. | ||
| mapping, fparam, aparam, comm_dict, charge_spin | ||
| Unused by this model. | ||
|
|
||
| Returns | ||
| ------- | ||
| dict | ||
| The three head outputs plus the two norm regularisers. | ||
| """ | ||
| del mapping, fparam, aparam, comm_dict, charge_spin | ||
| backbone = self.descriptor.forward_tokens(extended_coord, extended_atype, nlist) | ||
| return self.fitting_net.call_tokens(backbone) | ||
|
|
||
| def apply_out_stat(self, ret: dict[str, Array], atype: Array) -> dict[str, Array]: | ||
| """Return the head outputs untouched. | ||
|
|
||
| Self-supervised targets carry no per-element bias to add back: the | ||
| element head predicts a distribution, and the coordinate and distance | ||
| heads predict geometry the data already fixes. | ||
| """ | ||
| del atype | ||
| return ret |
This file contains hidden or bidirectional Unicode text that may be interpreted or compiled differently than what appears below. To review, open the file in an editor that reveals hidden Unicode characters.
Learn more about bidirectional Unicode characters
Oops, something went wrong.
Oops, something went wrong.
Add this suggestion to a batch that can be applied as a single commit.
This suggestion is invalid because no changes were made to the code.
Suggestions cannot be applied while the pull request is closed.
Suggestions cannot be applied while viewing a subset of changes.
Only one suggestion per line can be applied in a batch.
Add this suggestion to a batch that can be applied as a single commit.
Applying suggestions on deleted lines is not supported.
You must change the existing code in this line in order to create a valid suggestion.
Outdated suggestions cannot be applied.
This suggestion has been applied or marked resolved.
Suggestions cannot be applied from pending reviews.
Suggestions cannot be applied on multi-line comments.
Suggestions cannot be applied while the pull request is queued to merge.
Suggestion cannot be applied right now. Please check back later.
Uh oh!
There was an error while loading. Please reload this page.