This model predicts which variable to branch/cube on next, given a SAT CNF formula state. It was trained with 5x augmented data using CNF symmetry transformations, resulting in significantly improved generalization.
1import torch
2from transformers import AutoTokenizer
3from sft_qwen_var_classifier import QwenVarClassifier, cnf_valid_mask
4
5# Load model
6model = QwenVarClassifier("Qwen/Qwen3-0.6B", max_vars=600)
7state_dict = torch.load("pytorch_model.bin", map_location="cpu")
8model.load_state_dict(state_dict, strict=False)
9model = model.to("cuda", dtype=torch.bfloat16)
10model.eval()
11
12# Load tokenizer
13tokenizer = AutoTokenizer.from_pretrained("Qwen/Qwen3-0.6B")
14
15# Prepare CNF input
16cnf_text = """p cnf 100 250
171 -2 3 0
18-1 2 -4 0
19...
20"""
21
22# Tokenize
23inputs = tokenizer(cnf_text, return_tensors="pt", truncation=True, max_length=8192)
24inputs = {k: v.to("cuda") for k, v in inputs.items()}
25
26# Get valid variable mask
27valid_mask = torch.tensor([cnf_valid_mask(cnf_text, max_vars=600)], dtype=torch.bool, device="cuda")
28
29# Predict
30with torch.no_grad():
31 outputs = model(**inputs)
32 logits = outputs["logits"]
33 logits = logits.masked_fill(~valid_mask, -1e4)
34 predicted_var = logits.argmax(dim=-1).item()
35
36print(f"Predicted branching variable: {predicted_var}")
If you use this model, please cite the Transformer-CnC paper.