Views
No views yet
1import torch
2from transformers import AutoModelForCausalLM, AutoProcessor, GenerationConfig, StopStringCriteria
3from PIL import Image
4import time
5
6# For 2 x 24 GB. If using 1 x 48 GB or more (lucky you), you can just use device_map="auto"
7device_map = {
8 "model.vision_backbone": "cpu", # Seems to be required to not run out of memory at 48 GB
9 "model.transformer.wte": 0,
10 "model.transformer.ln_f": 0,
11 "model.transformer.ff_out": 1,
12}
13# For 2 x 24 GB, this works for *only* 38 or 39. Any higher or lower and it'll either only work for 1 token of output or fail completely.
14switch_point = 38 # layer index to switch to second GPU
15device_map |= {f"model.transformer.blocks.{i}": 0 for i in range(0, switch_point)}
16device_map |= {f"model.transformer.blocks.{i}": 1 for i in range(switch_point, 80)}
17
18model_name = "SeanScripts/Molmo-72B-0924-nf4"
19model = AutoModelForCausalLM.from_pretrained(
20 model_name,
21 use_safetensors=True,
22 device_map=device_map,
23 trust_remote_code=True, # Required for Molmo at the moment.
24)
25model.model.vision_backbone.float() # vision backbone needs to be in FP32 for this
26
27processor = AutoProcessor.from_pretrained(
28 model_name,
29 trust_remote_code=True, # Required for Molmo at the moment.
30)
31
32torch.cuda.empty_cache()
33
34image = Image.open("test.png")
35inputs = processor.process(images=image, text="Caption this image.")
36inputs = {k: v.to("cuda:0").unsqueeze(0) for k,v in inputs.items()}
37prompt_tokens = inputs["input_ids"].size(1)
38print("Prompt tokens:", prompt_tokens)
39
40t0 = time.time()
41output = model.generate_from_batch(
42 inputs,
43 generation_config=GenerationConfig(
44 max_new_tokens=256,
45 ),
46 stopping_criteria=[StopStringCriteria(tokenizer=processor.tokenizer, stop_strings=["<|endoftext|>"])],
47 tokenizer=processor.tokenizer,
48)
49t1 = time.time()
50total_time = t1 - t0
51generated_tokens = output.size(1) - prompt_tokens
52time_per_token = generated_tokens/total_time
53print(f"Generated {generated_tokens} tokens in {total_time:.3f} s ({time_per_token:.3f} tok/s)")
54
55response = processor.tokenizer.decode(output[0, prompt_tokens:], skip_special_tokens=True)
56print(response)
57
58torch.cuda.empty_cache()