File size: 9,908 Bytes
aa77bd5
 
 
 
 
 
 
 
 
 
 
 
 
 
8296a67
aa77bd5
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
8296a67
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
aa77bd5
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
8296a67
aa77bd5
 
 
 
8296a67
aa77bd5
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
8296a67
 
 
aa77bd5
 
 
 
 
 
 
 
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
114
115
116
117
118
119
120
121
122
123
124
125
126
127
128
129
130
131
132
133
134
135
136
137
138
139
140
141
142
143
144
145
146
147
148
149
150
151
152
153
154
155
156
157
158
159
160
161
162
163
164
165
166
167
168
169
170
171
172
173
174
175
176
177
178
179
180
181
182
183
184
185
186
187
188
189
190
191
192
193
194
195
196
197
198
199
200
201
202
203
204
205
206
207
208
209
210
211
212
213
214
215
216
217
218
219
220
221
222
223
224
225
226
227
228
229
230
"""Keep the official joint head while moving only vocabulary rows over PCIe."""

from __future__ import annotations

import gc
import importlib.util
import json
import sys
from contextlib import contextmanager
from functools import partial

import torch

from .bundle import bundle_manifest, read_bundle
from .core import PREPROCESSING_VERSION, PROFILES


class LexicalRows:
    def __init__(self, weight, *, device, dtype):
        if weight.device.type != "cpu" or weight.is_meta:
            raise ValueError("語彙行列には実体のあるCPUテンソルが必要です。")
        self.weight, self.device, self.dtype = weight, device, dtype

    def __getitem__(self, token_ids):
        shape = (*token_ids.shape, self.weight.shape[1])
        return self.weight.index_select(0, token_ids.reshape(-1).cpu()).reshape(shape).to(self.device, dtype=self.dtype)


class CPUEmbedding(torch.nn.Module):
    def __init__(self, weight, *, device, dtype):
        super().__init__()
        self.rows = LexicalRows(weight, device=device, dtype=dtype)
        self.weight = torch.nn.Parameter(weight, requires_grad=False)

    def forward(self, tokens):
        return self.rows[tokens]


def load_options(profile):
    settings = PROFILES[profile]
    result = {"device_map": {"": 0}}
    if settings["cpu_embeddings"]:
        result["device_map"] = {
            "visual": 0,
            "language_model.embed_tokens": "cpu",
            "language_model.layers": 0,
            "language_model.norm": 0,
            "language_model.rotary_emb": 0,
        }
    if settings["precision"] != "bf16":
        quantization = {
            "llm_int8_skip_modules": ["visual", "lm_head"],
            "llm_int8_enable_fp32_cpu_offload": settings["cpu_embeddings"],
        }
        if settings["precision"] == "int8":
            quantization["load_in_8bit"] = True
        else:
            quantization.update(
                load_in_4bit=True,
                bnb_4bit_quant_type="nf4",
                bnb_4bit_use_double_quant=True,
                bnb_4bit_compute_dtype=torch.bfloat16,
            )
        result["quantization"] = quantization
    return result


def image_options(processor, max_pixels):
    # Transformers 5 applies per-call pixel bounds through size, rather than
    # the legacy max_pixels constructor argument. Never change processor defaults.
    return {
        "images_kwargs": {
            "size": {
                "shortest_edge": min(processor.image_processor.size["shortest_edge"], max_pixels),
                "longest_edge": max_pixels,
            }
        }
    }


def image_measurement(processor, encoded):
    grid = (encoded.media or {}).get("image_grid_thw")
    if grid is None:
        return None, 0
    frames, height, width = map(int, grid[0])
    image_processor = processor.image_processor
    patch, merge = image_processor.patch_size, image_processor.merge_size
    return [width * patch, height * patch], frames * height * width // merge**2


def encode_complete(official, processor, record, max_length):
    # The release silently truncates state. Encode once without truncation and
    # refuse overflow before collating or allocating decoder activations.
    encoded = official.encode_record(processor.tokenizer, record, processor=processor, max_length=1_000_000)
    if len(encoded.input_ids) > max_length:
        raise ValueError(
            f"入力は{len(encoded.input_ids)}トークンです。上限{max_length}を超えます。補足や項目を短くするか処理解像度を下げてください。"
        )
    return encoded


def tensor_from_source(source, key):
    from safetensors import safe_open

    index = json.loads((source / "model.safetensors.index.json").read_text())
    with safe_open(source / index["weight_map"][key], framework="pt", device="cpu", backend="pread") as stream:
        return stream.get_tensor(key)


@contextmanager
def stream_weights():
    import transformers.modeling_utils as loading

    # Windows reserves commit for every copy-on-write mmap, including unread
    # pages. Keep TF's lazy tensor loading without retaining all shard maps.
    original = loading.safe_open
    loading.safe_open = partial(original, backend="pread")
    try:
        yield
    finally:
        loading.safe_open = original


