Views
No views yet
rom swift.llm import get_model_tokenizer, get_template
from swift.llm.dataset.preprocessor import MessagesPreprocessor
from swift.llm.template.template_inputs import TemplateInputs
from transformers import CLIPProcessor, CLIPModel
from PIL import Image
import torch
from swift.llm.utils import to_device
from modelscope import snapshot_download
from glavis import register_laq
image_features = torch.randn(1, 768)
system_prompt = """
#Role
You are a multimodal conversational recommendation assistant capable of understanding user needs and recommending suitable products through both textual and visual reasoning.
# Task
Based on the conversation history and user requirements, perform reasoning, generate retrieval clues (both textual and latent visual queries), and provide a natural response
# Output Format (Strictly Follow)
Your output MUST follow this exact structure:
<think> step-by-step reasoning process regarding user intent and retrieval strategy </think>
<ret> key textual retrieval keywords and recommended product descriptions </ret>
<|vis_emb_start|>latent visual queries<|vis_emb_end|>
<answer> natural response </answer>
# Critical Rules
- Token Usage: <|vis_emb_start|> and <|vis_emb_end|> wrap the latent visual queries.
- Visual Semantics: The content within visual tags must represent high-dimensional visual embeddings, capturing nuances hard to describe in text.
- Textual and latent visual queries should be complementary, not redundant
- Maintain coherent, helpful conversational tone in final response
"""
data = {
"messages":
[
{"role": "system", "content": system_prompt},
{"role": "user", "content": "Hi"},
{"role": "assistant", "content": "Hello, please tell me how can i help you today?", "loss": False},
{"role": "user", "content": "I would love to see some kurta having embroidered prints ."}],
"images": ["MMD/images//41kflO7sFBL.jpg"],
"pos_images": ["MMD/images//31oV%2Brdz-BL.jpg"],
"false_images": ["MMD/images//41kflO7sFBL.jpg", "MMD/images//41HBIbagzJL.jpg", "MMD/images//419CI7BfxCL.jpg", "MMD/images//41LL9ChGKGL.jpg", "MMD/images//51nriE9U4ML.jpg", "MMD/images//41kIuXyBi%2BL.jpg", "MMD/images//31on5t56uoL.jpg", "MMD/images//41bZg24ygcL.jpg", "MMD/images//41YrZ6ft97L.jpg", "MMD/images//51gwMOcnLQL.jpg", "MMD/images//41sS%2BQhPLFL.jpg", "MMD/images//41hLueZ0bDL.jpg", "MMD/images//41dN9-npvbL.jpg", "MMD/images//410uU0vpFgL.jpg", "MMD/images//41CT84Y98xL.jpg", "MMD/images//31LhdbzxCcL.jpg", "MMD/images//316ieWt8bqL.jpg", "MMD/images//41qtqygf1HL.jpg", "MMD/images//41yQIl6u0aL.jpg", "MMD/images//41dD7YodiJL.jpg", "MMD/images//41vPc9wxTpL.jpg", "MMD/images//41jJiSLDYPL.jpg", "MMD/images//41i5Nv6NWtL.jpg", "MMD/images//31yO8UyYIYL.jpg", "MMD/images//41i1jV7w4eL.jpg", "MMD/images//31PD6JjgOLL.jpg", "MMD/images//31OjBNe-3%2BL.jpg", "MMD/images//419xuMjU0LL.jpg", "MMD/images//41IUnJfcruL.jpg", "MMD/images//41eUZaBErKL.jpg", "MMD/images//31imbDMB9FL.jpg", "MMD/images//31YgitQ2IpL.jpg", "MMD/images//41DIhjLdOXL.jpg", "MMD/images//41c8dx6nGPL.jpg", "MMD/images//31wWwT1iskL.jpg", "MMD/images//41BGoy9GerL.jpg", "MMD/images//41wuXeX%2Bu9L.jpg", "MMD/images//41uHRGS4SkL.jpg", "MMD/images//41T%2BSQbrfZL.jpg", "MMD/images//41eIuFRfk%2BL.jpg", "MMD/images//319FaTWKUpL.jpg", "MMD/images//51nEx7wIBdL.jpg", "MMD/images//41aNC9mjdhL.jpg", "MMD/images//318oSeR8UpL.jpg", "MMD/images//412cgWIXBcL.jpg", "MMD/images//412XTLi8jgL.jpg", "MMD/images//41ereT-rc1L.jpg", "MMD/images//41sdH7w42JL.jpg", "MMD/images//41IOaLtxVWL.jpg", "MMD/images//41L8r-%2Bvo0L.jpg"]
}
model, processor = get_model_tokenizer('GLaVis/GLaVis_MMD',
model_type = 'qwen3_vl_laq',
torch_dtype = torch.bfloat16,
attn_implementation="flash_attention_2",
device_map="cuda",)
processor.tokenizer.padding_side = 'left'
template = get_template('qwen3_vl_laq', processor)
print("hf_device_map", model.hf_device_map)
print("device", next(model.parameters()).device)
from transformers import LogitsProcessor, LogitsProcessorList
class VectorizedVisualBlockProcessor(LogitsProcessor):
def __init__(self, start_id, end_id, pad_id, vis_len=4):
self.start_id = start_id
self.end_id = end_id
self.pad_id = pad_id
self.vis_len = vis_len
def __call__(self, input_ids: torch.LongTensor, scores: torch.FloatTensor) -> torch.FloatTensor:
batch_size, seq_len = input_ids.shape
device = input_ids.device
positions = torch.arange(seq_len, device=device)
last_start_idxs = torch.where(
input_ids == self.start_id,
positions,
torch.tensor(-1, device=device)
).max(dim=1).values
last_end_idxs = torch.where(
input_ids == self.end_id,
positions,
torch.tensor(-1, device=device)
).max(dim=1).values
is_active = (last_start_idxs != -1) & (last_start_idxs > last_end_idxs)
if is_active.any():
dist_from_start = (seq_len - 1) - last_start_idxs
need_pad_mask = is_active & (dist_from_start < self.vis_len)
need_end_mask = is_active & (dist_from_start >= self.vis_len)
if need_pad_mask.any():
scores[need_pad_mask, :] = -float("inf")
scores[need_pad_mask, self.pad_id] = 0
if need_end_mask.any():
scores[need_end_mask, :] = -float("inf")
scores[need_end_mask, self.end_id] = 0
is_inactive = ~is_active
if is_inactive.any():
scores[is_inactive, self.pad_id] = -float("inf")
scores[is_inactive, self.end_id] = -float("inf")
return scores
start_id = processor.tokenizer.convert_tokens_to_ids("<|vis_emb_start|>")
pad_id = processor.tokenizer.convert_tokens_to_ids("<vis_emb_pad>")
end_id = processor.tokenizer.convert_tokens_to_ids("<|vis_emb_end|>")
vis_processor = VectorizedVisualBlockProcessor(
start_id=start_id,
pad_id=pad_id,
end_id=end_id,
vis_len=model.config.num_visual_tokens,
)
logits_processor_list = LogitsProcessorList([vis_processor])
preprocessor = MessagesPreprocessor()
processed_data = preprocessor.preprocess(data)
encoded = template.encode(processed_data)
device = model.device
input_dict = {}
for k, v in encoded.items():
if isinstance(v, list):
input_dict[k] = torch.tensor(v, device=device, dtype=torch.long).unsqueeze(0)
elif isinstance(v, torch.Tensor):
input_dict[k] = v.to(device)
else:
input_dict[k] = v
with torch.no_grad(), torch.inference_mode():
generated_outputs = model.generate(
**input_dict,
max_new_tokens=1024,
use_cache=True,
logits_processor=logits_processor_list,
do_sample=True,
temperature=0.1,
repetition_penalty=1.1,
)
input_token_count = input_dict['input_ids'].shape[1]
end_time = time.time()
total_inference_time = end_time - start_time
total_token_count = generated_outputs.shape[1]
new_tokens_generated = total_token_count - input_token_count
if new_tokens_generated > 0:
time_per_token = total_inference_time / new_tokens_generated
tokens_per_second = new_tokens_generated / total_inference_time
else:
time_per_token = 0
tokens_per_second = 0
print(f"total_inference_time: {total_inference_time:.4f} s")
print(f"new_tokens_generated: {new_tokens_generated}")
print(f"time_per_token (Latency): {time_per_token * 1000:.2f} ms")
print(f"tokens_per_second (Throughput): {tokens_per_second:.2f} tokens/s")
generated_ids_trimmed = [out_ids[len(in_ids) :] for in_ids, out_ids in zip(input_dict['input_ids'], generated_outputs)]
output_texts = processor.batch_decode(generated_ids_trimmed, skip_special_tokens=False)
print(output_texts[0].encode('utf-8').decode('latin1'))