Skip to content

reason1

Reason1 (Qwen2.5-VL) text encoder.

Classes

fastvideo.models.encoders.reason1.Reason1TextEncoder

Reason1TextEncoder(config: Reason1Config, prefix: str = '', checkpoint_path: str | None = None)

Bases: TextEncoder

Reason1 (Qwen2.5-VL) text encoder.

Source code in fastvideo/models/encoders/reason1.py
def __init__(self, config: Reason1Config, prefix: str = "", checkpoint_path: str | None = None):
    super().__init__(config)

    self.prefix = prefix
    self.quant_config = None  # For future quantization support

    self.embedding_concat_strategy = config.arch_config.embedding_concat_strategy
    self.n_layers_per_group = config.arch_config.n_layers_per_group
    self.num_embedding_padding_tokens = config.arch_config.num_embedding_padding_tokens

    config_path = checkpoint_path if checkpoint_path else config.tokenizer_type

    logger.info("Initializing Reason1TextEncoder (Qwen2.5-VL) from %s", config_path)
    try:
        from transformers import AutoConfig as HFAutoConfig
        hf_config = HFAutoConfig.from_pretrained(
            config_path,
            trust_remote_code=True,
        )
    except Exception as e:
        logger.warning("Failed to load HF config from %s (%s). Using default Qwen2.5-VL-7B config.",
                       config_path, e)
        hf_config = Qwen2_5_VLConfig(
                hidden_size=3584,
                intermediate_size=18944,
                max_window_layers=28,
                num_attention_heads=28,
                num_hidden_layers=28,
                num_key_value_heads=4,
                tie_word_embeddings=False,
                vocab_size=152064,
        )

    hf_config.output_hidden_states = True

    if hasattr(config.arch_config, '_attn_implementation') and config.arch_config._attn_implementation:
        hf_config._attn_implementation = config.arch_config._attn_implementation
    else:
        hf_config._attn_implementation = "flash_attention_2"
    logger.info("Reason1 attention implementation: %s", getattr(hf_config, "_attn_implementation", None))

    with torch.device("meta"):
        self.model = Qwen2_5_VLForConditionalGenerationSimple(hf_config)

    self.processor = AutoProcessor.from_pretrained(
        config_path,
        trust_remote_code=True,
    )

    weights_override = os.getenv("FASTVIDEO_REASON1_WEIGHTS_PATH")
    if weights_override:
        self.secondary_weights = (
            _WeightsSource(
                model_or_path=weights_override,
                prefix="",
                fall_back_to_pt=True,
                allow_patterns_overrides=None,
            ),
        )
        logger.info("Reason1TextEncoder: overlaying weights from %s", weights_override)

    self._weights_loaded = False

Methods:

fastvideo.models.encoders.reason1.Reason1TextEncoder.compute_text_embeddings
compute_text_embeddings(prompts: list[str], device: str | device = 'cuda') -> Tensor

Compute embeddings for a list of prompts.

Source code in fastvideo/models/encoders/reason1.py
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

Functions: