File size: 14,993 Bytes
5a00cfd
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
6ba3914
5a00cfd
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
1
2
3
4
5
6
7
8
9
10
11
12
13
14
15
16
17
18
19
20
21
22
23
24
25
26
27
28
29
30
31
32
33
34
35
36
37
38
39
40
41
42
43
44
45
46
47
48
49
50
51
52
53
54
55
56
57
58
59
60
61
62
63
64
65
66
67
68
69
70
71
72
73
74
75
76
77
78
79
80
81
82
83
84
85
86
87
88
89
90
91
92
93
94
95
96
97
98
99
100
101
102
103
104
105
106
107
108
109
110
111
112
113
114
115
116
117
118
119
120
121
122
123
124
125
126
127
128
129
130
131
132
133
134
135
136
137
138
139
140
141
142
143
144
145
146
147
148
149
150
151
152
153
154
155
156
157
158
159
160
161
162
163
164
165
166
167
168
169
170
171
172
173
174
175
176
177
178
179
180
181
182
183
184
185
186
187
188
189
190
191
192
193
194
195
196
197
198
199
200
201
202
203
204
205
206
207
208
209
210
211
212
213
214
215
216
217
218
219
220
221
222
223
224
225
226
227
228
229
230
231
232
233
234
235
236
237
238
239
240
241
242
243
244
245
246
247
248
249
250
251
252
253
254
255
256
257
258
259
260
261
262
263
264
265
266
267
268
269
270
271
272
273
274
275
276
277
278
279
280
281
282
283
284
285
286
287
288
289
290
291
292
293
294
295
296
297
298
299
300
301
302
303
304
305
306
307
308
309
310
311
312
import json
from pprint import pprint
from typing import List, Dict, Optional, Generator
# Loading and printing out config file (want it to be at top of hf space logs)
def load_config(file_path: str = 'config.json') -> Dict:
    with open(file_path, 'r') as file:
        config = json.load(file)
    pprint(config)
    return config
config = load_config('config.json')

### Importing Dependencies, automatically installed from requirements.txt on huggingface space startup) ###
import os
import re
import shutil
import functools
from datetime import datetime

import torch
import numpy as np
import gradio as gr
import huggingface_hub
from datasets import Dataset
from transformers import AutoTokenizer, AutoModelForCausalLM, GenerationConfig, StoppingCriteria, StoppingCriteriaList
from peft import LoraConfig, get_peft_model, PeftModel, PeftConfig, get_peft_model_state_dict
from callbacks import Iteratorize

base_dir = os.path.dirname(os.path.abspath('__file__'))
print(f"Base directory: {base_dir}")

# Looking in "Repository secrets" to get huggingface api key
HF_TOKEN = os.environ.get("HF_TOKEN")
huggingface_hub.login(HF_TOKEN)
api = huggingface_hub.HfApi()

# matching base model to specified LoRA adapater, by using last 2 letters from hf_load_model_name to determine if 3b or 7b model if needed
def extract_base_model(config: Dict) -> str:
    model_variant = config['hf_load_model_name'][-2:].lower()
    if model_variant == "3b":
        return "togethercomputer/RedPajama-INCITE-Chat-3B-v1"
    else:
        return "togethercomputer/RedPajama-INCITE-7B-Chat"
base_model = extract_base_model(config)
print(f"Using base model: {base_model}")

# Creating huggingface tokenizer
tokenizer = AutoTokenizer.from_pretrained(base_model)


