bengoldberg0 commited on
Commit
9a33401
·
verified ·
1 Parent(s): ec848d1

Add inference script and remove draft status heading

Browse files
Files changed (3) hide show
  1. README.md +0 -5
  2. inference.py +197 -0
  3. requirements.txt +8 -0
README.md CHANGED
@@ -37,11 +37,6 @@ separate external source.
37
  This repository contains a LoRA adapter, not the full IBM base model.
38
  This is an independent student project, not an IBM-endorsed detector.
39
 
40
- ## Release status
41
-
42
- Private draft. Fresh-adapter inference verification and standalone usage
43
- instructions are pending. This section will be updated before public release.
44
-
45
  ## Model and training
46
 
47
  | Setting | Value |
 
37
  This repository contains a LoRA adapter, not the full IBM base model.
38
  This is an independent student project, not an IBM-endorsed detector.
39
 
 
 
 
 
 
40
  ## Model and training
41
 
42
  | Setting | Value |
inference.py ADDED
@@ -0,0 +1,197 @@
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
1
+ import argparse
2
+ import json
3
+ from pathlib import Path
4
+
5
+ import torch
6
+ from huggingface_hub import snapshot_download
7
+ from peft import PeftModel, prepare_model_for_kbit_training
8
+ from transformers import AutoModelForCausalLM, AutoTokenizer, BitsAndBytesConfig
9
+
10
+
11
+ class URLClassifier:
12
+ def __init__(self, repository, revision):
13
+ if not torch.cuda.is_available():
14
+ raise RuntimeError("A CUDA GPU is required.")
15
+ if not torch.cuda.is_bf16_supported():
16
+ raise RuntimeError("A BF16-capable GPU is required.")
17
+
18
+ root = Path(snapshot_download(
19
+ repo_id=repository,
20
+ revision=revision,
21
+ allow_patterns=[
22
+ "adapter_config.json",
23
+ "adapter_model.safetensors",
24
+ "tokenizer.json",
25
+ "tokenizer_config.json",
26
+ "special_tokens_map.json",
27
+ "added_tokens.json",
28
+ "chat_template.jinja",
29
+ "experiment_config/*.json",
30
+ ],
31
+ ))
32
+
33
+ def read_config(name):
34
+ return json.loads(
35
+ (root / "experiment_config" / name).read_text(
36
+ encoding="utf-8"
37
+ )
38
+ )
39
+
40
+ base = read_config("base_model.json")
41
+ self.policy = read_config("input_policy.json")
42
+ prompt = read_config("prompt_config.json")
43
+
44
+ self.instruction = prompt["instruction"]
45
+ self.tokenizer = AutoTokenizer.from_pretrained(
46
+ str(root), local_files_only=True
47
+ )
48
+ if self.tokenizer.pad_token_id is None:
49
+ self.tokenizer.pad_token = self.tokenizer.eos_token
50
+ self.tokenizer.padding_side = "left"
51
+
52
+ self.label_ids = prompt["label_token_ids"]
53
+ for label in ["A", "B"]:
54
+ ids = self.tokenizer.encode(label, add_special_tokens=False)
55
+ if ids != [self.label_ids[label]]:
56
+ raise RuntimeError(f"Unexpected tokenization for {label}.")
57
+
58
+ quantization = BitsAndBytesConfig(
59
+ load_in_4bit=True,
60
+ bnb_4bit_quant_type="nf4",
61
+ bnb_4bit_use_double_quant=True,
62
+ bnb_4bit_compute_dtype=torch.bfloat16,
63
+ )
64
+
65
+ model = AutoModelForCausalLM.from_pretrained(
66
+ base["model_id"],
67
+ revision=base["revision"],
68
+ quantization_config=quantization,
69
+ device_map={"": 0},
70
+ dtype=torch.bfloat16,
71
+ attn_implementation="sdpa",
72
+ )
73
+ model = prepare_model_for_kbit_training(
74
+ model, use_gradient_checkpointing=False
75
+ )
76
+ self.model = PeftModel.from_pretrained(
77
+ model, str(root), is_trainable=False
78
+ )
79
+ self.model.eval()
80
+ self.model.config.use_cache = False
81
+
82
+ def _encode(self, url):
83
+ if not isinstance(url, str) or not url.strip():
84
+ raise ValueError("URL must be a nonempty string.")
85
+
86
+ def encode(text):
87
+ rendered = self.tokenizer.apply_chat_template(
88
+ [{"role": "user", "content": self.instruction + text}],
89
+ tokenize=False,
90
+ add_generation_prompt=True,
91
+ enable_thinking=False,
92
+ )
93
+ return self.tokenizer.encode(
94
+ rendered, add_special_tokens=False
95
+ )
96
+
97
+ prompt_ids = encode(url)
98
+ original_length = len(prompt_ids)
99
+ maximum = self.policy["max_prompt_tokens"]
100
+
101
+ if original_length <= maximum:
102
+ return prompt_ids, False
103
+
104
+ url_ids = self.tokenizer.encode(url, add_special_tokens=False)
105
+
106
+ while len(prompt_ids) > maximum:
107
+ excess = len(prompt_ids) - maximum
108
+ keep = max(0, len(url_ids) - excess - 8)
109
+
110
+ if keep >= len(url_ids):
111
+ raise RuntimeError("Truncation made no progress.")
112
+
113
+ url_ids = url_ids[:keep]
114
+ text = self.tokenizer.decode(
115
+ url_ids,
116
+ skip_special_tokens=False,
117
+ clean_up_tokenization_spaces=False,
118
+ )
119
+ prompt_ids = encode(text)
120
+
121
+ if not url_ids and len(prompt_ids) > maximum:
122
+ raise ValueError("Instructions exceed the token budget.")
123
+
124
+ return prompt_ids, True
125
+
126
+ def predict_urls(self, urls, batch_size=4):
127
+ if not isinstance(batch_size, int) or batch_size < 1:
128
+ raise ValueError("batch_size must be a positive integer.")
129
+
130
+ urls = list(urls)
131
+ results = []
132
+
133
+ for start in range(0, len(urls), batch_size):
134
+ encoded = [
135
+ self._encode(url)
136
+ for url in urls[start:start + batch_size]
137
+ ]
138
+ features = [
139
+ {
140
+ "input_ids": ids,
141
+ "attention_mask": [1] * len(ids),
142
+ }
143
+ for ids, truncated in encoded
144
+ ]
145
+
146
+ inputs = self.tokenizer.pad(
147
+ features, padding=True, return_tensors="pt"
148
+ ).to("cuda")
149
+
150
+ with torch.inference_mode():
151
+ output = self.model(**inputs, use_cache=False)
152
+ logits = output.logits[
153
+ :, -1, [self.label_ids["A"], self.label_ids["B"]]
154
+ ].float()
155
+
156
+ if not torch.isfinite(logits).all().item():
157
+ raise RuntimeError("Non-finite label logits.")
158
+
159
+ margins = logits[:, 1] - logits[:, 0]
160
+ scores = torch.softmax(logits, dim=-1)[:, 1]
161
+
162
+ for margin, score, (_, truncated) in zip(
163
+ margins.cpu().tolist(),
164
+ scores.cpu().tolist(),
165
+ encoded,
166
+ ):
167
+ prediction = 0 if margin > 0 else 1
168
+ results.append({
169
+ "label": "phishing" if prediction == 0 else "legitimate",
170
+ "prediction": prediction,
171
+ "phishing_score_uncalibrated": score,
172
+ "phishing_logit_margin": margin,
173
+ "url_truncated": truncated,
174
+ })
175
+
176
+ del inputs, output, logits, margins, scores
177
+
178
+ return results
179
+
180
+
181
+ def main():
182
+ parser = argparse.ArgumentParser()
183
+ parser.add_argument(
184
+ "--repository",
185
+ default="bengoldberg0/granite-4.2-3b-phishing-url-qlora",
186
+ )
187
+ parser.add_argument("--revision", required=True)
188
+ parser.add_argument("--url", required=True)
189
+ args = parser.parse_args()
190
+
191
+ classifier = URLClassifier(args.repository, args.revision)
192
+ result = classifier.predict_urls([args.url])[0]
193
+ print(json.dumps(result, indent=2))
194
+
195
+
196
+ if __name__ == "__main__":
197
+ main()
requirements.txt ADDED
@@ -0,0 +1,8 @@
 
 
 
 
 
 
 
 
 
1
+ torch==2.11.0+cu128
2
+ transformers==5.17.0
3
+ peft==0.21.0
4
+ accelerate==1.15.0
5
+ bitsandbytes==0.50.2
6
+ huggingface_hub==1.31.0
7
+ safetensors==0.8.0
8
+ tokenizers==0.23.1