""" 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())