### Downloading LoRA adapter & additional instruction prompt config ###
additional_instruction = "The details below description the character, utilizing any relevant detail in your response.\n" # overridden by fine-tuned config
if config['hf_load_model_name']:
    if not os.path.exists("lora_weights"):
        os.mkdir("lora_weights")
    #Download finetuned model
    huggingface_hub.hf_hub_download(repo_id=f"sortxyz/{config['hf_load_model_repo']}",
                                    filename=f"{config['hf_load_model_name']}/adapter_model.bin",
                                    repo_type="model",
                                    local_dir="lora_weights")
    #Download finetuned config
    huggingface_hub.hf_hub_download(repo_id=f"sortxyz/{config['hf_load_model_repo']}",
                                    filename=f"{config['hf_load_model_name']}/adapter_config.json",
                                    repo_type="model",
                                    local_dir="lora_weights")
    huggingface_hub.hf_hub_download(repo_id=f"sortxyz/{config['hf_load_model_repo']}",
                                    filename=f"{config['hf_load_model_name']}/additional_instruction.json",
                                    repo_type="model",
                                    local_dir="lora_weights")
    # Using identicial prompt in inference to that which was used in training
    with open(f"{base_dir}/lora_weights/{config['hf_load_model_name']}/additional_instruction.json", 'r') as file:
        additional_instruction = json.load(file)['additional_instruction']
    os.remove(f"{base_dir}/lora_weights/{config['hf_load_model_name']}/additional_instruction.json")

### Prompt formatting ###
instruction_empty_prompt = """<human>:{}\n""" + additional_instruction + """{}\n<bot>:"""
response_empty_prompt = instruction_empty_prompt + """{}\n<human>:"""
print(f"{'#'*25} Empty response prompt {'#'*25}\n" + response_empty_prompt + f"\n{'#'*23} Empty response prompt END {'#'*23}")

def reformat_row(row):
    """
    Converts 'attributes' list of a given row into a dictionary, reducing the string length by approximately 2x.

    Parameters:
    row (dict): Contains 'id' (int) and 'attributes' (list of dictionaries).
    {'id': 979, 'attributes': [{'value': 'Chillbucks Apron', 'trait_type': 'Accessories'}, {'value': 'Black Shirt', 'trait_type': 'Apparel'}, {'value': 'Blue', 'trait_type': 'Background'}, {'value': 'Uhhh', 'trait_type': 'Expression'}, {'value': 'Short Curly Black', 'trait_type': 'Hair'}, {'value': 'Purple', 'trait_type': 'Skin'}]}

    Returns:
    dict: Original 'id' with 'trait_type' and 'value' pairs from 'attributes' as new keys and values.
    {'id': 979, 'Accessories': 'Chillbucks Apron', 'Apparel': 'Black Shirt', 'Background': 'Blue', 'Expression': 'Uhhh', 'Hair': 'Short Curly Black', 'Skin': 'Purple'}
    """
    return {**{'id': row['id']}, **{pair['trait_type'] : pair['value'] for pair in row['attributes']}}

### Downloading full dataset for character lookup & test dataset for loss metric evaluation ###
if not os.path.exists("downloaded_data"):
    os.mkdir("downloaded_data")
huggingface_hub.hf_hub_download(repo_id=f"sortxyz/{config['hf_complete_dataset_repo']}",
                                filename=config['hf_complete_dataset_name'],
                                repo_type="dataset",
                                local_dir="downloaded_data")
shutil.move(f"downloaded_data/{config['hf_complete_dataset_name']}", 'complete_dataset.jsonl')

### Load complete_dataset, format it, and create dict for characters to be referenced by ID ###
# Currently using 991 rows, hopefully should directly scale up to 20,000 without further finetuning
with open('complete_dataset.jsonl', 'r') as fp:
    complete_dataset_raw = [json.loads(x) for x in fp.readlines()]
complete_dataset = [reformat_row(data_dict) for data_dict in complete_dataset_raw]
data_indexed_by_id = {data_dict['id'] : data_dict for data_dict in complete_dataset} # Used for inference
print(f"length of complete_dataset: {len(complete_dataset)}")

def find_relevant_rows(instruction):
    """Used to append relevant information to character(s) referenced in prompt
    Given a string it extracts the corresponding rows
    e.g. "Describe ID 972 and 979" returns
    {'id': 972, 'Accessories': 'Smoke Frame Glasses', 'Apparel': 'Blue Button Up', 'Background': 'Blue', 'Expression': 'Chill Smile', 'Facial Features': 'Stuble Goatee', 'Hair': 'Short Blond', 'Skin': 'Orange'}
    {'id': 979, 'Accessories': 'Chillbucks Apron', 'Apparel': 'Black Shirt', 'Background': 'Blue', 'Expression': 'Uhhh', 'Hair': 'Short Curly Black', 'Skin': 'Purple'}
    """
    ids_found = [int(n) for n in re.findall(r'\d+', instruction)]
    rows_returned = []
    for id in ids_found:
        if id not in data_indexed_by_id.keys():
            rows_returned.append(f"ID {id} is not contained in dataset")
        else:
            rows_returned.append(str(data_indexed_by_id[id]))
    return "\n".join(rows_returned)

