feat(embeddings): expand input token from 145-D to 151-D with physical-functional payload
cb7920d unverified | from __future__ import annotations | |
| import json | |
| import logging | |
| import os | |
| import random | |
| import torch.multiprocessing as mp | |
| from pathlib import Path | |
| from typing import Any | |
| import numpy as np | |
| import torch | |
| from datasets import load_dataset | |
| from sentence_transformers import SentenceTransformer | |
| from transformers import Trainer, TrainingArguments | |
| from pino.pimt_model import FragranceTrajectoryDataset, DEFAULT_EMBEDDING_DIM | |
| from pino.pimt_model_hf import PIMTConfig, PhysicsInformedMixtureTransformer | |
| logger = logging.getLogger("pino.train_hf") | |
| _TEXT_ENCODER: SentenceTransformer | None = None | |
| def get_text_encoder() -> SentenceTransformer: | |
| global _TEXT_ENCODER | |
| if _TEXT_ENCODER is None: | |
| _TEXT_ENCODER = SentenceTransformer("sentence-transformers/all-MiniLM-L6-v2") | |
| return _TEXT_ENCODER | |
| def seed_everything(seed: int = 42) -> None: | |
| """Lock all RNGs for fully reproducible training runs.""" | |
| random.seed(seed) | |
| os.environ["PYTHONHASHSEED"] = str(seed) | |
| np.random.seed(seed) | |
| torch.manual_seed(seed) | |
| torch.cuda.manual_seed(seed) | |
| torch.backends.cudnn.deterministic = True | |
| def pad_trajectory_collate(batch: list) -> dict: | |
| """Custom collate for variable-length ingredient formulations with text conditioning.""" | |
| max_molecules = max(item["tokens"].size(0) for item in batch) | |
| max_timesteps = max(item["physics"].size(0) for item in batch) | |
| bsz = len(batch) | |
| tokens = torch.zeros(bsz, max_molecules, DEFAULT_EMBEDDING_DIM, dtype=torch.float32) | |
| physics = torch.zeros(bsz, max_timesteps, max_molecules, 2, dtype=torch.float32) | |
| src_key_padding_mask = torch.ones(bsz, max_molecules, dtype=torch.bool) | |
| labels_obj = torch.zeros(bsz, max_timesteps, 138, dtype=torch.float32) | |
| labels_sub = torch.zeros(bsz, 7, dtype=torch.float32) | |
| genre_labels = torch.zeros(bsz, dtype=torch.int64) | |
| # Encode text on-the-fly using sentence-transformers only if a record is missing | |
| # its pre-computed embedding. This avoids loading the encoder inside forked | |
| # dataloader workers (which can fail on CUDA re-init). | |
| text_encoder = None | |
| text_embeddings = torch.zeros(bsz, 384, dtype=torch.float32) | |
| for i, item in enumerate(batch): | |
| n_mol = item["tokens"].size(0) | |
| t_steps = item["physics"].size(0) | |
| tokens[i, :n_mol] = item["tokens"] | |
| physics[i, :t_steps, :n_mol] = item["physics"] | |
| src_key_padding_mask[i, :n_mol] = False | |
| labels_obj[i, :t_steps] = item["target_obj"] | |
| labels_sub[i] = item["target_sub"] | |
| if "text_embedding" in item: | |
| text_embeddings[i] = item["text_embedding"].clone().detach().float() if torch.is_tensor(item["text_embedding"]) else torch.tensor(item["text_embedding"], dtype=torch.float32) | |
| elif "text_conditioning" in item: | |
| if text_encoder is None: | |
| text_encoder = get_text_encoder() | |
| emb = text_encoder.encode(item["text_conditioning"], convert_to_numpy=True) | |
| text_embeddings[i] = torch.from_numpy(emb).float() | |
| # Fuse objective and subjective labels into a single (B, T, 151) tensor for HF Trainer (legacy shape). | |
| labels_sub_t = labels_sub.unsqueeze(1).expand(-1, max_timesteps, -1) | |
| labels = torch.cat([labels_obj, labels_sub_t], dim=-1) | |
| return { | |
| "tokens": tokens, | |
| "physics": physics, | |
| "src_key_padding_mask": src_key_padding_mask, | |
| "text_embedding": text_embeddings, | |
| "labels": labels, | |
| } | |
| def compute_metrics(eval_pred) -> dict[str, float]: | |
| """ | |
| Compute isolated objective and subjective metrics from eval predictions. | |
| """ | |
| predictions, labels = eval_pred | |
| obj_pred = predictions[:, :, :138] | |
| obj_true = labels[:, :, :138] | |
| sub_pred = predictions[:, 0, 138:] | |
| sub_true = labels[:, 0, 138:] | |
| obj_mse = float(np.mean((obj_pred - obj_true) ** 2)) | |
| sub_mae = float(np.mean(np.abs(sub_pred - sub_true))) | |
| return { | |
| "objective_mse": round(obj_mse, 6), | |
| "subjective_mae": round(sub_mae, 6), | |
| "eval_loss": round(obj_mse + 0.5 * sub_mae, 6), | |
| } | |
| def export_publication_metrics( | |
| trainer: Trainer, | |
| val_ds: FragranceTrajectoryDataset, | |
| output_path: str = "data/publication_metrics.json", | |
| ) -> None: | |
| """Run evaluation on the validation set and save raw prediction/target pairs.""" | |
| logger.info("Exporting publication validation metrics to %s", output_path) | |
| predictions = trainer.predict(val_ds) | |
| pred_arr = predictions.predictions | |
| true_arr = predictions.label_ids | |
| obj_pred = pred_arr[:, :, :138] | |
| obj_true = true_arr[:, :, :138] | |
| sub_pred = pred_arr[:, 0, 138:] | |
| sub_true = true_arr[:, 0, 138:] | |
| # Save a subset (first 200 records) for graphing predicted-vs-actual. | |
| subset_size = min(200, obj_pred.shape[0]) | |
| metrics = { | |
| "objective": { | |
| "predictions": obj_pred[:subset_size].tolist(), | |
| "targets": obj_true[:subset_size].tolist(), | |
| "mse": float(np.mean((obj_pred - obj_true) ** 2)), | |
| }, | |
| "subjective": { | |
| "predictions": sub_pred[:subset_size].tolist(), | |
| "targets": sub_true[:subset_size].tolist(), | |
| "mae": float(np.mean(np.abs(sub_pred - sub_true))), | |
| }, | |
| } | |
| Path(output_path).parent.mkdir(parents=True, exist_ok=True) | |
| Path(output_path).write_text(json.dumps(metrics, indent=2)) | |
| logger.info("Publication metrics saved") | |
| def run_training() -> None: | |
| """Entry point used by the Hugging Face training Space.""" | |
| main() | |
| def main() -> None: | |
| # Use spawn for dataloader workers so CUDA is safe with multiprocessing. | |
| try: | |
| mp.set_start_method("spawn", force=True) | |
| except Exception: | |
| pass | |
| logging.basicConfig( | |
| level=logging.INFO, | |
| format="%(asctime)s %(levelname)s %(name)s: %(message)s", | |
| ) | |
| logger.info("Training PIMT with Hugging Face transformers") | |
| seed = 42 | |
| seed_everything(seed) | |
| # Load the verified stratified partitions from the Hugging Face Hub. | |
| logger.info("Loading PINO synthetic dataset from Hugging Face Hub") | |
| hub_dataset = load_dataset("mattbitzesty/pino-synthetic-dataset") | |
| train_records = list(hub_dataset["train"]) | |
| val_records = list(hub_dataset["validation"]) | |
| train_ds = FragranceTrajectoryDataset(records=train_records, use_embedding_fallback=True) | |
| val_ds = FragranceTrajectoryDataset(records=val_records, use_embedding_fallback=True) | |
| config = PIMTConfig( | |
| embedding_dim=DEFAULT_EMBEDDING_DIM, | |
| objective_dim=138, | |
| state_dim=2, | |
| hidden_dim=256, | |
| num_heads=8, | |
| num_layers=4, | |
| num_classes_sub=7, | |
| ) | |
| model = PhysicsInformedMixtureTransformer(config) | |
| training_args = TrainingArguments( | |
| output_dir="./models/pino_publication_run", | |
| do_train=True, | |
| do_eval=True, | |
| evaluation_strategy="epoch", | |
| num_train_epochs=5, | |
| per_device_train_batch_size=32, | |
| per_device_eval_batch_size=32, | |
| fp16=True, | |
| gradient_accumulation_steps=4, | |
| dataloader_num_workers=2, | |
| dataloader_pin_memory=True, | |
| logging_steps=10, | |
| logging_dir="./logs/tensorboard", | |
| logging_strategy="steps", | |
| report_to=["tensorboard"], | |
| save_strategy="epoch", | |
| save_total_limit=2, | |
| load_best_model_at_end=True, | |
| metric_for_best_model="eval_loss", | |
| greater_is_better=False, | |
| disable_tqdm=False, | |
| seed=seed, | |
| remove_unused_columns=False, | |
| push_to_hub=True, | |
| hub_model_id="mattbitzesty/pino-pimt", | |
| hub_strategy="end", | |
| ) | |
| trainer = Trainer( | |
| model=model, | |
| args=training_args, | |
| train_dataset=train_ds, | |
| eval_dataset=val_ds, | |
| data_collator=pad_trajectory_collate, | |
| compute_metrics=compute_metrics, | |
| ) | |
| trainer.train() | |
| export_publication_metrics(trainer, val_ds) | |
| logger.info("HF training complete") | |
| if __name__ == "__main__": | |
| main() | |