An easy-to-use LLM controllable generation tool
Project description
FlexyGen
English | 中文版
FlexyGen is an easy-to-use controllable generation tool for LLMs. Through the idea of inversion of control (IoC), developers can inject a series of triggers into the model to control the model generation process. In the trigger, the generated content can be modified according to the current generation state of the model (currently only supports splicing new content after the current generated sentence).
Installation
pip install flexygen
Getting Started
Demo Code: examples/emoji.py
0. Import Dependencies
import random
from transformers import AutoTokenizer, AutoModelForCausalLM
from flexygen import FlexyGen, GenerationState
1. Load HF Tokenizer and Model
tokenizer = AutoTokenizer.from_pretrained("Qwen/Qwen2.5-0.5B-Instruct")
model = AutoModelForCausalLM.from_pretrained("Qwen/Qwen2.5-0.5B-Instruct")
2. Wrap the Model with FlexyGen Interface
model = FlexyGen.wrap(model, tokenizer)
3. Inject a Splicer Trigger
Inject a splicer trigger named emoji into the model.
The trigger will be called after each token is generated by the model.
If the current sentence ends with ",", ".", "!", "?", the trigger returns a random emoji character and the emoji will be attached after the current sentence. Else the trigger returns None and the current sentence will not be modified.
state stores the current generation state. The current sentence can be accessed through state.input_ids.
@model.splicer("emoji")
def emoji_trigger(state: GenerationState) -> bool:
def random_emoji():
ranges = [
(0x1F600, 0x1F64F),
(0x1F300, 0x1F5FF),
(0x1F680, 0x1F6FF),
(0x2700, 0x27BF),
]
start, end = random.choice(ranges)
code_point = random.randint(start, end)
return chr(code_point)
sentence = tokenizer.batch_decode(state.input_ids)[0].strip()
if sentence.endswith((",", ".", "!", "?")):
return random_emoji() # Returns a random emoji character
4. Generate
input_text = tokenizer.apply_chat_template([
{"role": "user", "content": "Why the sky is blue?"},
], tokenize=False, add_generation_prompt=True)
inputs = tokenizer(input_text, return_tensors="pt").to(model.device)
outputs = model.generate(**inputs, do_sample=True, max_new_tokens=128)
# Emoji attached after each sentence.
print(tokenizer.batch_decode(outputs)[0])
Access other contents in state
In addition to state.input_ids, state can also access other contents:
state.model_kwargs: model keyword arguments, includingattention_mask,past_key_valuesand other keyword args passed to the modelstate.current_length: current sentence length (number of tokens)state.next_tokens: token currently generated by the model, equivalent tostate.input_ids[:, -1]state.next_token_logits: logits of the current tokenstate.next_token_scores: The scores of the current token (logits processed by logit_processor)
If return_dict_in_generate=True and output_scores=True are specified in the generate() method, the state.scores, that is, the scores of all tokens, can be accessed in the trigger.
If return_dict_in_generate=True and output_logits=True are specified in the generate() method, the state.raw_logits, that is, the attention weights of the model, can be accessed in the trigger.
If return_dict_in_generate=True and output_attentions=True are specified in the generate() method, state.decoder_attentions and state.cross_attentions can be accessed in the trigger, which are the attention weights and cross attention weights of the decoder (only for encoder-decoder architecture), respectively
If return_dict_in_generate=True and output_hidden_states=True are specified in the generate() method, state.decoder_hidden_states can be accessed in the trigger, which represents the hidden state vector output by each Transformer layer of the decoder.
Using SentenceLevelFlexyGen
In some applications, it is necessary to trigger certain calls based on the probability of a sentence or the probability of certain tokens in a sentence. For example, adaptive RAG will determine the probability of a token in a sentence to decide whether to trigger a retrieval.
At this time, you can use SentenceLevelFlexyGen:
# ... Omitting model definitions
from flexygen import SentenceLevelFlexyGen, SentenceLevelGenerationState
model = SentenceLevelFlexyGen.wrap(model, tokenizer)
@model.splicer("prob")
def prob_trigger(state: SentenceLevelGenerationState):
if state.end_of_sentences[0]:
if min(state.sentence_token_probs[0]) < 0.1:
# Returns reflection words when the minimum
# token probability of a sentence is lower than 0.1
return " ... Wait, I'm not sure. "
# ... Omitting generation
When using SentenceLevelFlexyGen, state in the trigger becomes a SentenceLevelGenerationState object, which has more content than GenerationState:
state.end_of_sentences: List[bool]: a list of Boolean variables indicating whether each output in a batch generates a complete sentence (by default, ending with the six characters.?!.?!indicates that a sentence has been generated)state.sentence_tokens: List[List[int]]: the current sentencestate.sentence_token_probs: List[List[int]]: the probability of each token in the current sentence
Project details
Download files
Download the file for your platform. If you're not sure which to choose, learn more about installing packages.
Source Distribution
Built Distribution
Filter files by name, interpreter, ABI, and platform.
If you're not sure about the file name format, learn more about wheel file names.
Copy a direct link to the current filters
File details
Details for the file flexygen-0.0.4.tar.gz.
File metadata
- Download URL: flexygen-0.0.4.tar.gz
- Upload date:
- Size: 23.0 kB
- Tags: Source
- Uploaded using Trusted Publishing? No
- Uploaded via: twine/6.0.1 CPython/3.10.13
File hashes
| Algorithm | Hash digest | |
|---|---|---|
| SHA256 |
a9d57c6fb2845caf5315bf68df81bb2bc54dea865a917185cc65721cafa33c94
|
|
| MD5 |
0e16965183f91d0ef00765b7107904d0
|
|
| BLAKE2b-256 |
4fc01ec56c67a510ca8c6c83acefac3317b021d64fa0b3c345fef96dac95e3f0
|
File details
Details for the file flexygen-0.0.4-py3-none-any.whl.
File metadata
- Download URL: flexygen-0.0.4-py3-none-any.whl
- Upload date:
- Size: 25.0 kB
- Tags: Python 3
- Uploaded using Trusted Publishing? No
- Uploaded via: twine/6.0.1 CPython/3.10.13
File hashes
| Algorithm | Hash digest | |
|---|---|---|
| SHA256 |
82148f0f9b147c18ed9b62e4874535830a12e174b14da785fc8841ce44296737
|
|
| MD5 |
73547a6faf0f0211225fe1234703a311
|
|
| BLAKE2b-256 |
eab55f6d818c97773c661f7f9c6cf31f627fd55ae08e9ec269691652daffe0fe
|