### Load persistent dataset repo ###
persistent_dataset_repo_url = "https://huggingface.co/datasets/sortxyz/persistent-space-dataset"
persistent_data_filename = f"{base_dir}/data/{config['persistent_data_filename']}"

repo = huggingface_hub.Repository(local_dir="data", clone_from=persistent_dataset_repo_url, use_auth_token=HF_TOKEN)
if not os.path.exists(persistent_data_filename):
    persistent_data = {"data": []} # user_data.json has not been created
else:
    with open(persistent_data_filename, 'r') as f:
        persistent_data = json.load(f) # load existing data, which will then be appended to

def store_persistent_information(user_prompt: str, model_prompt: str, generate_output: str, print_commit_url=False):
    """Appends message to persistent data .json file and uploads it to hf dataset repo on every call.
    Note: repo is only retrieved at the start of the runtime, so will not work for spaces in parallele"""
    persistent_data['data'].append({"user_prompt": user_prompt, "model_prompt": model_prompt, "generate_output": generate_output, "time": datetime.now().strftime('%Y-%m-%d %H:%M:%S')})
    with open(persistent_data_filename, 'w') as f:
        json.dump(persistent_data, f, indent=4)
    commit_url = repo.push_to_hub()
    if print_commit_url:
        print(commit_url)

class StopWordsCriteria(StoppingCriteria):
    """
    Class for stopping output generation when the "<human>" token is encountered.
    Prevents the display of any remaining unwanted tokens during streaming.
    Buffering would be faster, but it is incompatible with multiple beams,
    so we must decode after every new token.
    """

    def __init__(self, tokenizer, stop_words: list = ["<human>"], prompt_length: int = -1, stream_callback: bool = False):
        self._tokenizer = tokenizer
        self._stop_words = stop_words
        self.prompt_length = prompt_length
        self._stream_callback = stream_callback

    def __call__(self, input_ids: torch.LongTensor, scores: torch.FloatTensor, **kwargs) -> bool:
        if self.prompt_length == -1:
            self.prompt_length = len(input_ids[0]) - 1

        # Decode model output into text, excluding the prompt
        decoded_text = self._tokenizer.decode(input_ids[0][self.prompt_length:])
        for stop_word in self._stop_words:
            if stop_word in decoded_text:
                return True

        if self._stream_callback:
            # Skips last characters if they are part of "<human>"
            last_characters_to_miss = 0
            for stop_word in self._stop_words:
                for i in range(1, len(stop_word)):
                    if decoded_text.endswith(stop_word[0:i]):
                        last_characters_to_miss = max(i, last_characters_to_miss)

            # Sends data to Gradio using callback, allows interruption of huggingface .generate function after each token
            self._stream_callback(decoded_text) if last_characters_to_miss == 0 else self._stream_callback(decoded_text[:-last_characters_to_miss])
        return False

