Skip to content

Commit d71ae5e

Browse files
authored
Merge pull request #14 from zhaiwenxi/merge-preserve-both
Merge dpa-adapt updates
2 parents 510392a + 53bbeec commit d71ae5e

8 files changed

Lines changed: 308 additions & 199 deletions

File tree

doc/dpa_adapt/README.md

Lines changed: 21 additions & 21 deletions
Original file line numberDiff line numberDiff line change
@@ -1,6 +1,6 @@
1-
# ADAPT: Atomistic DPA Adaptation for Property Tasks
1+
# DPA-ADAPT: Atomistic DPA Adaptation for Property Tasks
22

3-
**ADAPT** is a scikit-learn-style Python package for fine-tuning pre-trained DPA models on your own materials or molecular property dataset. No DeePMD-kit JSON configs or `dp train` pipelines to write.
3+
**DPA-ADAPT** (`dpa-adapt`, Python import `dpa_adapt`) is a toolkit for adapting pretrained DPA models to downstream atomistic property prediction tasks. The main CLI is `dpa-adapt`; the optional short alias is `dpaad`. No DeePMD-kit JSON configs or `dp train` pipelines to write.
44

55
## Installation
66

@@ -194,35 +194,35 @@ X = extract_descriptors(
194194

195195
## CLI
196196

197-
| Command | Description |
198-
| --------------------------- | -------------------------------------------------------------------- |
199-
| `dpaad fit` | Fine-tune (`--strategy frozen_sklearn\|linear_probe\|finetune\|mft`) |
200-
| `dpaad predict` | Predict with a frozen `.pth` bundle |
201-
| `dpaad evaluate` | Evaluate against stored labels |
202-
| `dpaad extract-descriptors` | Extract pooled DPA descriptors to `.npy` |
203-
| `dpaad cv` | Cross-validate |
204-
| `dpaad data convert` | Convert structure / CSV / formula → `deepmd/npy` |
205-
| `dpaad data validate` | Sanity-check `deepmd/npy` directories |
206-
| `dpaad data attach-labels` | Inject `.npy` label arrays |
197+
| Command | Description |
198+
|---------|-------------|
199+
| `dpa-adapt fit` / `dpaad fit` | Fine-tune (`--strategy frozen_sklearn\|linear_probe\|finetune\|mft`) |
200+
| `dpa-adapt predict` / `dpaad predict` | Predict with a frozen `.pth` bundle |
201+
| `dpa-adapt evaluate` / `dpaad evaluate` | Evaluate against stored labels |
202+
| `dpa-adapt extract-descriptors` / `dpaad extract-descriptors` | Extract pooled DPA descriptors to `.npy` |
203+
| `dpa-adapt cv` / `dpaad cv` | Cross-validate |
204+
| `dpa-adapt data convert` / `dpaad data convert` | Convert structure / CSV / formula → `deepmd/npy` |
205+
| `dpa-adapt data validate` / `dpaad data validate` | Sanity-check `deepmd/npy` directories |
206+
| `dpa-adapt data attach-labels` / `dpaad data attach-labels` | Inject `.npy` label arrays |
207207

208208
```bash
209209
# Data conversion
210-
dpaad data convert --input POSCAR --output ./npy
210+
dpa-adapt data convert --input POSCAR --output ./npy
211211
dpaad data convert --input data.csv --output ./npy --property-name homo
212-
dpaad data convert --input comps.csv --output ./npy \
213-
--fmt formula --poscar template.POSCAR --sets 3
212+
dpa-adapt data convert --input comps.csv --output ./npy \
213+
--fmt formula --poscar template.POSCAR --sets 3
214214

215215
# Fine-tune
216-
dpaad fit --train-data ./npy/train --pretrained DPA-3.1-3M \
217-
--strategy frozen_sklearn --predictor rf --target-key homo --output model.pth
216+
dpa-adapt fit --train-data ./npy/train --pretrained DPA-3.1-3M \
217+
--strategy frozen_sklearn --predictor rf --target-key homo --output model.pth
218218

219219
# MFT
220220
dpaad fit --train-data /data/qm9 --aux-data /data/spice2 \
221-
--pretrained /path/to/DPA-3.1-3M.pt --strategy mft --target-key homo
221+
--pretrained /path/to/DPA-3.1-3M.pt --strategy mft --target-key homo
222222

223223
# Predict / evaluate
224-
dpaad predict --model model.pth --data ./npy/test
225-
dpaad evaluate --model model.pth --data ./npy/test
224+
dpa-adapt predict --model model.pth --data ./npy/test
225+
dpa-adapt evaluate --model model.pth --data ./npy/test
226226
```
227227

228-
`dpaad --help` does not load torch — all heavy imports are lazy.
228+
`dpa-adapt --help` and `dpaad --help` do not load torch — all heavy imports are lazy.

doc/dpa_adapt/input_formats.md

Lines changed: 122 additions & 99 deletions
Large diffs are not rendered by default.

doc/index.rst

Lines changed: 1 addition & 0 deletions
Original file line numberDiff line numberDiff line change
@@ -44,6 +44,7 @@ DeePMD-kit is a package written in Python/C++, designed to minimize the effort r
4444
test/index
4545
inference/index
4646
dpa_adapt/README
47+
dpa_adapt/input_formats
4748
cli
4849
third-party/index
4950
agent-skills

dpa_adapt/cli.py

Lines changed: 20 additions & 22 deletions
Original file line numberDiff line numberDiff line change
@@ -280,7 +280,9 @@ def _cmd_data_convert(args: argparse.Namespace) -> int:
280280
train_ratio=args.train_ratio,
281281
smiles_col=args.smiles_col,
282282
mol_dir=args.mol_dir,
283-
seed=args.seed,
283+
mol_template=args.mol_template,
284+
split_seed=args.split_seed,
285+
conformer_seed=args.conformer_seed,
284286
poscar=args.poscar,
285287
formula_col=args.formula_col,
286288
base_element=args.base_element,
@@ -631,28 +633,24 @@ def get_parser() -> argparse.ArgumentParser:
631633
parser_data_convert.add_argument("--property-col", default="Property")
632634
parser_data_convert.add_argument("--smiles-col", default="SMILES")
633635
parser_data_convert.add_argument("--mol-dir", default=None)
636+
parser_data_convert.add_argument("--mol-template", default="id{row}.mol",
637+
help="Filename template under --mol-dir; use {row} for the CSV row index.")
634638
parser_data_convert.add_argument("--train-ratio", type=float, default=0.9)
635-
parser_data_convert.add_argument("--seed", type=int, default=42)
636-
parser_data_convert.add_argument(
637-
"--poscar", default=None, help="Template POSCAR for fmt=formula."
638-
)
639-
parser_data_convert.add_argument(
640-
"--base-element",
641-
default=None,
642-
help="Sublattice element to substitute "
643-
"(fmt=formula). Auto-inferred if omitted.",
644-
)
645-
parser_data_convert.add_argument(
646-
"--formula-col",
647-
default=0,
648-
help="Column index or name for the formula (fmt=formula, default: 0).",
649-
)
650-
parser_data_convert.add_argument(
651-
"--sets",
652-
type=int,
653-
default=1,
654-
help="Random structures per formula (fmt=formula, default: 1).",
655-
)
639+
parser_data_convert.add_argument("--split-seed", type=int, default=None,
640+
help="Random seed for train/valid split (SMILES input).")
641+
parser_data_convert.add_argument("--conformer-seed", type=int, default=None,
642+
help="Random seed for RDKit conformer generation (SMILES input).")
643+
parser_data_convert.add_argument("--poscar", default=None,
644+
help="Template POSCAR for fmt=formula.")
645+
parser_data_convert.add_argument("--base-element", default=None,
646+
help="Sublattice element to substitute "
647+
"(fmt=formula). Auto-inferred if omitted.")
648+
parser_data_convert.add_argument("--formula-col", default="formula",
649+
help="Column index or name for the formula "
650+
"(fmt=formula, default: formula).")
651+
parser_data_convert.add_argument("--sets", type=int, default=1,
652+
help="Random structures per formula "
653+
"(fmt=formula, default: 1).")
656654
parser_data_convert.add_argument("--overwrite", action="store_true")
657655

658656
# data validate

dpa_adapt/data/convert.py

Lines changed: 7 additions & 2 deletions
Original file line numberDiff line numberDiff line change
@@ -104,9 +104,12 @@ def auto_convert(
104104
train_ratio: float = 0.9,
105105
smiles_col: str = "SMILES",
106106
mol_dir: str | None = None,
107+
mol_template: str = "id{row}.mol",
108+
split_seed: int | None = None,
109+
conformer_seed: int | None = None,
107110
seed: int = 42,
108111
poscar: str | None = None,
109-
formula_col: int | str = 0,
112+
formula_col: str = "formula",
110113
base_element: str | None = None,
111114
sets: int = 1,
112115
overwrite: bool = False,
@@ -147,7 +150,9 @@ def auto_convert(
147150
property_col=property_col,
148151
train_ratio=train_ratio,
149152
smiles_col=smiles_col,
150-
seed=seed,
153+
mol_template=mol_template,
154+
split_seed=split_seed,
155+
conformer_seed=conformer_seed,
151156
overwrite=overwrite,
152157
)
153158
converted = {

dpa_adapt/data/formula.py

Lines changed: 33 additions & 40 deletions
Original file line numberDiff line numberDiff line change
@@ -217,20 +217,18 @@ def formula_to_npy(
217217
csv_path: str,
218218
output_dir: str,
219219
poscar: str,
220-
formula_col: int | str = 0,
221-
property_col: int | str = 1,
222-
property_name: str = "property",
220+
formula_col: str = "formula",
221+
property_col: str = "Property",
222+
property_name: str = "Property",
223223
base_element: str | None = None,
224224
sets: int = 1,
225225
seed: int = 42,
226226
) -> list[str]:
227227
"""Convert a formula CSV + template POSCAR to ``deepmd/npy`` systems.
228228
229-
CSV format: two or more columns. The formula column holds composition
229+
CSV format: two or more named columns. The formula column holds composition
230230
strings (e.g. ``Ni0.65Gd0.15Fe0.10Co0.05Yb0.05O2H1``); the property
231-
column holds the scalar target value. Header auto-detected: if the first
232-
data row's property column cannot be parsed as ``float``, that row is
233-
skipped as a header.
231+
column holds the scalar target value.
234232
235233
For each CSV row, *sets* random doped structures are generated. Each
236234
structure is written as a ``deepmd/npy`` system under
@@ -244,13 +242,13 @@ def formula_to_npy(
244242
Destination directory for ``deepmd/npy`` output.
245243
poscar : str
246244
Path to template POSCAR (VASP format).
247-
formula_col : int | str
248-
Column index (0-based) or column name for the formula. Default: 0.
249-
property_col : int | str
250-
Column index (0-based) or column name for the property value. Default: 1.
245+
formula_col : str
246+
Column name for the formula. Default: ``"formula"``.
247+
property_col : str
248+
Column name for the property value. Default: ``"Property"``.
251249
property_name : str
252250
Label key written into each system (``set.000/{property_name}.npy``).
253-
Default: ``"property"``.
251+
Default: ``"Property"``.
254252
base_element : str | None
255253
Host element for random substitution. Auto-inferred from the template
256254
POSCAR when ``None``.
@@ -288,21 +286,25 @@ def formula_to_npy(
288286
break
289287
delimiter = "\t" if "\t" in first_line else ","
290288
fh.seek(0)
291-
reader = csv.reader(fh, delimiter=delimiter)
289+
reader = csv.DictReader(fh, delimiter=delimiter)
290+
if reader.fieldnames is None:
291+
raise ValueError(f"No header row found in formula CSV: {csv_path!r}")
292+
formula_header = _resolve_col(formula_col, reader.fieldnames)
293+
property_header = _resolve_col(property_col, reader.fieldnames)
292294
for raw_row in reader:
293-
if not raw_row or all(c.strip() == "" for c in raw_row):
295+
if raw_row is None or all((v or "").strip() == "" for v in raw_row.values()):
294296
continue
295-
row_values = [c.strip() for c in raw_row]
296-
# Resolve column indices from names if needed.
297-
fidx = _resolve_col(formula_col, row_values, allow_name=True)
298-
pidx = _resolve_col(property_col, row_values, allow_name=True)
299-
formula_str = row_values[fidx]
300-
prop_str = row_values[pidx]
297+
formula_str = (raw_row.get(formula_header) or "").strip()
298+
prop_str = (raw_row.get(property_header) or "").strip()
299+
if not formula_str:
300+
raise ValueError(f"Empty formula value in column {formula_header!r}")
301301
try:
302302
prop_val = float(prop_str)
303303
except ValueError:
304-
# Likely a header row — skip.
305-
continue
304+
raise ValueError(
305+
f"Could not parse property value {prop_str!r} "
306+
f"from column {property_header!r}"
307+
) from None
306308
rows.append((formula_str, prop_val))
307309

308310
if not rows:
@@ -367,21 +369,12 @@ def formula_to_npy(
367369

368370

369371
def _resolve_col(
370-
spec: int | str,
371-
row_values: list[str],
372-
allow_name: bool = False,
373-
) -> int:
374-
"""Resolve a column specifier to an integer index.
375-
376-
- *int* → used directly.
377-
- *str* + ``allow_name=True`` → looks up the column name in *row_values*
378-
(case-insensitive), falling back to ``int(spec)``.
379-
"""
380-
if isinstance(spec, int):
381-
return spec
382-
if allow_name:
383-
lower_map = {v.lower(): i for i, v in enumerate(row_values)}
384-
key = spec.lower()
385-
if key in lower_map:
386-
return lower_map[key]
387-
return int(spec)
372+
spec: str,
373+
fieldnames: list[str],
374+
) -> str:
375+
"""Resolve a case-insensitive column name to the exact CSV header."""
376+
lower_map = {name.lower(): name for name in fieldnames if name is not None}
377+
key = str(spec).lower()
378+
if key in lower_map:
379+
return lower_map[key]
380+
raise KeyError(f"Column {spec!r} not found in CSV header {fieldnames}")

0 commit comments

Comments
 (0)