Skip to main content

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, including attention_mask, past_key_values and other keyword args passed to the model
  • state.current_length: current sentence length (number of tokens)
  • state.next_tokens: token currently generated by the model, equivalent to state.input_ids[:, -1]
  • state.next_token_logits: logits of the current token
  • state.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 sentence
  • state.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

flexygen-0.0.5.tar.gz (23.0 kB view details)

Uploaded Source

Built Distribution

If you're not sure about the file name format, learn more about wheel file names.

flexygen-0.0.5-py3-none-any.whl (25.0 kB view details)

Uploaded Python 3

File details

Details for the file flexygen-0.0.5.tar.gz.

File metadata

  • Download URL: flexygen-0.0.5.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

Hashes for flexygen-0.0.5.tar.gz
Algorithm Hash digest
SHA256 330aa95119059000d74618e120d453761c171d5786bd8ab26496d9d43b4f2549
MD5 e93dacd56acb0e8fca63bb40f583dcdb
BLAKE2b-256 6c27afb50f595b3e6e8fcc64967a28b07bbc3130e6fa5988793b8a9e36182bc8

See more details on using hashes here.

File details

Details for the file flexygen-0.0.5-py3-none-any.whl.

File metadata

  • Download URL: flexygen-0.0.5-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

Hashes for flexygen-0.0.5-py3-none-any.whl
Algorithm Hash digest
SHA256 a3db57df074381ed7ae4cf293a7e7b60e775e3efd6047db7cd4e360b4165976b
MD5 14c5badbed919af09e9ac15c754d5244
BLAKE2b-256 25f7415d6e8882a27d6d6450fd7287a69039a428c4afd28daf759853b77a4f07

See more details on using hashes here.

Supported by

AWS Cloud computing and Security Sponsor Datadog Monitoring Depot Continuous Integration Fastly CDN Google Download Analytics Pingdom Monitoring Sentry Error logging StatusPage Status page