Views
No views yet
Qwen/Qwen3-4B (causal language model)1import torch
2from transformers import AutoTokenizer
3from sft_qwen_var_classifier import QwenVarClassifier, cnf_valid_mask
4
5# Load model
6model = QwenVarClassifier("Qwen/Qwen3-4B", max_vars=500)
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-4B")
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=500)], 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}")pytorch_model.bin - Model weights (8GB, bfloat16)sft_qwen_var_classifier.py - Model class definition (required for loading)inference_demo.py - Example inference script| Metric | Value |
|---|---|
| Validation Accuracy | 16.36% |
| Validation Loss | 3.87 |
| Random Baseline | ~1% |
| Improvement | 16x |