Views
No views yet
QwenVarClassifier built on Qwen3-4B backbone, trained in two stages:| Model | Avg Reward (Cube Score) | Reward Gap | Top-1 Accuracy |
|---|---|---|---|
| DPO-7500 (this model) | 1.489 | 0.239 | 18.51% |
| SFT Baseline | 1.368 | 0.218 | 25.06% |
| GRPO-5x | 1.418 | 0.221 | 16.51% |
1class QwenVarClassifier(nn.Module):
2 """
3 Input: CNF formula in DIMACS format
4 Output: Logits for each variable (1 to max_vars)
5 """
6 def __init__(self, model_name="Qwen/Qwen3-4B", max_vars=600):
7 self.qwen = AutoModelForCausalLM.from_pretrained(model_name)
8 self.head = nn.Sequential(
9 nn.LayerNorm(hidden_size),
10 nn.Linear(hidden_size, max_vars + 1)
11 )
12
13 def forward(self, input_ids, attention_mask):
14 # Get last token representation
15 outputs = self.qwen(input_ids, attention_mask, output_hidden_states=True)
16 last_hidden = outputs.hidden_states[-1] # (batch, seq, hidden)
17
18 # Pool using last token
19 seq_lengths = attention_mask.sum(dim=1) - 1
20 last_token_hidden = last_hidden[range(batch_size), seq_lengths]
21
22 # Classify
23 logits = self.head(last_token_hidden)
24 return logits1import torch
2from transformers import AutoTokenizer
3
4# Load tokenizer
5tokenizer = AutoTokenizer.from_pretrained("Yale-ROSE/Qwen3-4B-SAT-VarSelector-Sym-Aug-DPO")
6
7# Load model
8checkpoint = torch.load("model.pt", map_location="cuda")
9
10# Initialize model architecture
11from sft_qwen_var_classifier import QwenVarClassifier
12model = QwenVarClassifier("Qwen/Qwen3-4B", max_vars=600)
13model.load_state_dict(checkpoint)
14model.eval()1def predict_variable(cnf_text: str, model, tokenizer, device="cuda"):
2 """Predict the best variable to branch on."""
3 # Tokenize
4 inputs = tokenizer(cnf_text, return_tensors="pt", max_length=8192, truncation=True)
5 inputs = {k: v.to(device) for k, v in inputs.items()}
6
7 # Get valid variable mask from CNF
8 valid_mask = cnf_valid_mask(cnf_text, max_vars=600)
9 valid_mask = torch.tensor(valid_mask, device=device)
10
11 # Predict
12 with torch.no_grad():
13 logits = model(**inputs)
14 # Mask invalid variables
15 logits = logits.masked_fill(~valid_mask.bool(), float('-inf'))
16 pred_var = logits.argmax(dim=-1).item()
17
18 return pred_var
19
20# Example CNF (DIMACS format)
21cnf = """p cnf 5 3
221 -2 3 0
23-1 2 0
242 -3 0
25"""
26
27pred = predict_variable(cnf, model, tokenizer)
28print(f"Predicted variable: {pred}")p cnf <num_vars> <num_clauses>
<literal1> <literal2> ... 0
<literal1> <literal2> ... 0
...1# Training hyperparameters
2beta: 0.1 # DPO temperature
3learning_rate: 1e-6 # Lower LR for DPO
4epochs: 3
5batch_size: 1
6gradient_accumulation: 16
7max_length: 8192
8
9# Data
10train_pairs: 48,982 # Filtered (margin ≥ 0.3)
11valid_pairs: 5,3931def dpo_loss(policy_chosen, policy_rejected, ref_chosen, ref_rejected, beta=0.1):
2 """DPO loss for classifier architecture."""
3 policy_diff = policy_chosen - policy_rejected
4 ref_diff = ref_chosen - ref_rejected
5 logits = beta * (policy_diff - ref_diff)
6 return -F.logsigmoid(logits).mean()| Checkpoint | Avg Reward | Reward Gap |
|---|---|---|
| 1000 | 1.335 | 0.192 |
| 3500 | 1.427 | 0.226 |
| 5000 | 1.460 | 0.235 |
| 7500 | 1.489 | 0.239 |
| 9000 | 1.461 | 0.234 |
Yale-ROSE/SAT-CNF-Rewards (sym_augmented split)Yale-ROSE/SAT-CNF-Rewards (dpo/ folder)
1@misc{sat-var-dpo-2026,
2 title={Neural SAT Variable Selection with Direct Preference Optimization},
3 author={Yale ROSE Lab},
4 year={2026},
5 howpublished={HuggingFace Models},
6 url={https://huggingface.co/Yale-ROSE/Qwen3-4B-SAT-VarSelector-Sym-Aug-DPO}
7}