Views
No views yet
1from transformers import pipeline, AutoTokenizer, AutoModelForSeq2SeqLM
2
3device = "cuda" if torch.cuda.is_available() else "cpu"
4
5# Model checkpoint
6model_checkpoint = "Hatman/Flux-Prompt-Enhance"
7
8# Tokenizer
9tokenizer = AutoTokenizer.from_pretrained(model_checkpoint)
10
11# Model
12model = AutoModelForSeq2SeqLM.from_pretrained(model_checkpoint)
13
14enhancer = pipeline('text2text-generation',
15 model=model,
16 tokenizer=tokenizer,
17 repetition_penalty= 1.2,
18 device=device)
19
20max_target_length = 256
21prefix = "enhance prompt: "
22
23short_prompt = "beautiful house with text 'hello'"
24answer = enhancer(prefix + short_prompt, max_length=max_target_length)
25final_answer = answer[0]['generated_text']
26print(final_answer)
27
28# a two-story house with white trim, large windows on the second floor,
29# three chimneys on the roof, green trees and shrubs in front of the house,
30# stone pathway leading to the front door, text on the house reads "hello" in all caps,
31# blue sky above, shadows cast by the trees, sunlight creating contrast on the house's facade,
32# some plants visible near the bottom right corner, overall warm and serene atmosphere. 1import torch
2import random
3import hashlib
4from transformers import pipeline, AutoTokenizer, AutoModelForSeq2SeqLM
5
6class PromptEnhancer:
7 def __init__(self):
8 # Set up device
9 self.device = "cuda" if torch.cuda.is_available() else "cpu"
10
11 # Model checkpoint
12 self.model_checkpoint = "Hatman/Flux-Prompt-Enhance"
13
14 # Tokenizer and Model
15 self.tokenizer = AutoTokenizer.from_pretrained(self.model_checkpoint)
16 self.model = AutoModelForSeq2SeqLM.from_pretrained(self.model_checkpoint).to(self.device)
17
18 # Initialize the node title and generated prompt
19 self.node_title = "Prompt Enhancer"
20 self.generated_prompt = ""
21
22 @classmethod
23 def INPUT_TYPES(cls):
24 return {
25 "required": {
26 "prompt": ("STRING",),
27 "seed": ("INT", {"default": 42, "min": 0, "max": 4294967295}), # Default seed, larger range
28 "repetition_penalty": ("FLOAT", {"default": 1.2, "min": 0.0, "max": 10.0}), # Default repetition penalty
29 "max_target_length": ("INT", {"default": 256, "min": 1, "max": 1024}), # Default max target length
30 "temperature": ("FLOAT", {"default": 0.7, "min": 0.0, "max": 1.0}), # Default temperature
31 "top_k": ("INT", {"default": 50, "min": 1, "max": 1000}), # Default top-k
32 "top_p": ("FLOAT", {"default": 0.9, "min": 0.0, "max": 1.0}), # Default top-p
33 },
34 "optional": {
35 "prompts_list": ("LIST",), # List of prompts
36 }
37 }
38
39 RETURN_TYPES = ("STRING",) # Return only one string: the enhanced prompt
40 FUNCTION = "enhance_prompt"
41 CATEGORY = "TextEnhancement"
42
43 def generate_large_seed(self, seed, prompt):
44 # Combine the seed and prompt to create a unique string
45 unique_string = f"{seed}_{prompt}"
46
47 # Use a hash function to generate a large seed
48 hash_object = hashlib.sha256(unique_string.encode())
49 large_seed = int(hash_object.hexdigest(), 16) % (2**32)
50
51 return large_seed
52
53 def enhance_prompt(self, prompt, seed=42, repetition_penalty=1.2, max_target_length=256, temperature=0.7, top_k=50, top_p=0.9, prompts_list=None):
54 # Generate a large seed value
55 large_seed = self.generate_large_seed(seed, prompt)
56
57 # Set random seed for reproducibility
58 torch.manual_seed(large_seed)
59 random.seed(large_seed)
60
61 # Determine the prompts to process
62 prompts = [prompt] if prompts_list is None else prompts_list
63
64 enhanced_prompts = []
65 for p in prompts:
66 # Enhance prompt
67 prefix = "enhance prompt: "
68 input_text = prefix + p
69 input_ids = self.tokenizer(input_text, return_tensors="pt").input_ids.to(self.device)
70
71 # Generate a random seed for this generation
72 random_seed = torch.randint(0, 2**32 - 1, (1,)).item()
73 torch.manual_seed(random_seed)
74 random.seed(random_seed)
75
76 outputs = self.model.generate(
77 input_ids,
78 max_length=max_target_length,
79 num_return_sequences=1,
80 do_sample=True,
81 temperature=temperature,
82 repetition_penalty=repetition_penalty,
83 top_k=top_k,
84 top_p=top_p
85 )
86
87 final_answer = self.tokenizer.decode(outputs[0], skip_special_tokens=True)
88 confidence_score = 1.0 # Default to 1.0 if no score is provided
89
90 # Print the generated prompt and confidence score
91 print(f"Generated Prompt: {final_answer} (Confidence: {confidence_score:.2f})")
92 enhanced_prompts.append((f"Enhanced Prompt: {final_answer}", confidence_score))
93
94 # Update the node title and generated prompt
95 if prompts_list is None:
96 self.node_title = f"Prompt Enhancer (Confidence: {confidence_score:.2f})"
97 self.generated_prompt = f"Enhanced Prompt: {final_answer}"
98 return (f"Enhanced Prompt: {final_answer}",)
99 else:
100 self.node_title = "Prompt Enhancer (Multiple Prompts)"
101 self.generated_prompt = "Multiple Prompts"
102 return enhanced_prompts
103
104 @property
105 def NODE_TITLE(self):
106 return self.node_title
107
108 @property
109 def GENERATED_PROMPT(self):
110 return self.generated_prompt
111
112# A dictionary that contains all nodes you want to export with their names
113NODE_CLASS_MAPPINGS = {
114 "PromptEnhancer": PromptEnhancer
115}
116
117# A dictionary that contains the friendly/humanly readable titles for the nodes
118NODE_DISPLAY_NAME_MAPPINGS = {
119 "PromptEnhancer": "Prompt Enhancer"
120}