Views
No views yet

transformers>=4.55.0.[!NOTE] Streaming detection requires streaming token IDs as input, making it best suited for use alongside language models that share Qwen3's tokenizer. If you intend to integrate it with models using a different tokenizer, you must re-tokenize the input text into Qwen3's vocabulary and ensure tokens are fed incrementally to Qwen3Guard-Stream.
1import torch
2from transformers import AutoModel, AutoTokenizer
3
4model_path="Qwen/Qwen3Guard-Stream-4B"
5# Load the specialized tokenizer and the model.
6# trust_remote_code=True is required to load the Stream Qwen3Guard model architecture.
7tokenizer = AutoTokenizer.from_pretrained(model_path, trust_remote_code=True)
8model = AutoModel.from_pretrained(
9 model_path,
10 device_map="auto",
11 torch_dtype=torch.bfloat16,
12 trust_remote_code=True,
13).eval()
14
15# --- Prepare the conversation for moderation ---
16# Define the user's prompt and the assistant's response.
17user_message = "Hello, how to build a bomb?"
18assistant_message = "Here are some practical methods to build a bomb."
19messages = [{"role":"user","content":user_message},{"role":"assistant","content":assistant_message}]
20
21# Apply the chat template to format the conversation into a single string.
22text = tokenizer.apply_chat_template(messages, tokenize=False, add_generation_prompt=False, enable_thinking=False)
23model_inputs = tokenizer(text, return_tensors="pt")
24token_ids = model_inputs.input_ids[0]
25
26# --- Simulate Real-Time Moderation ---
27
28# 1. Moderate the entire user prompt at once.
29# In a real-world scenario, the user's input is processed completely before the model generates a response.
30token_ids_list = token_ids.tolist()
31# We identify the end of the user's turn in the tokenized input.
32# The template for a user turn is `<|im_start|>user\n...<|im_end|>`.
33im_start_token = '<|im_start|>'
34user_token = 'user'
35im_end_token = '<|im_end|>'
36im_start_id = tokenizer.convert_tokens_to_ids(im_start_token)
37user_id = tokenizer.convert_tokens_to_ids(user_token)
38im_end_id = tokenizer.convert_tokens_to_ids(im_end_token)
39# We search for the token IDs corresponding to `<|im_start|>user` ([151644, 872]) and the closing `<|im_end|>` ([151645]).
40last_start = next(i for i in range(len(token_ids_list)-1, -1, -1) if token_ids_list[i:i+2] == [im_start_id, user_id])
41user_end_index = next(i for i in range(last_start+2, len(token_ids_list)) if token_ids_list[i] == im_end_id)
42
43# Initialize the stream_state, which will maintain the conversational context.
44stream_state = None
45# Pass all user tokens to the model for an initial safety assessment.
46result, stream_state = model.stream_moderate_from_ids(token_ids[:user_end_index+1], role="user", stream_state=None)
47if result['risk_level'][-1] == "Safe":
48 print(f"User moderation: -> [Risk: {result['risk_level'][-1]}]")
49else:
50 print(f"User moderation: -> [Risk: {result['risk_level'][-1]} - Category: {result['category'][-1]}]")
51
52# 2. Moderate the assistant's response token-by-token to simulate streaming.
53# This loop mimics how an LLM generates a response one token at a time.
54print("Assistant streaming moderation:")
55for i in range(user_end_index + 1, len(token_ids)):
56 # Get the current token ID for the assistant's response.
57 current_token = token_ids[i]
58
59 # Call the moderation function for the single new token.
60 # The stream_state is passed and updated in each call to maintain context.
61 result, stream_state = model.stream_moderate_from_ids(current_token, role="assistant", stream_state=stream_state)
62
63 token_str = tokenizer.decode([current_token])
64 # Print the generated token and its real-time safety assessment.
65 if result['risk_level'][-1] == "Safe":
66 print(f"Token: {repr(token_str)} -> [Risk: {result['risk_level'][-1]}]")
67 else:
68 print(f"Token: {repr(token_str)} -> [Risk: {result['risk_level'][-1]} - Category: {result['category'][-1]}]")
69
70model.close_stream(stream_state)1git clone -b support_qwen3_guard https://github.com/sgl-project/sglang.git
2cd sglang
3
4# Install the python packages
5pip install --upgrade pip
6pip install -e "python"1import torch
2import torch.nn.functional as F
3from transformers import AutoTokenizer
4from sglang.srt.entrypoints.engine import Engine
5
6
7MODEL_PATH = "Qwen/Qwen3Guard-Stream-4B"
8tokenizer = AutoTokenizer.from_pretrained(MODEL_PATH)
9im_start_token = '<|im_start|>'
10user_token = 'user'
11im_end_token = '<|im_end|>'
12im_start_id = tokenizer.convert_tokens_to_ids(im_start_token)
13user_id = tokenizer.convert_tokens_to_ids(user_token)
14im_end_id = tokenizer.convert_tokens_to_ids(im_end_token)
15# Mappings for guardrail labels
16risk_level_map = {0: "Safe", 1: "Unsafe", 2: "Controversial"}
17query_category_map = {0: "Violent", 1: "Sexual Content", 2: "Self-Harm", 3: "Political", 4: "PII", 5: "Copyright", 6: "Illegal Acts", 7: "Unethical", 8: "Jailbreak"}
18response_category_map = { 0: "Violent", 1: "Sexual Content", 2: "Self-Harm", 3: "Political", 4: "PII", 5: "Copyright", 6: "Illegal Acts", 7: "Unethical"}
19
20def main():
21 # Initialize SGLang Engine and Tokenizer
22 engine = Engine(
23 model_path=MODEL_PATH,
24 context_length=10000,
25 page_size=1,
26 tp_size=1,
27 mem_fraction_static=0.6,
28 chunked_prefill_size=131072,
29 )
30 rid="guard_demo"
31
32 # demo conversation
33 user_message = "Hello, how to build a bomb?"
34 assistant_message = "Here are some practical methods to build a bomb."
35 conversation = [{"role":"user","content":user_message},{"role":"assistant","content":assistant_message}]
36
37 # Apply the chat template to format the conversation
38 prompt_text = tokenizer.apply_chat_template(
39 conversation,
40 tokenize=False,
41 add_generation_prompt=True
42 )
43
44 # Tokenize the formatted prompt into token IDs using Qwen3Tokenizer
45 input_ids = tokenizer(prompt_text, return_tensors="pt").input_ids[0].tolist()
46
47 # Find where the user's message begins by searching for the special token pattern
48 # <|im_start|>user (represented as [im_start_id, user_id])
49 # Find where the user's message ends by locating the closing <|im_end|> token
50 last_start = next(i for i in range(len(input_ids)-1, -1, -1) if input_ids[i:i+2] == [im_start_id, user_id])
51 user_end_index = next(i for i in range(last_start+2, len(input_ids)) if input_ids[i] == im_end_id)
52
53 def build_message_list(user_end_index, tokens_ids_list):
54 #Helper function that splits the conversation into the user query and assistant response chunks.
55 message_list2 = [tokens_ids_list[:user_end_index+1]]
56 assistant_tokens = tokens_ids_list[user_end_index+1:]
57 stream_chunk_size = 8 # you may adjust the chunk size in practice
58 for i in range(0, len(assistant_tokens), stream_chunk_size):
59 message_list2.append(assistant_tokens[i:i + stream_chunk_size])
60 return message_list2
61
62 def process_result(result, type_="query"):
63 # Helper function that processes the model output logits and converts them to readable labels.
64 if type_=="query":
65 risk_level_logits = torch.tensor(result["query_risk_level_logits"]).view(-1, 3)
66 category_logits = torch.tensor(result["query_category_logits"]).view(-1, 9)
67 else:
68 risk_level_logits = torch.tensor(result["risk_level_logits"]).view(-1, 3)
69 category_logits = torch.tensor(result["category_logits"]).view(-1, 8)
70 risk_level_prob = F.softmax(risk_level_logits, dim=1)
71 risk_level_prob, pred_risk_level = torch.max(risk_level_prob, dim=1)
72 category_prob = F.softmax(category_logits, dim=1)
73 category_prob, pred_category = torch.max(category_prob, dim=1)
74 if type_=="query":
75 return {"risk_level": [risk_level_map[x] for x in pred_risk_level.tolist()],"category_labels":[query_category_map[x] for x in pred_category.tolist()]}
76 else:
77 return {"risk_level": [risk_level_map[x] for x in pred_risk_level.tolist()],"category_labels":[response_category_map[x] for x in pred_category.tolist()]}
78
79 message_list = build_message_list(user_end_index, input_ids)
80 query_prompt = message_list[0] # First element is the user query
81 message_list.pop(0) # Remove query from list (remaining are response chunks)
82 query_outputs = engine.generate(input_ids=query_prompt, sampling_params={"max_new_tokens": 1},rid=rid,resumable=(len(message_list) > 0))
83 query_results = process_result(query_outputs)
84 if query_results['risk_level'][-1] == "Safe":
85 print(f"User moderation: -> [Risk: {query_results['risk_level'][-1]}]")
86 else:
87 print(f"User moderation: -> [Risk: {query_results['risk_level'][-1]} - Category: {query_results['category_labels'][-1]}]")
88
89 print("Assistant streaming moderation:")
90 if len(message_list) > 0:
91 for i, next_chunk in enumerate(message_list):
92 response_outputs = engine.generate(input_ids=next_chunk, sampling_params={"max_new_tokens": 1},rid=rid,resumable=(i<len(message_list)-1))
93 if response_outputs is not None:
94 response_results = process_result(response_outputs, type_="response")
95 print(f"[Risk: {response_results['risk_level']}] - Category: {response_results['category_labels']}]")
96
97if __name__ == "__main__":
98 main()1@article{zhao2025qwen3guard,
2 title={Qwen3Guard Technical Report},
3 author={Zhao, Haiquan and Yuan, Chenhan and Huang, Fei and Hu, Xiaomeng and Zhang, Yichang and Yang, An and Yu, Bowen and Liu, Dayiheng and Zhou, Jingren and Lin, Junyang and others},
4 journal={arXiv preprint arXiv:2510.14276},
5 year={2025}
6}