Download finetune.py from OneScience-Group/Equiformer_v3: direct link, hf CLI and curl.
- Browser
- Download file 14.5 kB
-
https://huggingface.co/OneScience-Group/Equiformer_v3/resolve/main/finetune.py
- Command line
-
hf download hf://OneScience-Group/Equiformer_v3/finetune.py
-
curl -L -o finetune.py https://huggingface.co/OneScience-Group/Equiformer_v3/resolve/main/finetune.py
14.5 kB
| """Compatibility entrypoint for Equiformer V3 checkpoint fine-tuning. | |
| New training configurations are handled by :mod:`train`, which supports | |
| scratch training, checkpoint initialization, and full-state resume. The | |
| loader remains here for compatibility with ``evaluate.py`` and external code. | |
| """ | |
| from __future__ import annotations | |
| import argparse | |
| import json | |
| import os | |
| from dataclasses import dataclass | |
| from pathlib import Path | |
| os.environ.setdefault( | |
| "ONESCIENCE_EQUIFORMER_V3_JD_PATH", | |
| str(Path(__file__).resolve().parent / "weight" / "Jd.pt"), | |
| ) | |
| import torch | |
| import yaml | |
| from torch.nn.parallel import DistributedDataParallel | |
| from torch.utils.data import DataLoader, Subset | |
| from torch.utils.data.distributed import DistributedSampler | |
| from onescience.datapipes.materials.custom_stack import data_list_collater | |
| from onescience.datapipes.materials.custom_stack.storage.ase_datasets import AseDBDataset | |
| from onescience.utils.equiformer_v3 import ( | |
| EquiformerV3CheckpointTransforms, | |
| load_equiformer_v3_checkpoint, | |
| ) | |
| from onescience.utils.uma.normalization.element_references import ( | |
| fit_linear_references, | |
| ) | |
| class DistributedContext: | |
| """Runtime information for a normal Python process or a torchrun worker.""" | |
| rank: int = 0 | |
| world_size: int = 1 | |
| local_rank: int = 0 | |
| def enabled(self) -> bool: | |
| return self.world_size > 1 | |
| def is_main(self) -> bool: | |
| return self.rank == 0 | |
| def _init_distributed(device_name: str, backend: str) -> DistributedContext: | |
| world_size = int(os.environ.get("WORLD_SIZE", "1")) | |
| if world_size == 1: | |
| if device_name.startswith("cuda") and torch.cuda.is_available(): | |
| torch.cuda.set_device(0) | |
| return DistributedContext() | |
| if not torch.distributed.is_available(): | |
| raise RuntimeError("torch.distributed is required for multi-device fine-tuning.") | |
| rank = int(os.environ["RANK"]) | |
| local_rank = int(os.environ.get("LOCAL_RANK", rank)) | |
| if device_name.startswith("cuda"): | |
| if not torch.cuda.is_available(): | |
| raise RuntimeError( | |
| "torchrun requested multiple CUDA/DCU devices, but CUDA is unavailable." | |
| ) | |
| torch.cuda.set_device(local_rank) | |
| torch.distributed.init_process_group(backend=backend, rank=rank, world_size=world_size) | |
| return DistributedContext(rank=rank, world_size=world_size, local_rank=local_rank) | |
| def _close_distributed(context: DistributedContext) -> None: | |
| if context.enabled and torch.distributed.is_initialized(): | |
| torch.distributed.barrier() | |
| torch.distributed.destroy_process_group() | |
| def _loader( | |
| path: str | list[str], | |
| batch_size: int, | |
| workers: int, | |
| max_samples: int | None = None, | |
| context: DistributedContext | None = None, | |
| train: bool = False, | |
| seed: int = 0, | |
| ) -> DataLoader: | |
| dataset = AseDBDataset( | |
| { | |
| "src": path, | |
| "a2g_args": { | |
| "r_edges": False, | |
| "r_energy": True, | |
| "r_forces": True, | |
| "r_stress": True, | |
| }, | |
| } | |
| ) | |
| if max_samples is not None: | |
| sample_count = min(max_samples, len(dataset)) | |
| generator = torch.Generator().manual_seed(seed) | |
| indices = torch.randperm(len(dataset), generator=generator)[:sample_count].tolist() | |
| dataset = Subset(dataset, indices) | |
| context = context or DistributedContext() | |
| sampler = None | |
| if context.enabled: | |
| sampler = DistributedSampler( | |
| dataset, | |
| num_replicas=context.world_size, | |
| rank=context.rank, | |
| shuffle=train, | |
| drop_last=False, | |
| ) | |
| return DataLoader( | |
| dataset, | |
| batch_size=batch_size, | |
| shuffle=sampler is None and train, | |
| sampler=sampler, | |
| num_workers=workers, | |
| collate_fn=lambda items: data_list_collater(items, otf_graph=True), | |
| ) | |
| def _loss( | |
| pred, | |
| batch, | |
| energy_weight: float, | |
| force_weight: float, | |
| stress_weight: float, | |
| transforms: EquiformerV3CheckpointTransforms, | |
| ): | |
| losses = {} | |
| if energy_weight: | |
| energy_target = transforms.normalize_target( | |
| "energy", batch.energy, pred["energy"], batch | |
| ) | |
| energy_error = pred["energy"] - energy_target | |
| natoms_shape = (-1,) + (1,) * (energy_error.ndim - 1) | |
| natoms = batch.natoms.to(energy_error).reshape(natoms_shape) | |
| losses["energy"] = (energy_error / natoms).square().mean() | |
| if force_weight: | |
| force_target = transforms.normalize_target( | |
| "forces", batch.forces, pred["forces"], batch | |
| ) | |
| losses["forces"] = (pred["forces"] - force_target).square().mean() | |
| if stress_weight: | |
| if not hasattr(batch, "stress"): | |
| raise ValueError("stress_weight is nonzero, but the batch has no stress labels") | |
| if "stress" not in pred: | |
| raise ValueError("stress_weight is nonzero, but the model returned no stress") | |
| stress_target = transforms.normalize_target( | |
| "stress", batch.stress, pred["stress"], batch | |
| ) | |
| losses["stress"] = (pred["stress"] - stress_target).square().mean() | |
| total = energy_weight * losses.get("energy", 0.0) | |
| total = total + force_weight * losses.get("forces", 0.0) | |
| total = total + stress_weight * losses.get("stress", 0.0) | |
| return total, {key: float(value.detach()) for key, value in losses.items()} | |
| def _run_epoch( | |
| model, | |
| loader, | |
| device, | |
| optimizer, | |
| weights, | |
| transforms: EquiformerV3CheckpointTransforms, | |
| context: DistributedContext, | |
| ): | |
| training = optimizer is not None | |
| model.train(training) | |
| total = 0.0 | |
| batches = 0 | |
| metric_names = tuple( | |
| name | |
| for name, weight in zip(("energy", "forces", "stress"), weights) | |
| if weight | |
| ) | |
| metrics = {name: 0.0 for name in metric_names} | |
| for batch in loader: | |
| batch = batch.to(device) | |
| if training: | |
| optimizer.zero_grad(set_to_none=True) | |
| prediction = model(batch) | |
| loss, batch_metrics = _loss(prediction, batch, *weights, transforms) | |
| if training: | |
| loss.backward() | |
| optimizer.step() | |
| total += float(loss.detach()) | |
| batches += 1 | |
| for key, value in batch_metrics.items(): | |
| metrics[key] = metrics.get(key, 0.0) + value | |
| if batches == 0: | |
| raise RuntimeError("The dataset contains no samples.") | |
| values = torch.tensor( | |
| [total, *metrics.values(), float(batches)], | |
| dtype=torch.float64, | |
| device=device, | |
| ) | |
| if context.enabled: | |
| torch.distributed.all_reduce(values, op=torch.distributed.ReduceOp.SUM) | |
| global_batches = values[-1].item() | |
| return { | |
| "loss": values[0].item() / global_batches, | |
| **{ | |
| key: values[index].item() / global_batches | |
| for index, key in enumerate(metrics, start=1) | |
| }, | |
| } | |
| def main() -> None: | |
| parser = argparse.ArgumentParser(description=__doc__) | |
| parser.add_argument("--config", help="YAML configuration path") | |
| parser.add_argument("--checkpoint") | |
| parser.add_argument("--train", help="ASE DB or ASE-LMDB training path") | |
| parser.add_argument("--val", help="ASE DB or ASE-LMDB validation path") | |
| parser.add_argument("--output") | |
| parser.add_argument("--device") | |
| parser.add_argument("--epochs", type=int) | |
| parser.add_argument("--batch-size", type=int) | |
| parser.add_argument("--workers", type=int) | |
| parser.add_argument("--lr", type=float) | |
| parser.add_argument("--energy-weight", type=float) | |
| parser.add_argument("--force-weight", type=float) | |
| parser.add_argument("--stress-weight", type=float) | |
| parser.add_argument("--max-train-samples", type=int) | |
| parser.add_argument("--max-val-samples", type=int) | |
| parser.add_argument("--backend", help="torch.distributed backend for torchrun") | |
| parser.add_argument("--seed", type=int) | |
| parser.add_argument( | |
| "--fit-element-references", | |
| action=argparse.BooleanOptionalAction, | |
| default=None, | |
| help="fit energy element references on the training data", | |
| ) | |
| args = parser.parse_args() | |
| if not args.config: | |
| parser.error("--config is required; use a YAML file from demo/configs") | |
| config_path = args.config | |
| with Path(config_path).expanduser().open() as handle: | |
| config = yaml.safe_load(handle) or {} | |
| for key, value in config.items(): | |
| if getattr(args, key.replace("-", "_"), None) is None: | |
| setattr(args, key.replace("-", "_"), value) | |
| for key in ("checkpoint", "train", "val", "output"): | |
| value = getattr(args, key) | |
| if value is not None: | |
| if isinstance(value, list): | |
| value = [ | |
| os.path.expandvars(os.path.expanduser(str(item))) | |
| for item in value | |
| ] | |
| else: | |
| value = os.path.expandvars(os.path.expanduser(str(value))) | |
| setattr(args, key, value) | |
| required = ("checkpoint", "train", "val", "output") | |
| missing = [key for key in required if not getattr(args, key)] | |
| if missing: | |
| parser.error("missing required config fields: " + ", ".join(missing)) | |
| args.backend = args.backend or "nccl" | |
| args.seed = 0 if args.seed is None else args.seed | |
| args.fit_element_references = bool(args.fit_element_references) | |
| if not any((args.energy_weight, args.force_weight, args.stress_weight)): | |
| parser.error( | |
| "at least one of energy_weight, force_weight, or stress_weight " | |
| "must be nonzero" | |
| ) | |
| if args.device.startswith("cuda") and not torch.cuda.is_available(): | |
| raise RuntimeError("CUDA/DCU was requested but torch.cuda.is_available() is false.") | |
| context = _init_distributed(args.device, args.backend) | |
| try: | |
| if args.device.startswith("cuda"): | |
| device = torch.device(f"cuda:{context.local_rank}") | |
| else: | |
| device = torch.device(args.device) | |
| torch.manual_seed(args.seed + context.rank) | |
| model = load_equiformer_v3_checkpoint(args.checkpoint).to(device) | |
| transforms = EquiformerV3CheckpointTransforms.from_checkpoint( | |
| args.checkpoint | |
| ) | |
| if args.fit_element_references: | |
| reference_dataset = _loader( | |
| args.train, args.batch_size, args.workers | |
| ).dataset | |
| fitted_references = fit_linear_references( | |
| targets=["energy"], | |
| dataset=reference_dataset, | |
| batch_size=args.batch_size, | |
| num_workers=args.workers, | |
| log_metrics=False, | |
| shuffle=False, | |
| ) | |
| transforms.elementrefs["energy"] = fitted_references["energy"] | |
| if context.is_main: | |
| print("fitted energy element references from training data", flush=True) | |
| transforms = transforms.to(device) | |
| if context.enabled: | |
| model = DistributedDataParallel( | |
| model, | |
| device_ids=[context.local_rank] if device.type == "cuda" else None, | |
| output_device=context.local_rank if device.type == "cuda" else None, | |
| ) | |
| optimizer = torch.optim.AdamW(model.parameters(), lr=args.lr) | |
| train_loader = _loader( | |
| args.train, | |
| args.batch_size, | |
| args.workers, | |
| args.max_train_samples, | |
| context=context, | |
| train=True, | |
| seed=args.seed, | |
| ) | |
| val_loader = _loader( | |
| args.val, | |
| args.batch_size, | |
| args.workers, | |
| args.max_val_samples, | |
| context=context, | |
| train=False, | |
| seed=args.seed + 1, | |
| ) | |
| weights = (args.energy_weight, args.force_weight, args.stress_weight) | |
| history = [] | |
| for epoch in range(args.epochs): | |
| if isinstance(train_loader.sampler, DistributedSampler): | |
| train_loader.sampler.set_epoch(epoch) | |
| train_metrics = _run_epoch( | |
| model, train_loader, device, optimizer, weights, transforms, context | |
| ) | |
| # Force/stress outputs are gradients of the energy, so validation also | |
| # needs autograd even though model parameters are not updated. | |
| val_metrics = _run_epoch( | |
| model, val_loader, device, None, weights, transforms, context | |
| ) | |
| record = {"epoch": epoch, "train": train_metrics, "val": val_metrics} | |
| if context.is_main: | |
| history.append(record) | |
| print(json.dumps(record, sort_keys=True), flush=True) | |
| if context.is_main: | |
| output = Path(args.output) | |
| output.parent.mkdir(parents=True, exist_ok=True) | |
| source = torch.load(args.checkpoint, map_location="cpu", weights_only=False) | |
| base_model = model.module if context.enabled else model | |
| checkpoint = { | |
| "config": source["config"], | |
| "normalizers": source.get("normalizers", {}), | |
| "state_dict": { | |
| key: value.detach().cpu() | |
| for key, value in base_model.state_dict().items() | |
| }, | |
| "metadata": dict(source.get("metadata", {})), | |
| } | |
| checkpoint["elementrefs"] = { | |
| name: { | |
| key: value.detach().cpu() | |
| for key, value in elementref.state_dict().items() | |
| } | |
| for name, elementref in transforms.elementrefs.items() | |
| } | |
| checkpoint["metadata"].update( | |
| { | |
| "onescience_equiformer_v3_history": history, | |
| "source_checkpoint": args.checkpoint, | |
| "world_size": context.world_size, | |
| "loss_space": "checkpoint_normalized", | |
| "element_references": ( | |
| "fitted_from_training_data" | |
| if args.fit_element_references | |
| else "source_checkpoint" | |
| ), | |
| } | |
| ) | |
| del source | |
| torch.save(checkpoint, output) | |
| print(f"saved checkpoint: {output}", flush=True) | |
| finally: | |
| _close_distributed(context) | |
| if __name__ == "__main__": | |
| from train import main as training_main | |
| training_main() | |