mirror of
https://github.com/RYDE-WORK/lnp_ml.git
synced 2026-10-01 21:33:22 +08:00
103 lines
3.6 KiB
Python
Executable File
103 lines
3.6 KiB
Python
Executable File
#!/usr/bin/env python3
|
|
"""Generate reproducible /predict/batch payloads for GPU memory tests."""
|
|
|
|
from __future__ import annotations
|
|
|
|
import argparse
|
|
import csv
|
|
import json
|
|
import random
|
|
from pathlib import Path
|
|
|
|
|
|
# Used only when data/interim/internal.csv is not available. These are valid,
|
|
# deliberately varied SMILES; the project dataset is preferred for realistic RAG prompts.
|
|
FALLBACK_SMILES = [
|
|
"CC(C)NCCNC(C)C",
|
|
"CCCCN(CC)CC",
|
|
"CCCCCCCCN(C)C",
|
|
"CCN(CC)CCOC(=O)CCCC",
|
|
"CCCCCCCCCCN(CC)CC",
|
|
"CC(C)(C)OC(=O)NCCN(C)C",
|
|
"CCCCCCCCCCCCN(C)C",
|
|
"CCOC(=O)CN(C)C",
|
|
"CCCCCCN(CC)CC",
|
|
"CN(C)CCOC(=O)CCCCCCCC",
|
|
]
|
|
|
|
|
|
def load_smiles(source: Path | None) -> list[str]:
|
|
if source and source.is_file():
|
|
with source.open(encoding="utf-8-sig", newline="") as handle:
|
|
rows = csv.DictReader(handle)
|
|
values = []
|
|
seen = set()
|
|
for row in rows:
|
|
value = (row.get("smiles") or "").strip()
|
|
if value and value.lower() != "nan" and value not in seen:
|
|
seen.add(value)
|
|
values.append(value)
|
|
if values:
|
|
return values
|
|
return FALLBACK_SMILES
|
|
|
|
|
|
def make_composition(rng: random.Random) -> tuple[float, float, float, float]:
|
|
# Match the ranges used by the optimizer and guarantee an exact 100% sum.
|
|
for _ in range(1000):
|
|
peg = round(rng.uniform(1.0, 5.0), 2)
|
|
cationic = round(rng.uniform(25.0, 53.0), 2)
|
|
phospholipid = round(rng.uniform(8.0, 30.0), 2)
|
|
cholesterol = round(100.0 - peg - cationic - phospholipid, 2)
|
|
if 15.0 <= cholesterol <= 46.0:
|
|
return cationic, phospholipid, cholesterol, peg
|
|
raise RuntimeError("unable to generate a valid composition")
|
|
|
|
|
|
def main() -> None:
|
|
parser = argparse.ArgumentParser(description=__doc__)
|
|
parser.add_argument("--count", type=int, required=True, choices=range(1, 501))
|
|
parser.add_argument("--batch-size", type=int, default=32, choices=range(1, 257))
|
|
parser.add_argument("--use-llm", action=argparse.BooleanOptionalAction, default=True)
|
|
parser.add_argument("--seed", type=int, default=20260922)
|
|
parser.add_argument("--source-csv", type=Path)
|
|
parser.add_argument("--output", type=Path, required=True)
|
|
args = parser.parse_args()
|
|
|
|
rng = random.Random(args.seed)
|
|
smiles_pool = load_smiles(args.source_csv)
|
|
rng.shuffle(smiles_pool)
|
|
items = []
|
|
for index in range(args.count):
|
|
cationic, phospholipid, cholesterol, peg = make_composition(rng)
|
|
items.append(
|
|
{
|
|
"smiles": smiles_pool[index % len(smiles_pool)],
|
|
"cationic_lipid_to_mrna_ratio": round(rng.uniform(7.0, 20.0), 2),
|
|
"cationic_lipid_mol_ratio": cationic,
|
|
"phospholipid_mol_ratio": phospholipid,
|
|
"cholesterol_mol_ratio": cholesterol,
|
|
"peg_lipid_mol_ratio": peg,
|
|
"helper_lipid": ("DOPE", "DSPC")[index % 2],
|
|
"route": ("intravenous", "intramuscular")[index % 2],
|
|
}
|
|
)
|
|
|
|
payload = {
|
|
"items": items,
|
|
"batch_size": args.batch_size,
|
|
"use_llm": args.use_llm,
|
|
}
|
|
args.output.parent.mkdir(parents=True, exist_ok=True)
|
|
args.output.write_text(
|
|
json.dumps(payload, ensure_ascii=False, indent=2) + "\n", encoding="utf-8"
|
|
)
|
|
print(
|
|
f"generated {args.output}: items={args.count}, unique_smiles={len(set(x['smiles'] for x in items))}, "
|
|
f"batch_size={args.batch_size}, use_llm={args.use_llm}"
|
|
)
|
|
|
|
|
|
if __name__ == "__main__":
|
|
main()
|