class Runner:
    @classmethod
    def from_directory(cls, directory, profile):
        return cls(None, profile, directory=directory)

    def __init__(self, root, profile, *, directory=None):
        from safetensors.torch import load_file
        from transformers import AutoProcessor, BitsAndBytesConfig, Qwen3_5Model

        if not torch.cuda.is_available() or not torch.cuda.is_bf16_supported():
            raise RuntimeError("ClefにはBF16対応のNVIDIA CUDA GPUが必要です。")
        self.settings, self.profile = PROFILES[profile], profile
        total = torch.cuda.get_device_properties(0).total_memory / 2**30
        if total < self.settings["min_vram_gib"]:
            raise RuntimeError(
                f"{self.settings['label']}のGPUメモリが不足しています。16GB設定またはFlashを選んでください。"
            )
        source, _ = read_bundle(directory, profile) if directory is not None else bundle_manifest(root, profile)
        spec = importlib.util.spec_from_file_location(
            "clef_release_" + self.settings["model"].replace("-", "_"), source / "joint_schema_model.py"
        )
        self.official = importlib.util.module_from_spec(spec)
        sys.modules[spec.name] = self.official
        spec.loader.exec_module(self.official)
        options = load_options(profile)
        saved = source if self.settings["precision"] != "bf16" else None
        if "quantization" in options:
            options["quantization_config"] = BitsAndBytesConfig(**options.pop("quantization"))
        print(f"{self.settings['label']}: {'保存済み' if saved else '公式重み'}を読み込み中", flush=True)  # noqa: T201
        with stream_weights():
            self.backbone, loading = Qwen3_5Model.from_pretrained(
                saved or source,
                dtype=torch.bfloat16,
                local_files_only=True,
                attn_implementation="sdpa",
                output_loading_info=True,
                **options,
            )
        expected_extra = {"lm_head.weight"}
        if loading["missing_keys"] or set(loading["unexpected_keys"]) != expected_extra:
            raise RuntimeError(f"Clef本体のキーが一致しません: {loading}")
        self.backbone.eval()
        self.backbone.config.use_cache = False
        if self.settings["cpu_embeddings"]:
            # Read the bundle's BF16 rows directly; CPU offload may cast to FP32.
            # Discard the old embedding together with its hook. Detaching a
            # hook can temporarily restore the entire vocabulary onto CUDA.
            embedding = tensor_from_source(source, "language_model.embed_tokens.weight")
            self.backbone.language_model.embed_tokens = CPUEmbedding(embedding, device="cuda:0", dtype=torch.bfloat16)
            gc.collect()
        self.lexical = LexicalRows(tensor_from_source(source, "lm_head.weight"), device="cuda:0", dtype=torch.bfloat16)
        self.head = self.official.JointSchemaHead(**json.loads((source / "joint_head_config.json").read_text()))
        self.head.load_state_dict(load_file(source / "joint_head.safetensors", backend="pread"), strict=True)
        self.head = self.head.to(device="cuda:0", dtype=torch.bfloat16).eval()
        self.processor = AutoProcessor.from_pretrained(source, local_files_only=True)
        self.source = source
        self.root, self.saved = root, saved

    @torch.inference_mode()
    def decide(self, request, image=None):
        record = {
            "model": self.settings["model"],
            "state": request["state"],
            "questions": request["questions"],
            "media_kwargs": image_options(self.processor, request["max_pixels"]),
        }
        if image is not None:
            record["images"] = [image]
        encoded = encode_complete(self.official, self.processor, record, request["max_length"])
        processing_size, image_tokens = image_measurement(self.processor, encoded)
        batch = self.official.collate_records([encoded], self.processor.tokenizer.pad_token_id, torch.device("cuda:0"))
        torch.cuda.reset_peak_memory_stats()
        media = batch.get("media") or {}
        model = self.backbone if media else self.backbone.language_model
        outputs = model(
            input_ids=batch["input_ids"],
            attention_mask=batch["attention_mask"],
            use_cache=False,
            return_dict=True,
            **media,
        )
        logits = self.head(
            outputs.last_hidden_state, batch["input_ids"], batch["attention_mask"], batch["records"], self.lexical
        )[0]
        answers = {
            q.question_id: self.official.systemone_answer(
                request["questions"][q.question_id],
                dict(zip(q.option_ids, logit.float().softmax(-1).tolist(), strict=True)),
            )
            for q, logit in zip(encoded.questions, logits, strict=True)
        }
        torch.cuda.synchronize()
        return {
            "model": self.settings["model"],
            "answers": answers,
            "preprocessing": PREPROCESSING_VERSION,
            "processing_size": processing_size,
            "usage": {"input_tokens": len(encoded.input_ids), "image_tokens": image_tokens, "output_tokens": 0},
            "peak_allocated_gib": torch.cuda.max_memory_allocated() / 2**30,
            "peak_reserved_gib": torch.cuda.max_memory_reserved() / 2**30,
        }

    def close(self):
        self.backbone = self.head = self.lexical = self.processor = None
        gc.collect()
        torch.cuda.empty_cache()