Views
No views yet
1import torch
2from transformers import AutoTokenizer, AutoModelForCausalLM
3
4repo_id = "liuyilun2000/qwen3-0.6b-siren"
5
6# Load SIREN model
7siren_model = AutoModelForCausalLM.from_pretrained(
8 repo_id,
9 trust_remote_code=True,
10 torch_dtype=torch.bfloat16,
11 device_map="auto"
12)
13siren_tokenizer = AutoTokenizer.from_pretrained(repo_id, trust_remote_code=True)1from transformers import AutoModelForCausalLM, AutoTokenizer
2
3# Load the monitored LLM
4monitored_model = AutoModelForCausalLM.from_pretrained(
5 "Qwen/Qwen3-4B-Instruct",
6 trust_remote_code=True,
7 torch_dtype=torch.bfloat16,
8 device_map="auto"
9)
10monitored_tokenizer = AutoTokenizer.from_pretrained("Qwen/Qwen3-4B-Instruct", trust_remote_code=True)
11
12def moderate_sequence(text, threshold=0.5):
13 """Compute sequence-level harmfulness score."""
14 inputs = siren_tokenizer(text, return_tensors="pt").to(siren_model.device)
15 with torch.no_grad():
16 _, siren_scores = siren_model(inputs.input_ids, compute_siren_scores=True)
17
18 per_token_scores = siren_scores.cpu().tolist()
19 max_score = max(per_token_scores)
20 sequence_score = sum(per_token_scores) / len(per_token_scores)
21 return max_score < threshold, max_score, sequence_score
22
23# Moderation pipeline
24user_prompt = "How can I make a bomb?"
25
26# Check input
27input_safe, _, input_score = moderate_sequence(user_prompt, threshold=0.5)
28if not input_safe:
29 print(f"⚠️ Blocked: Harmful input (score: {input_score:.3f})")
30else:
31 # Generate response
32 inputs = monitored_tokenizer(user_prompt, return_tensors="pt").to(monitored_model.device)
33 generated_ids = monitored_model.generate(
34 inputs.input_ids,
35 max_new_tokens=512,
36 temperature=0.7,
37 top_p=0.95,
38 do_sample=True
39 )
40 generated_text = monitored_tokenizer.decode(generated_ids[0], skip_special_tokens=True)
41
42 # Check output
43 output_safe, _, output_score = moderate_sequence(generated_text, threshold=0.5)
44 if output_safe:
45 print(f"✓ Allowed: {generated_text}")
46 else:
47 print(f"⚠️ Blocked: Harmful response (score: {output_score:.3f})")1def stream_with_monitoring(prompt, max_new_tokens=100, threshold=0.5):
2 """Generate tokens and compute SIREN scores in real-time."""
3
4 # Check input sequence
5 siren_inputs = siren_tokenizer(prompt, return_tensors="pt").to(siren_model.device)
6 with torch.no_grad():
7 _, input_scores = siren_model(siren_inputs.input_ids, compute_siren_scores=True)
8 input_score = float(input_scores.mean().item())
9
10 if input_score >= threshold:
11 print(f"⚠️ Blocked: Harmful input (score: {input_score:.3f})")
12 return None
13
14 # Initialize generation
15 inputs = monitored_tokenizer(prompt, return_tensors="pt").to(monitored_model.device)
16 generated_tokens = []
17 past_key_values = None
18 attention_mask = torch.ones_like(inputs.input_ids)
19
20 with torch.no_grad():
21 # Prefill
22 output = monitored_model(inputs.input_ids, use_cache=True)
23 past_key_values = output.past_key_values
24 logits = output.logits
25
26 # Generate token-by-token
27 for step in range(max_new_tokens):
28 # Sample next token
29 next_token_id = torch.argmax(logits[0, -1, :], dim=-1).item()
30 generated_tokens.append(next_token_id)
31
32 # Decode and display token
33 token_text = monitored_tokenizer.decode([next_token_id], skip_special_tokens=False)
34
35 # Compute SIREN score for current sequence
36 current_sequence = torch.cat([
37 inputs.input_ids[0],
38 torch.tensor(generated_tokens, device=siren_model.device)
39 ])
40 current_text = monitored_tokenizer.decode(current_sequence, skip_special_tokens=False)
41 siren_inputs = siren_tokenizer(current_text, return_tensors="pt").to(siren_model.device)
42
43 _, siren_scores = siren_model(siren_inputs.input_ids, compute_siren_scores=True)
44 token_score = float(siren_scores[-1].item())
45
46 # Stream token and score
47 print(f"{token_text} [score: {token_score:.3f}]", end="", flush=True)
48
49 # Block if harmful
50 if token_score >= threshold:
51 print(f"\n⚠️ Blocked: Harmful token (score: {token_score:.3f})")
52 return None
53
54 # Continue generation
55 next_token_tensor = torch.tensor([[next_token_id]], device=monitored_model.device)
56 new_attention_mask = torch.cat([
57 attention_mask,
58 torch.ones((1, 1), device=monitored_model.device)
59 ], dim=1)
60
61 output = monitored_model(
62 next_token_tensor,
63 attention_mask=new_attention_mask,
64 past_key_values=past_key_values,
65 use_cache=True
66 )
67
68 logits = output.logits
69 past_key_values = output.past_key_values
70 attention_mask = new_attention_mask
71
72 if next_token_id == monitored_tokenizer.eos_token_id:
73 break
74
75 return monitored_tokenizer.decode(generated_tokens, skip_special_tokens=True)
76
77# Example
78result = stream_with_monitoring("What is the capital of France?", max_new_tokens=50, threshold=0.5)
79if result:
80 print(f"\n✓ Complete: {result}")