File size: 7,539 Bytes
31f7037
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
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
"""Marker based decision model and checkpoint serialization."""
import json
from pathlib import Path

import torch
from torch import nn
from torch.nn import functional as F
from torch.utils.checkpoint import checkpoint
from safetensors.torch import load_file, save_file
from transformers import AutoConfig, AutoModel


class JuliaDecisionModel(nn.Module):
    def __init__(self, encoder, head_layers=2, n_act=2, dropout=.1):
        super().__init__()
        self.encoder = encoder
        width = encoder.config.hidden_size
        self.settings = dict(head_layers=head_layers, n_act=n_act, dropout=dropout)
        layer = nn.TransformerEncoderLayer(width, max(1, width // 64), 4 * width,
                                          dropout, batch_first=True, norm_first=True)
        self.head = nn.TransformerEncoder(layer, head_layers, enable_nested_tensor=False) if head_layers else None
        self.type_emb = nn.Embedding(3, width)
        self.scorer = nn.Sequential(nn.LayerNorm(width), nn.Linear(width, width), nn.GELU(), nn.Linear(width, 1))
        self.act_head = nn.Sequential(nn.Linear(width + 4, 256), nn.GELU(), nn.Linear(256, n_act))
        self.register_buffer('temperature', torch.ones(3))
        self.head_checkpointing = False
        self.marker_only_head = False
        self.encoder.config.reference_compile = False

    @classmethod
    def from_backbone(cls, path, revision=None, head_layers=2):
        encoder = AutoModel.from_pretrained(path, revision=revision, trust_remote_code=False,
                                          attn_implementation='sdpa')
        return cls(encoder, head_layers=head_layers)

    def forward(self, input_ids, attention_mask, marker_pos, marker_mask, qtype, return_actions=False):
        hidden = self.encoder(input_ids=input_ids, attention_mask=attention_mask).last_hidden_state
        hidden = hidden + self.type_emb(qtype)[:, None, :]
        # Only option positions (and CLS for the action head) are consumed.
        # The final block can keep full K/V context while avoiding unused queries
        # and feed-forward outputs. Training retains the original dropout path.
        selected = marker_pos
        if return_actions:
            selected = torch.cat((torch.zeros_like(marker_pos[:, :1]), marker_pos), dim=1)
        sparse = (self.marker_only_head and not self.training and self.head is not None
                  and len(self.head.layers) > 0
                  and type(self.head.layers[-1]) is nn.TransformerEncoderLayer)
        if self.head is not None:
            padding = ~attention_mask.bool()
            for i, layer in enumerate(self.head.layers):
                if sparse and i == len(self.head.layers) - 1:
                    hidden = self._selected_head(layer, hidden, selected, attention_mask)
                elif self.head_checkpointing and self.training:
                    hidden = checkpoint(layer, hidden, src_key_padding_mask=padding, use_reentrant=False)
                else:
                    hidden = layer(hidden, src_key_padding_mask=padding)
        if not sparse:
            positions = selected[:, :, None].expand(-1, -1, hidden.shape[-1])
            hidden = hidden.gather(1, positions)
        markers = hidden[:, 1:] if return_actions else hidden
        scores = self.scorer(markers).squeeze(-1).float()
        scores = scores.masked_fill(~marker_mask, -1e4)
        if not return_actions:
            return scores
        probability = scores.detach().softmax(-1)
        top = probability.topk(min(2, probability.shape[1]), dim=-1).values
        if top.shape[1] == 1:
            top = torch.cat((top, torch.zeros_like(top)), dim=-1)
        count = marker_mask.sum(-1).clamp_min(2).float()
        entropy = -(probability * probability.clamp_min(1e-9).log()).sum(-1) / count.log()
        features = torch.stack((top[:, 0], top[:, 0] - top[:, 1], entropy, count / 255), -1)
        actions = self.act_head(torch.cat((hidden[:, 0].float(), features), -1))
        return scores, actions

    @staticmethod
    def _selected_head(layer, hidden, selected, attention_mask):
        """Exact pre-norm block restricted to selected output positions (eval only)."""
        width = hidden.shape[-1]
        positions = selected[:, :, None].expand(-1, -1, width)
        normalized = layer.norm1(hidden)
        attn = layer.self_attn
        query = F.linear(normalized.gather(1, positions),
                         attn.in_proj_weight[:width], attn.in_proj_bias[:width])
        key, value = F.linear(normalized, attn.in_proj_weight[width:],
                              attn.in_proj_bias[width:]).chunk(2, dim=-1)
        batch = hidden.shape[0]
        heads = attn.num_heads
        split = lambda x: x.reshape(batch, -1, heads, width // heads).transpose(1, 2)
        attended = F.scaled_dot_product_attention(
            split(query), split(key), split(value),
            attn_mask=attention_mask[:, None, None, :].bool(), dropout_p=0.)
        attended = attended.transpose(1, 2).reshape(batch, selected.shape[1], width)
        result = hidden.gather(1, positions) + attn.out_proj(attended)
        return result + layer.linear2(layer.activation(layer.linear1(layer.norm2(result))))

    def save_pretrained(self, directory):
        root = Path(directory)
        root.mkdir(parents=True, exist_ok=True)
        self.encoder.config.save_pretrained(root / 'encoder')
        config = dict(format_version=1, architecture='JuliaDecisionModel',
                      weight_dtype=str(next(self.parameters()).dtype).replace('torch.', ''), **self.settings)
        (root / 'julia_config.json').write_text(json.dumps(config, indent=2) + '\n')
        save_file({k: v.detach().cpu().contiguous() for k, v in self.state_dict().items()},
                  str(root / 'model.safetensors'), metadata={'format': 'pt', 'family': 'julia'})

    @classmethod
    def from_pretrained(cls, directory, *, memory_map=False):
        root = Path(directory)
        if (root / 'julia_config.json').exists():
            config = json.loads((root / 'julia_config.json').read_text())
            if config['format_version'] != 1:
                raise ValueError('Unsupported Julia checkpoint format')
            settings = {k: config[k] for k in ('head_layers', 'n_act', 'dropout')}
        else:
            config = json.loads((root / 'rl_agent_config.json').read_text())
            settings = dict(head_layers=config['head_layers'], n_act=len(config.get('act_costs', {})) + 1)
        encoder_config = AutoConfig.from_pretrained(root / 'encoder', trust_remote_code=False)
        encoder_config.reference_compile = False
        try:
            from transformers.initialization import no_init_weights
        except ImportError:
            from transformers.modeling_utils import no_init_weights
        with no_init_weights():
            encoder = AutoModel.from_config(encoder_config, attn_implementation='sdpa', trust_remote_code=False)
            model = cls(encoder, **settings)
        if not memory_map and config.get('weight_dtype') == 'bfloat16':
            model.to(dtype=torch.bfloat16)
        # Inference can retain safetensors' private file-backed storage rather than
        # copying the entire (mostly unused vocabulary) embedding into anonymous RAM.
        # The default copy path remains available for training and mutable checkpoints.
        model.load_state_dict(load_file(str(root / 'model.safetensors')),
                              strict=True, assign=memory_map)
        return model