Skip to article frontmatterSkip to article content
Site not loading correctly?

This may be due to an incorrect BASE_URL configuration. See the MyST Documentation for reference.

Fine-tuning

This repo provides a number of scripts to quickly fine-tune a model using a custom ASE LMDB dataset. These scripts are merely for convenience and fine-tuning uses the exact same tooling and infrastructure as our standard training (see Training section). Training in the fairchem repo uses the fairchem CLI tool and configs are in Hydra yaml format.

Generating Training/Fine-tuning Datasets

First we need to generate a dataset in the aselmdb format for fine-tuning.

First you should checkout the fairchem repo and install it to access the scripts

Run this script to create the aselmdbs as well as a set of templated yamls for finetuning, we will use a few dummy structures for demonstration purposes

Warp 1.18.0 initialized:
   CUDA Toolkit 13.4, Driver 13.1
   Devices:
     "cpu"      : "x86_64"
     "cuda:0"   : "Tesla T4" (16 GiB, sm_75, mempool enabled)
   Kernel cache:
     /home/runner/.cache/warp/1.18.0







0it [00:00, ?it/s]

  0%|                                                    | 0/1 [00:00<?, ?it/s]

0it [00:00, ?it/s]











0it [00:00, ?it/s]




0it [00:00, ?it/s]
0it [00:00, ?it/s]



0it [00:00, ?it/s]
100%|████████████████████████████████████████████| 1/1 [00:00<00:00, 51.41it/s]
100%|████████████████████████████████████████████| 1/1 [00:00<00:00, 41.27it/s]
Computing normalizer values.: 100%|█████████████| 2/2 [00:00<00:00, 343.57it/s]


0it [00:00, ?it/s]
0it [00:00, ?it/s]                                       | 0/1 [00:00<?, ?it/s]





0it [00:00, ?it/s]







  0%|                                                    | 0/1 [00:00<?, ?it/s]



0it [00:00, ?it/s]
0it [00:00, ?it/s]



0it [00:00, ?it/s]






0it [00:00, ?it/s]
100%|████████████████████████████████████████████| 1/1 [00:00<00:00, 29.31it/s]
100%|████████████████████████████████████████████| 1/1 [00:00<00:00, 25.66it/s]
INFO:root:Generated dataset and data config yaml in /tmp/bulk
INFO:root:To run finetuning, run fairchem -c /tmp/bulk/uma_sm_finetune_template.yaml

This will generate a folder of LMDBs and a uma_sm_finetune_template.yaml that you can run directly with the fairchem CLI to start training.

Model Fine-tuning (Default Settings)

The previous step should have generated some YAML files to get you started on fine-tuning. You can simply run this with the fairchem CLI. The default is configured to run locally on 1 GPU.

Advanced Configuration

The scripts provide a simple way to get started on fine-tuning, but likely for your own use cases you will need to modify the parameters. The configuration uses Hydra-style YAMLs.

The basic YAML configuration looks like the following:

job:
  device_type: CUDA
  scheduler:
    mode: LOCAL
    ranks_per_node: 1
    num_nodes: 1
  debug: True
  run_dir: /tmp/uma_finetune_runs/
  run_name: uma_finetune
  logger:
    _target_: fairchem.core.common.logger.WandBSingletonLogger.init_wandb
    _partial_: true
    entity: example
    project: uma_finetune


base_model_name: uma-s-1p2p1
max_neighbors: 300
epochs: 1
steps: null
batch_size: 2
lr: 4e-4

train_dataloader ...
eval_dataloader ...
runner ...

Logging and Artifacts

For logging and checkpoints, all artifacts are stored in the location specified in job.run_dir. The visual logger we support is Weights and Biases.

Distributed Training

We support multi-GPU distributed training without additional infrastructure and multi-node distributed training on SLURM only.

Multi-GPU locally: Simply set job.scheduler.ranks_per_node=N where N is the number of GPUs you want to train on.

Multi-node on SLURM: Change job.scheduler.mode=SLURM and set both job.scheduler.ranks_per_node and job.scheduler.num_nodes to the desired values.

Resuming Runs

To resume from a checkpoint in the middle of a run, find the checkpoint folder at the step you want and use the same fairchem command:

Running Inference on the Fine-tuned Model

Inference is run in the same way as the UMA models, except you need to load the checkpoint from a local path.