def evaluate(
        model,
        instruction,
        temperature=0.1,
        top_p=0.75,
        top_k=40,
        num_beams=4,
        max_new_tokens=256,
        stream_output=False,
        store_outputs=True,
        print_progress=True,
        **kwargs,
    ):
    """Evaluate the model using the provided instruction and generation parameters.
    Yields:
        The generated output as a response to the provided instruction (Directly or via streaming).
    """
    prompt = instruction_empty_prompt.format(instruction, find_relevant_rows(instruction))
    if print_progress:
        print("\n" + "#"*50 + "\n" + prompt)

    # tokenizer input string
    inputs = tokenizer(prompt, return_tensors="pt")
    input_ids = inputs["input_ids"].to("cuda:0") #  torch.device(

    #configuration setting for model generation
    generate_params = {
        "input_ids": input_ids,
        "generation_config": GenerationConfig(
                                temperature=temperature,
                                top_p=top_p,
                                top_k=top_k,
                                num_beams=num_beams,
                                **kwargs),
        "return_dict_in_generate": True,
        "output_scores": True,
        "max_new_tokens": max_new_tokens,
    }

    if stream_output:
        ### Generate with streaming ###
        # Streaming the reply 1 token at a time, based on the trick of using 'stopping_criteria' to create an iterator, which is then tracked with a callback
        # ref - https://github.com/oobabooga/text-generation-webui/blob/ad37f396fc8bcbab90e11ecf17c56c97bfbd4a9c/modules/text_generation.py#L216-L243.

        def generate_with_callback(callback=None, **kwargs):
            kwargs.setdefault(
                "stopping_criteria", StoppingCriteriaList(
                    [StopWordsCriteria(tokenizer,
                                       stream_callback=callback)])
            )
            with torch.no_grad():
                model.generate(**kwargs)

        def generate_with_streaming(**kwargs):
            return Iteratorize(generate_with_callback, kwargs, callback=None)

        with generate_with_streaming(**generate_params) as generator:
            for output in generator:
                yield output
    else:
        ### Generate without streaming ###
        with torch.no_grad():
            generation_output = model.generate(
                **generate_params,
                stopping_criteria=StoppingCriteriaList([StopWordsCriteria(tokenizer)])
        )
        unformatted_output = tokenizer.decode(generation_output.sequences[0])
        # Truncate the input prompt, remove unwanted final tokens, strip whitespace and newlines from ends)
        output = unformatted_output[len(prompt):-len('<human>:<|endoftext|>')].strip(' \n')
    if print_progress:
        print(output)
    if store_outputs:
        store_persistent_information(instruction, prompt, output, print_commit_url=True)
    yield output

def create_gradio_interface(
    model: PeftModel,
    config: Dict,
    ) -> gr.Interface:
    """Creates and launches a Gradio Interface using given model and configuration parameters.

    Args:
        model (PeftModel): The model to be evaluated.
        config (dict): Configuration parameters for the Gradio Interface.

    Returns:
        gr.Interface: Gradio interface object."""

    gradio_inputs = [
        gr.components.Textbox(lines=2, label="Instruction", placeholder=config['gradio_placeholder'], value=config['gradio_value']),
        gr.components.Slider(minimum=0, maximum=1, value=0.1, label="Temperature"),
        gr.components.Slider(minimum=0, maximum=1, value=0.75, label="Top p"),
        gr.components.Slider(minimum=0, maximum=100, step=1, value=40, label="Top k"),
        gr.components.Slider(minimum=1, maximum=4, step=1, value=4, label="Beams"),
        gr.components.Slider(minimum=1, maximum=2000, step=1, value=256, label="Max tokens"),
        gr.components.Checkbox(label="Stream output", value=True),
        gr.components.Checkbox(label="Store user prompts (to improve future model versions)", value=False),
    ]

    interface = gr.Interface(
        fn=functools.partial(evaluate, model),
        inputs=gradio_inputs,
        outputs=[gr.inputs.Textbox(lines=5, label="Generation")],
        title=config['gradio_title'],
        description=config['gradio_description'],
        examples=config['gradio_examples'],
        cache_examples=False
    )

    interface.queue(concurrency_count=config['concurrency_count'])
    return interface

### Load base model ###
model = AutoModelForCausalLM.from_pretrained(
    base_model,
    load_in_8bit=True,
    device_map="auto")

### Loading LoRA adapter (if specified) ###
if config['hf_load_model_name']:
    model = PeftModel.from_pretrained(
                model,
                f"lora_weights/{config['hf_load_model_name']}")

# Loss dataset values (each model trained for 3 epochs on 10k prompt dataset)
# 7B Model: 2.40430->1.11667
# 3B Model: 3.95312->1.07058

### indefinitely run gradio interface ###
create_gradio_interface(model, config).launch(debug=True)