Views
No views yet
1import sys
2sys.path.insert(1, '/path/to/CogVLM')
3from sat.model import AutoModel
4import argparse
5from utils.models import CogAgentModel, CogVLMModel, FineTuneTestCogAgentModel
6import torch
7from sat.model.mixins import CachedAutoregressiveMixin
8from sat.quantization.kernels import quantize
9from sat.model import AutoModel
10from utils.utils import chat, llama2_tokenizer, llama2_text_processor_inference, get_image_processor
11from utils.models import CogAgentModel, CogVLMModel
12from tqdm import tqdm
13import os
14import argparse
15
16parser = argparse.ArgumentParser()
17parser.add_argument('--temperature', type=float, default=0.5)
18parser.add_argument('--repetition_penalty', type=float, default=1.1)
19args = parser.parse_args()
20args.bf16 = True
21args.stream_chat = False
22args.version = "chat"
23
24# You can download the testset from https://huggingface.co/datasets/SALT-NLP/Design2Code
25test_data_dir = "/path/to/Design2Code"
26predictions_dir = "/path/to/design2code_18b_v0_predictions"
27if not os.path.exists(predictions_dir):
28 try:
29 os.makedirs(predictions_dir)
30 except:
31 pass
32
33filename_list = [filename for filename in os.listdir(test_data_dir) if filename.endswith(".png")]
34world_size = 1
35model, model_args = FineTuneTestCogAgentModel.from_pretrained(
36 f"/path/to/design2code-18b-v0",
37 args=argparse.Namespace(
38 deepspeed=None,
39 local_rank=0,
40 rank=0,
41 world_size=1,
42 model_parallel_size=1,
43 mode='inference',
44 skip_init=True,
45 use_gpu_initialization=True,
46 device='cuda',
47 bf16=True,
48 fp16=None), overwrite_args={'model_parallel_size': world_size} if world_size != 1 else {})
49model = model.eval()
50model.add_mixin('auto-regressive', CachedAutoregressiveMixin())
51
52language_processor_version = model_args.text_processor_version if 'text_processor_version' in model_args else args.version
53print("[Language processor version]:", language_processor_version)
54tokenizer = llama2_tokenizer("lmsys/vicuna-7b-v1.5", signal_type=language_processor_version)
55image_processor = get_image_processor(model_args.eva_args["image_size"][0])
56cross_image_processor = get_image_processor(model_args.cross_image_pix) if "cross_image_pix" in model_args else None
57text_processor_infer = llama2_text_processor_inference(tokenizer, 2048, model.image_length)
58
59def get_html(image_path):
60 with torch.no_grad():
61 history = None
62 cache_image = None
63 # We use an empty string as the query
64 query = ''
65
66 response, history, cache_image = chat(
67 image_path,
68 model,
69 text_processor_infer,
70 image_processor,
71 query,
72 history=history,
73 cross_img_processor=cross_image_processor,
74 image=cache_image,
75 max_length=4096,
76 top_p=1.0,
77 temperature=args.temperature,
78 top_k=1,
79 invalid_slices=text_processor_infer.invalid_slices,
80 repetition_penalty=args.repetition_penalty,
81 args=args
82 )
83
84 return response
85
86for filename in tqdm(filename_list):
87 image_path = os.path.join(test_data_dir, filename)
88 generated_text = get_html(image_path)
89 with open(os.path.join(predictions_dir, filename.replace(".png", ".html")), "w", encoding='utf-8') as f:
90 f.write(generated_text)