Not a member of Pastebin yet?
Sign Up,
it unlocks many cool features!
- import os
- import time
- import torch
- from torch import Tensor
- from transformers import AutoTokenizer, T5Gemma2Model, BatchEncoding
- # Parameters
- prompt = "A category-five hurricane, viewed from inside the eye, reveals a circular stadium of cloud walls rising to fifty thousand feet with an eerie disk of blue sky directly overhead. Shot from a NOAA reconnaissance aircraft mounted camera, the perspective looks outward toward the eyewall — a near-vertical curtain of rotating cloud and lightning that is simultaneously terrifying and transcendent. The inner surface of the eyewall catches the setting sun, painting it in improbable shades of peach and rose. The camera slowly pans 360 degrees to complete one full revolution, capturing the entire coliseum of the storm. Below, the ocean surface is a white blur of foam and spray. The documentary-style cinematography strips away all artifice to present the storm as an entity of pure elemental power."
- negative_prompt = ""
- model_path = "./"
- device = "cuda"
- dtype = torch.bfloat16
- # Load Components
- tokenizer = AutoTokenizer.from_pretrained(model_path, subfolder="tokenizer")
- text_encoder = T5Gemma2Model.from_pretrained(model_path, subfolder="text_encoder", torch_dtype=dtype)
- text_encoder.eval()
- # ==========================================
- # Step 1: Encode Text
- # ==========================================
- def encode_prompt(prompt, tokenizer, text_encoder, device):
- prompt = [prompt] if isinstance(prompt, str) else prompt
- text_inputs = tokenizer(
- prompt,
- padding="max_length",
- max_length=512,
- truncation=True,
- add_special_tokens=True,
- return_attention_mask=True,
- return_tensors="pt",
- )
- #input_ids = text_inputs.input_ids.to(device)
- #attention_mask = text_inputs.attention_mask.to(device)
- text_inputs = BatchEncoding(
- {k: v.to(device) if isinstance(v, torch.Tensor) else v for k, v in text_inputs.items()}
- )
- attention_mask = text_inputs.attention_mask.to(device)
- '''# Encode text sequence
- # T5Gemma2Model bundles encoder and decoder/LM head, while _get_default_embeds expects an encoder-only model (similar to T5EncoderModel/T5GemmaEncoderModel), so use the encoder submodule explicitly here
- outputs = text_encoder.encoder(input_ids=input_ids, attention_mask=attention_mask, output_hidden_states=False, return_dict=True)
- # T5Gemma2Model returns an encoder_last_hidden_state or last_hidden_state depending on version — check both:
- if hasattr(outputs, "encoder_last_hidden_state") and outputs.encoder_last_hidden_state is not None:
- prompt_embeds = outputs.encoder_last_hidden_state
- else:
- prompt_embeds = outputs.last_hidden_state'''
- prompt_embeds = text_encoder.encoder(**text_inputs)[0]
- prompt_embeds = prompt_embeds.to(device=device)
- pooled_prompt_embeds = average_pool(prompt_embeds, attention_mask)
- return prompt_embeds, attention_mask, pooled_prompt_embeds
- def average_pool(last_hidden_states: Tensor, attention_mask: Tensor) -> Tensor:
- last_hidden = last_hidden_states.masked_fill(~attention_mask[..., None].bool(), 0.0)
- denom = attention_mask.sum(dim=1, keepdim=True).clamp(min=1) # avoid div by zero
- return last_hidden.sum(dim=1) / denom
- print("Encoding prompt...")
- text_encoder.to(device)
- with torch.no_grad():
- prompt_embeds, attn_mask, pooled_prompt_embeds = encode_prompt(prompt, tokenizer, text_encoder, device)
- neg_prompt_embeds, neg_attn_mask, neg_pooled_prompt_embeds = encode_prompt(negative_prompt, tokenizer, text_encoder, device)
- text_encoder.to("cpu")
- torch.cuda.empty_cache()
Advertisement