File size: 5,923 Bytes
fcf4209
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
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
"""Weight-free GLiNER2 boundary candidate selection between Core ML stages.

The learned start/end projections are outputs of the first Core ML stage. This
module preserves GLiNER2 2.0.0's stable ranking and deduplication on the host.
"""

import math

import torch
from gliner2.models.boundary.constants import MASK_LOGIT
from gliner2.models.boundary.indexing import gather_rows
from gliner2.models.boundary.pool import PooledCandidates, _deduplicate_pool
from gliner2.models.boundary.proposal import select_top_boundaries


def select_candidates(
    start_projection: torch.Tensor,
    end_projection: torch.Tensor,
    boundary_mask: torch.Tensor,
    query_mask: torch.Tensor,
    start_logits: torch.Tensor,
    end_logits: torch.Tensor,
    *,
    boundary_top_k: int,
    pool_size: int,
    min_pool_per_query: int,
) -> PooledCandidates:
    """Select the native shared pool using already projected Core ML states."""
    batch, n_boundaries, dim = start_projection.shape
    n_queries = query_mask.shape[1]
    if end_projection.shape != start_projection.shape:
        raise ValueError("Start and end projections must have the same shape")
    if start_logits.shape != (batch, n_queries, n_boundaries):
        raise ValueError("Start logits do not match boundary and query dimensions")
    if end_logits.shape != start_logits.shape:
        raise ValueError("End logits do not match start logits")
    if boundary_mask.shape != (batch, n_boundaries) or query_mask.shape != (batch, n_queries):
        raise ValueError("Boundary or query mask has an unexpected shape")

    floor = torch.full_like(start_logits, MASK_LOGIT)
    valid_boundary = boundary_mask.unsqueeze(1) & query_mask.unsqueeze(-1)
    union_start = torch.where(valid_boundary, start_logits, floor).amax(1)
    union_end = torch.where(valid_boundary, end_logits, floor).amax(1)
    union_valid = boundary_mask & query_mask.any(-1, keepdim=True)
    _, starts, starts_valid = select_top_boundaries(union_start.unsqueeze(1), union_valid.unsqueeze(1), boundary_top_k)
    _, ends, ends_valid = select_top_boundaries(union_end.unsqueeze(1), union_valid.unsqueeze(1), boundary_top_k)
    starts, ends = starts[:, 0], ends[:, 0]
    starts_valid, ends_valid = starts_valid[:, 0], ends_valid[:, 0]
    n_starts, n_ends = starts.shape[1], ends.shape[1]
    pair_start = starts.unsqueeze(-1).expand(batch, n_starts, n_ends).reshape(batch, -1)
    pair_end = ends.unsqueeze(1).expand(batch, n_starts, n_ends).reshape(batch, -1)
    pair_valid = (
        starts_valid.unsqueeze(-1) & ends_valid.unsqueeze(1) & (ends.unsqueeze(1) > starts.unsqueeze(-1))
    ).reshape(batch, -1)

    selected_start = gather_rows(start_projection, pair_start)
    selected_end = gather_rows(end_projection, pair_end)
    compatibility = (selected_start * selected_end).sum(-1) / math.sqrt(dim)
    union_pair_score = (
        compatibility
        + union_start.gather(1, pair_start.clamp(0, n_boundaries - 1))
        + union_end.gather(1, pair_end.clamp(0, n_boundaries - 1))
    )

    quota = min(min_pool_per_query, pair_start.shape[-1])
    quota_keys = pair_start.new_zeros((batch, 0))
    quota_scores = union_pair_score.new_zeros((batch, 0))
    quota_valid = pair_valid.new_zeros((batch, 0))
    if quota:
        start_idx = pair_start.clamp(0, start_logits.shape[2] - 1).unsqueeze(1).expand(batch, n_queries, -1)
        end_idx = pair_end.clamp(0, end_logits.shape[2] - 1).unsqueeze(1).expand(batch, n_queries, -1)
        per_query = start_logits.gather(2, start_idx) + end_logits.gather(2, end_idx) + compatibility.unsqueeze(1)
        per_query_valid = pair_valid.unsqueeze(1) & query_mask.unsqueeze(-1)
        ranked = torch.argsort(
            per_query.masked_fill(~per_query_valid, MASK_LOGIT),
            dim=-1,
            descending=True,
            stable=True,
        )[..., :quota]
        quota_start = start_idx.gather(-1, ranked)
        quota_end = end_idx.gather(-1, ranked)
        quota_valid = per_query_valid.gather(-1, ranked).reshape(batch, -1)
        quota_keys = (quota_start * n_boundaries + quota_end).reshape(batch, -1)
        rank_bonus = torch.arange(quota, 0, -1, device=start_projection.device, dtype=union_pair_score.dtype)
        quota_scores = (
            union_pair_score.new_full((batch, n_queries, quota), -MASK_LOGIT * 0.5) + rank_bonus.view(1, 1, quota)
        ).reshape(batch, -1)

    global_keys = pair_start * n_boundaries + pair_end
    all_keys = torch.cat((quota_keys, global_keys), -1)
    all_scores = torch.cat((quota_scores, union_pair_score.detach()), -1)
    all_valid = torch.cat((quota_valid, pair_valid), -1)
    selected_keys, selected_valid = _deduplicate_pool(all_keys, all_scores, all_valid, pool_size, n_boundaries)
    selected_keys = torch.where(selected_valid, selected_keys, torch.zeros_like(selected_keys))
    selected_s = torch.div(selected_keys, n_boundaries, rounding_mode="floor")
    selected_e = selected_keys - selected_s * n_boundaries
    indices = torch.stack((selected_s, selected_e), -1)
    indices = torch.where(selected_valid.unsqueeze(-1), indices, torch.zeros_like(indices))

    gathered_start = gather_rows(start_projection, selected_s)
    gathered_end = gather_rows(end_projection, selected_e)
    selected_compat = (gathered_start * gathered_end).sum(-1) / math.sqrt(dim)
    selected_score = (
        selected_compat
        + union_start.gather(1, selected_s.clamp(0, n_boundaries - 1))
        + union_end.gather(1, selected_e.clamp(0, n_boundaries - 1))
    )
    selected_score = selected_score.masked_fill(~selected_valid, MASK_LOGIT)
    selected_compat = torch.where(selected_valid, selected_compat, torch.zeros_like(selected_compat))
    return PooledCandidates(
        indices=indices,
        mask=selected_valid,
        proposal_logits=selected_score,
        gold_mask=None,
        compat_logits=selected_compat,
        stats=None,
    )