File size: 4,097 Bytes
ad6f7a0
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
1
2
3
4
5
6
7
8
9
10
11
12
13
14
15
16
17
18
19
20
21
22
23
24
25
26
27
28
29
30
31
32
33
34
35
36
37
38
39
40
41
42
43
44
45
46
47
48
49
50
51
52
53
54
55
56
57
58
59
60
61
62
63
64
65
66
67
68
69
70
71
72
73
74
75
76
77
78
79
80
81
82
83
84
85
86
87
88
89
90
91
92
93
94
95
96
97
98
99
100
101
102
103
104
105
106
107
108
109
110
111
112
113
from typing import Any, List, Optional
import os

import torch
import torch.nn.functional as F

from huggingface_hub import snapshot_download
from transformers import AutoTokenizer
from transformers.modeling_utils import PreTrainedModel
from transformers.models.qwen3 import Qwen3Config, Qwen3Model
from peft import PeftMixedModel, PeftConfig

class JinaEmbeddingsV5Model(PeftMixedModel):
    @classmethod
    def register_for_auto_class(cls, auto_class="AutoModel"):
        return PreTrainedModel.register_for_auto_class.__func__(cls, auto_class)
    
    @classmethod
    def from_pretrained(cls, pretrained_model_name_or_path: str, *args, **kwargs):
        base_config = Qwen3Config.from_pretrained(
            pretrained_model_name_or_path,
        )
        base_model = Qwen3Model.from_pretrained(
            pretrained_model_name_or_path,
            config=base_config,
            attn_implementation='flash_attention_2',
            dtype=torch.bfloat16,
        )
        kwargs = dict[str, Any](kwargs)
        kwargs.pop("config", None)
        if os.path.isdir(base_model.name_or_path):
            adapters_dir = os.path.join(base_model.name_or_path, "adapters")
        else:
            adapter_cache_path = snapshot_download(
                repo_id=base_model.name_or_path,
                allow_patterns=["adapters/*"],
            )
            adapters_dir = os.path.join(adapter_cache_path, "adapters")
        adapter_names = ["retrieval", "text-matching", "classification", "clustering"]

        adapter_paths = {
            name: os.path.join(adapters_dir, name)
            for name in adapter_names
        }

        peft_config = PeftConfig.from_pretrained(adapter_paths["retrieval"], **kwargs)
        model = cls(base_model, peft_config, adapter_name="retrieval")
        model._pretrained_path = pretrained_model_name_or_path
        for adapter_name in adapter_names:
            model.load_adapter(
                adapter_paths[adapter_name],
                adapter_name=adapter_name,
                **kwargs,
            )

        model.tokenizer = AutoTokenizer.from_pretrained(
            pretrained_model_name_or_path,
            trust_remote_code=True,
        )
        return model

    def encode(
        self,
        texts: List[str],
        task: str,
        prompt_name: Optional[str] = "document",
        truncate_dim: Optional[int] = None,
        max_length: Optional[int] = None,
    ) -> List[torch.Tensor]:
        if task not in {"retrieval", "classification", "text-matching", "clustering"}:
            raise ValueError(f"Unknown task: {task}")

        if prompt_name is None:
            prompt_name = "document"
        if prompt_name not in {"query", "document"}:
            raise ValueError(f"Unknown prompt_name: {prompt_name}")

        prefix = "Query: " if prompt_name == "query" else "Document: "
        inputs = [f"{prefix}{text}" for text in texts]

        if not hasattr(self, "tokenizer") or self.tokenizer is None:
            raise ValueError("Tokenizer not found on model. Load with from_pretrained().")

        batch = self.tokenizer(
            inputs,
            return_tensors="pt",
            padding=True,
            truncation=True,
            max_length=max_length,
        )
        device = next(self.parameters()).device
        batch = {k: v.to(device) for k, v in batch.items()}
        print(batch['input_ids'])
        self.set_adapter([task])
        with torch.no_grad():
            outputs = self(**batch)
            hidden = outputs.last_hidden_state
            mask = batch.get("attention_mask")
            if mask is None:
                pooled = hidden[:, -1]
            else:
                sequence_lengths = mask.sum(dim=1) - 1
                pooled = hidden[
                    torch.arange(hidden.shape[0], device=hidden.device),
                    sequence_lengths,
                ]

            if truncate_dim is not None:
                pooled = pooled[:, :truncate_dim]
            embeddings = F.normalize(pooled, p=2, dim=-1)

        return embeddings