def compute_text_embeddings(
self,
prompts: list[str],
device: str | torch.device = "cuda",
) -> torch.Tensor:
"""Compute embeddings for a list of prompts."""
input_ids_batch = []
tok = getattr(self.processor, "tokenizer", None)
if tok is None:
raise RuntimeError("Reason1TextEncoder requires processor.tokenizer")
pad_id = getattr(tok, "pad_id", None)
if pad_id is None:
pad_id = getattr(tok, "pad_token_id", None)
if pad_id is None:
pad_id = getattr(self.model.config, "pad_token_id", None)
if pad_id is None:
pad_id = 0
for prompt in prompts:
conversations = [
{
"role": "system",
"content": [
{
"type": "text",
"text": "You are a helpful assistant who will provide prompts to an image generator.",
}
],
},
{
"role": "user",
"content": [
{
"type": "text",
"text": prompt,
}
],
},
]
try:
tokenizer_output = tok.apply_chat_template(
conversations,
tokenize=True,
add_generation_prompt=False,
add_vision_id=False,
)
except TypeError:
tokenizer_output = tok.apply_chat_template(
conversations,
tokenize=True,
add_generation_prompt=False,
)
if isinstance(tokenizer_output, dict) and "input_ids" in tokenizer_output:
input_ids = tokenizer_output["input_ids"]
if hasattr(input_ids, "tolist"):
input_ids = input_ids.tolist()
else:
input_ids = tokenizer_output
if hasattr(input_ids, "tolist"):
input_ids = input_ids.tolist()
if isinstance(input_ids, list) and len(input_ids) == 1 and isinstance(
input_ids[0], list):
input_ids = input_ids[0]
if not isinstance(input_ids, list):
raise RuntimeError(
f"Unexpected chat_template output type: {type(tokenizer_output)}"
)
if self.num_embedding_padding_tokens > len(input_ids):
pad_len = self.num_embedding_padding_tokens - len(input_ids)
input_ids = input_ids + [pad_id] * pad_len
else:
input_ids = input_ids[:self.num_embedding_padding_tokens]
input_ids = torch.LongTensor(input_ids).to(device=device)
input_ids_batch.append(input_ids)
input_ids_batch = torch.stack(input_ids_batch, dim=0)
# Cosmos2.5 alignment: keep attention_mask=None.
target_device = input_ids_batch.device
try:
embed_device = self.model.model.embed_tokens.weight.device # type: ignore[attr-defined]
except Exception:
embed_device = None
if embed_device is not None and embed_device != target_device:
self.model = self.model.to(target_device)
with torch.no_grad():
position_ids, _ = get_rope_index(
self.model.config,
input_ids_batch,
image_grid_thw=None,
video_grid_thw=None,
second_per_grid_ts=None,
attention_mask=None,
)
position_ids = position_ids.to(target_device)
outputs = self.model.model(
input_ids=input_ids_batch,
position_ids=position_ids,
attention_mask=None,
output_hidden_states=True,
return_dict=True,
use_cache=False,
)
hidden_states = outputs.hidden_states
normalized_hidden_states = []
for layer_idx in range(1, len(hidden_states)):
normalized_state = self._mean_normalize(hidden_states[layer_idx])
normalized_hidden_states.append(normalized_state)
if self.embedding_concat_strategy == "full_concat":
text_embeddings = torch.cat(normalized_hidden_states, dim=-1)
elif self.embedding_concat_strategy == "mean_pooling":
text_embeddings = torch.stack(normalized_hidden_states).mean(dim=0)
elif self.embedding_concat_strategy == "pool_every_n_layers_and_concat":
pooled_embeddings = []
for i in range(0, len(normalized_hidden_states), self.n_layers_per_group):
group = normalized_hidden_states[i : i + self.n_layers_per_group]
pooled = torch.stack(group).mean(dim=0)
pooled_embeddings.append(pooled)
text_embeddings = torch.cat(pooled_embeddings, dim=-1)
else:
raise ValueError(
f"Unknown embedding_concat_strategy: {self.embedding_concat_strategy}"
)
return text_embeddings