RAIC / test_upload.py
thinkhong's picture
Upload FTT UKB CVD classification model (20240603_1_UKB_CVD_cls_hp_search)
86f5ee3 verified
Raw History Blame Contribute Delete
5.42 kB
"""
Smoke test for the thinkhong/RAIC checkpoint on the Hugging Face Hub.
Downloads config/model/tokenizers via `trust_remote_code=True` and runs a forward pass on a
single synthetic example covering every continuous and categorical variable the model was
trained with. Raw values for a few named variables are min-max normalized with the provided
`train_continuous_variables_min_max_values.json` to show real preprocessing; all other
variables get a neutral placeholder value since this is a connectivity/parity smoke test,
not a clinical prediction.
Usage:
python test_upload.py [--repo-id thinkhong/RAIC] [--revision main]
"""
import argparse
import json
import sys
import torch
from huggingface_hub import hf_hub_download
from transformers import AutoConfig, AutoModelForSequenceClassification, AutoTokenizer
CLS_ID = 1
SPECIALS = {"[UNK]", "[CLS]", "[SEP]", "[PAD]", "[MASK]"}
def minmax_normalize(x, x_min, x_max, new_min=1.0, new_max=3.0):
return (x - x_min) / (x_max - x_min) * (new_max - new_min) + new_min
def build_example(bias_tok, weights_tok, min_max, raw_values=None):
raw_values = raw_values or {}
bias_vocab = bias_tok.get_vocab()
weight_vocab = weights_tok.get_vocab()
bias_names = [n for n in bias_vocab if n not in SPECIALS]
cont_vars = sorted(n for n in bias_names if n.endswith("_Continuous") or n.endswith("_Integer"))
cat_vars = sorted(n for n in bias_names if n not in cont_vars)
num_input_val = [1.0] # CLS weight
for v in cont_vars:
if v in raw_values and v in min_max:
x_min, x_max = min_max[v]
num_input_val.append(minmax_normalize(raw_values[v], x_min, x_max))
else:
num_input_val.append(2.0) # neutral placeholder (mid of the [1, 3] normalization range)
num_input_ids = [CLS_ID] + [weight_vocab[v] for v in cont_vars]
num_variable_ids = [CLS_ID] + [bias_vocab[v] for v in cont_vars]
cat_value_tokens = []
for cv in cat_vars:
token = f"{cv}_{int(raw_values[cv])}" if cv in raw_values else None
if token is None or token not in weight_vocab:
token = sorted(w for w in weight_vocab if w.startswith(cv + "_"))[0]
cat_value_tokens.append(token)
cat_input_ids = [weight_vocab[t] for t in cat_value_tokens]
cat_variable_ids = [bias_vocab[v] for v in cat_vars]
seq_len = len(num_input_ids) + len(cat_input_ids)
batch = {
"num_input_val": torch.tensor([num_input_val], dtype=torch.float32),
"num_input_ids": torch.tensor([num_input_ids], dtype=torch.long),
"cat_input_ids": torch.tensor([cat_input_ids], dtype=torch.long),
"num_variable_ids": torch.tensor([num_variable_ids], dtype=torch.long),
"cat_variable_ids": torch.tensor([cat_variable_ids], dtype=torch.long),
"attention_mask": torch.tensor([[1] * seq_len], dtype=torch.long),
}
return batch, {"cont_vars": cont_vars, "cat_vars": cat_vars, "seq_len": seq_len}
def main():
parser = argparse.ArgumentParser()
parser.add_argument("--repo-id", default="thinkhong/RAIC")
parser.add_argument("--revision", default=None)
args = parser.parse_args()
print(f"Loading config/model from '{args.repo_id}' (trust_remote_code=True) ...")
config = AutoConfig.from_pretrained(args.repo_id, revision=args.revision, trust_remote_code=True)
model = AutoModelForSequenceClassification.from_pretrained(
args.repo_id, revision=args.revision, trust_remote_code=True
)
model.eval()
print(f" architecture: {config.architectures}")
print(f" num_labels: {config.num_labels} id2label: {config.id2label}")
print(f" params: {sum(p.numel() for p in model.parameters()):,}")
print("Loading tokenizers ...")
bias_tok = AutoTokenizer.from_pretrained(
args.repo_id, revision=args.revision, subfolder="ftt_variable_bias_tokenizer_20240530"
)
weights_tok = AutoTokenizer.from_pretrained(
args.repo_id, revision=args.revision, subfolder="ftt_variable_weights_tokenizer_20240530"
)
print("Loading min/max normalization stats ...")
min_max_path = hf_hub_download(
args.repo_id, "train_continuous_variables_min_max_values.json", revision=args.revision
)
min_max = json.load(open(min_max_path))
# A few named, semi-realistic raw values; every other variable gets a neutral placeholder.
raw_values = {
"AGE_Continuous": 55.0,
"SBP_Continuous": 130.0,
"SEX_Categorical": 0,
}
batch, meta = build_example(bias_tok, weights_tok, min_max, raw_values)
print(
f"Built example: seq_len={meta['seq_len']} "
f"({len(meta['cont_vars'])} continuous + {len(meta['cat_vars'])} categorical + 1 CLS)"
)
with torch.no_grad():
out1 = model(**batch)
out2 = model(**batch) # determinism check
assert torch.allclose(out1.logits, out2.logits), "Model is non-deterministic in eval mode!"
probs = torch.softmax(out1.logits, dim=-1)[0]
pred_id = int(probs.argmax())
print()
print("Logits:", out1.logits.tolist())
print("Probs: ", probs.tolist())
print(f"Predicted class: {pred_id} ({config.id2label[pred_id]})")
print()
print(
"OK: model downloaded from the Hub, loaded with trust_remote_code=True, "
"and produced a deterministic forward pass."
)
if __name__ == "__main__":
sys.exit(main())