Skip to content

Ltnet

Transformers-style LT-Net package.

LTNetConfig

Bases: PretrainedConfig

Architecture and processor metadata for LT-Net checkpoints.

Parameters:

Name Type Description Default
dataset_name str

Dataset slug for the converted checkpoint.

'coco'
vocab_size int

Mixed special/object/predicate token vocabulary size.

206
obj_classes_size int

Object-id embedding/classifier vocabulary size.

155
hidden_size int

Transformer hidden dimension.

256
num_hidden_layers int

Number of relation encoder layers.

4
num_attention_heads int

Number of relation encoder attention heads.

4
dropout float

Dropout used by embeddings, encoder, and bbox heads.

0.1
enable_noise bool

Whether the original config enabled relation noise.

False
noise_size int

Original relation-noise size.

64
decoder_head_type DecoderHeadType | str

GMM or Linear bbox head.

gmm
decoder_box_loss BoxLossType | str

PDF or Reg objective family.

pdf
decoder_schedule_sample bool

Whether training used scheduled sampling.

False
decoder_two_path bool

Whether the original config enabled two-path decoding.

False
decoder_global_feature bool

Whether to concatenate max-pooled global features.

True
decoder_greedy bool

Whether inference uses GMM means instead of sampling.

True
xy_temperature float

GMM mixture temperature for center coordinates.

1.0
wh_temperature float

GMM mixture temperature for box size.

1.0
refine bool

Whether a refinement head is present.

False
refine_head_type DecoderHeadType | str

Refinement bbox head type.

linear
refine_box_loss BoxLossType | str

Refinement objective family.

reg
refine_x_softmax bool

Whether the decoder records XY PDF scores for refine.

True
max_sequence_length int

Processor/model sequence length.

128
id2label Mapping[int, str] | Mapping[str, str] | None

Public object label mapping.

None
relation_id2label Mapping[int, str] | Mapping[str, str] | None

Public relation label mapping.

None
model_type str | None

Ignored compatibility field from serialized configs.

None
transformers_version str | None

Ignored compatibility field.

None
kwargs str | int | float | bool | None

Additional PretrainedConfig fields.

{}

Examples:

>>> config = LTNetConfig(hidden_size=32, num_attention_heads=4)
>>> config.model_type
'ltnet'
Source code in models/ltnet/src/ltnet/configuration_ltnet.py
 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
class LTNetConfig(PretrainedConfig):
    """Architecture and processor metadata for LT-Net checkpoints.

    Args:
        dataset_name: Dataset slug for the converted checkpoint.
        vocab_size: Mixed special/object/predicate token vocabulary size.
        obj_classes_size: Object-id embedding/classifier vocabulary size.
        hidden_size: Transformer hidden dimension.
        num_hidden_layers: Number of relation encoder layers.
        num_attention_heads: Number of relation encoder attention heads.
        dropout: Dropout used by embeddings, encoder, and bbox heads.
        enable_noise: Whether the original config enabled relation noise.
        noise_size: Original relation-noise size.
        decoder_head_type: ``GMM`` or ``Linear`` bbox head.
        decoder_box_loss: ``PDF`` or ``Reg`` objective family.
        decoder_schedule_sample: Whether training used scheduled sampling.
        decoder_two_path: Whether the original config enabled two-path decoding.
        decoder_global_feature: Whether to concatenate max-pooled global features.
        decoder_greedy: Whether inference uses GMM means instead of sampling.
        xy_temperature: GMM mixture temperature for center coordinates.
        wh_temperature: GMM mixture temperature for box size.
        refine: Whether a refinement head is present.
        refine_head_type: Refinement bbox head type.
        refine_box_loss: Refinement objective family.
        refine_x_softmax: Whether the decoder records XY PDF scores for refine.
        max_sequence_length: Processor/model sequence length.
        id2label: Public object label mapping.
        relation_id2label: Public relation label mapping.
        model_type: Ignored compatibility field from serialized configs.
        transformers_version: Ignored compatibility field.
        kwargs: Additional ``PretrainedConfig`` fields.

    Examples:
        >>> config = LTNetConfig(hidden_size=32, num_attention_heads=4)
        >>> config.model_type
        'ltnet'
    """

    model_type = "ltnet"

    def __init__(
        self,
        *,
        dataset_name: str = "coco",
        vocab_size: int = 206,
        obj_classes_size: int = 155,
        hidden_size: int = 256,
        num_hidden_layers: int = 4,
        num_attention_heads: int = 4,
        dropout: float = 0.1,
        enable_noise: bool = False,
        noise_size: int = 64,
        decoder_head_type: DecoderHeadType | str = DecoderHeadType.gmm,
        decoder_box_loss: BoxLossType | str = BoxLossType.pdf,
        decoder_schedule_sample: bool = False,
        decoder_two_path: bool = False,
        decoder_global_feature: bool = True,
        decoder_greedy: bool = True,
        xy_temperature: float = 1.0,
        wh_temperature: float = 1.0,
        refine: bool = False,
        refine_head_type: DecoderHeadType | str = DecoderHeadType.linear,
        refine_box_loss: BoxLossType | str = BoxLossType.reg,
        refine_x_softmax: bool = True,
        max_sequence_length: int = 128,
        id2label: Mapping[int, str] | Mapping[str, str] | None = None,
        relation_id2label: Mapping[int, str] | Mapping[str, str] | None = None,
        bos_token_id: int = 1,
        eos_token_id: int = 2,
        pad_token_id: int = 0,
        mask_token_id: int = 3,
        model_type: str | None = None,
        transformers_version: str | None = None,
        **kwargs: str | int | float | bool | None,
    ) -> None:
        """Initialize LT-Net architecture and metadata fields."""
        _ = (model_type, transformers_version)

        self.dataset_name = dataset_name
        self.vocab_size = vocab_size
        self.obj_classes_size = obj_classes_size
        self.hidden_size = hidden_size
        self.num_hidden_layers = num_hidden_layers
        self.num_attention_heads = num_attention_heads
        self.dropout = dropout
        self.enable_noise = enable_noise
        self.noise_size = noise_size

        self.decoder_head_type = str(_normalize_head_type(decoder_head_type))
        self.decoder_box_loss = str(_normalize_box_loss(decoder_box_loss))
        self.decoder_schedule_sample = decoder_schedule_sample
        self.decoder_two_path = decoder_two_path
        self.decoder_global_feature = decoder_global_feature
        self.decoder_greedy = decoder_greedy
        self.xy_temperature = xy_temperature
        self.wh_temperature = wh_temperature
        self.refine = refine

        self.refine_head_type = str(_normalize_head_type(refine_head_type))
        self.refine_box_loss = str(_normalize_box_loss(refine_box_loss))
        self.refine_x_softmax = refine_x_softmax
        self.max_sequence_length = max_sequence_length
        self.mask_token_id = mask_token_id
        _ = kwargs

        super().__init__()
        normalized_id2label = id2label or DEFAULT_ID2LABEL
        normalized_relation_id2label = relation_id2label or DEFAULT_RELATION_ID2LABEL
        self.id2label = {
            int(key): str(value) for key, value in normalized_id2label.items()
        }
        self.relation_id2label = {
            int(key): str(value) for key, value in normalized_relation_id2label.items()
        }
        self.bos_token_id = bos_token_id
        self.eos_token_id = eos_token_id
        self.pad_token_id = pad_token_id
        self.label2id = {value: key for key, value in self.id2label.items()}

__init__

__init__(
    *,
    dataset_name: str = "coco",
    vocab_size: int = 206,
    obj_classes_size: int = 155,
    hidden_size: int = 256,
    num_hidden_layers: int = 4,
    num_attention_heads: int = 4,
    dropout: float = 0.1,
    enable_noise: bool = False,
    noise_size: int = 64,
    decoder_head_type: DecoderHeadType
    | str = DecoderHeadType.gmm,
    decoder_box_loss: BoxLossType | str = BoxLossType.pdf,
    decoder_schedule_sample: bool = False,
    decoder_two_path: bool = False,
    decoder_global_feature: bool = True,
    decoder_greedy: bool = True,
    xy_temperature: float = 1.0,
    wh_temperature: float = 1.0,
    refine: bool = False,
    refine_head_type: DecoderHeadType
    | str = DecoderHeadType.linear,
    refine_box_loss: BoxLossType | str = BoxLossType.reg,
    refine_x_softmax: bool = True,
    max_sequence_length: int = 128,
    id2label: Mapping[int, str]
    | Mapping[str, str]
    | None = None,
    relation_id2label: Mapping[int, str]
    | Mapping[str, str]
    | None = None,
    bos_token_id: int = 1,
    eos_token_id: int = 2,
    pad_token_id: int = 0,
    mask_token_id: int = 3,
    model_type: str | None = None,
    transformers_version: str | None = None,
    **kwargs: str | int | float | bool | None,
) -> None

Initialize LT-Net architecture and metadata fields.

Source code in models/ltnet/src/ltnet/configuration_ltnet.py
 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
def __init__(
    self,
    *,
    dataset_name: str = "coco",
    vocab_size: int = 206,
    obj_classes_size: int = 155,
    hidden_size: int = 256,
    num_hidden_layers: int = 4,
    num_attention_heads: int = 4,
    dropout: float = 0.1,
    enable_noise: bool = False,
    noise_size: int = 64,
    decoder_head_type: DecoderHeadType | str = DecoderHeadType.gmm,
    decoder_box_loss: BoxLossType | str = BoxLossType.pdf,
    decoder_schedule_sample: bool = False,
    decoder_two_path: bool = False,
    decoder_global_feature: bool = True,
    decoder_greedy: bool = True,
    xy_temperature: float = 1.0,
    wh_temperature: float = 1.0,
    refine: bool = False,
    refine_head_type: DecoderHeadType | str = DecoderHeadType.linear,
    refine_box_loss: BoxLossType | str = BoxLossType.reg,
    refine_x_softmax: bool = True,
    max_sequence_length: int = 128,
    id2label: Mapping[int, str] | Mapping[str, str] | None = None,
    relation_id2label: Mapping[int, str] | Mapping[str, str] | None = None,
    bos_token_id: int = 1,
    eos_token_id: int = 2,
    pad_token_id: int = 0,
    mask_token_id: int = 3,
    model_type: str | None = None,
    transformers_version: str | None = None,
    **kwargs: str | int | float | bool | None,
) -> None:
    """Initialize LT-Net architecture and metadata fields."""
    _ = (model_type, transformers_version)

    self.dataset_name = dataset_name
    self.vocab_size = vocab_size
    self.obj_classes_size = obj_classes_size
    self.hidden_size = hidden_size
    self.num_hidden_layers = num_hidden_layers
    self.num_attention_heads = num_attention_heads
    self.dropout = dropout
    self.enable_noise = enable_noise
    self.noise_size = noise_size

    self.decoder_head_type = str(_normalize_head_type(decoder_head_type))
    self.decoder_box_loss = str(_normalize_box_loss(decoder_box_loss))
    self.decoder_schedule_sample = decoder_schedule_sample
    self.decoder_two_path = decoder_two_path
    self.decoder_global_feature = decoder_global_feature
    self.decoder_greedy = decoder_greedy
    self.xy_temperature = xy_temperature
    self.wh_temperature = wh_temperature
    self.refine = refine

    self.refine_head_type = str(_normalize_head_type(refine_head_type))
    self.refine_box_loss = str(_normalize_box_loss(refine_box_loss))
    self.refine_x_softmax = refine_x_softmax
    self.max_sequence_length = max_sequence_length
    self.mask_token_id = mask_token_id
    _ = kwargs

    super().__init__()
    normalized_id2label = id2label or DEFAULT_ID2LABEL
    normalized_relation_id2label = relation_id2label or DEFAULT_RELATION_ID2LABEL
    self.id2label = {
        int(key): str(value) for key, value in normalized_id2label.items()
    }
    self.relation_id2label = {
        int(key): str(value) for key, value in normalized_relation_id2label.items()
    }
    self.bos_token_id = bos_token_id
    self.eos_token_id = eos_token_id
    self.pad_token_id = pad_token_id
    self.label2id = {value: key for key, value in self.id2label.items()}

LTNetForLayoutGeneration

Bases: PreTrainedModel

Transformers PreTrainedModel for LT-Net relation-to-layout inference.

Source code in models/ltnet/src/ltnet/modeling_ltnet.py
 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
class LTNetForLayoutGeneration(PreTrainedModel):
    """Transformers ``PreTrainedModel`` for LT-Net relation-to-layout inference."""

    config_class = LTNetConfig
    base_model_prefix = "ltnet"
    main_input_name = "input_token"
    _tied_weights_keys: dict[str, str] = {}

    def __init__(self, config: LTNetConfig) -> None:
        """Initialize relation encoder and bbox head."""
        super().__init__(config)
        self.encoder = RelEncoder(config)
        self.bbox_head = BBoxHead(config)
        self.all_tied_weights_keys = dict(self._tied_weights_keys)

    def forward(
        self,
        input_token: Int[torch.Tensor, "batch sequence"],
        input_obj_id: Int[torch.Tensor, "batch sequence"],
        segment_label: Int[torch.Tensor, "batch sequence"],
        token_type: Int[torch.Tensor, "batch sequence"],
        src_mask: Bool[torch.Tensor, "batch 1 sequence"] | None = None,
        global_mask: Bool[torch.Tensor, "batch sequence"] | None = None,
        bbox: Float[torch.Tensor, "batch sequence 4"] | None = None,
        bbox_mask: Bool[torch.Tensor, "batch sequence"] | None = None,
        inference: bool = False,
        generator: torch.Generator | None = None,
        output_hidden_states: bool = False,
        return_dict: bool = True,
    ) -> (
        LTNetModelOutput
        | tuple[Float[torch.Tensor, "batch sequence feature"] | None, ...]
    ):
        """Run LT-Net relation encoding and bbox prediction.

        Args:
            input_token: Mixed object/predicate token ids.
            input_obj_id: Stable object ids for object-token positions.
            segment_label: Relation segment ids.
            token_type: Token type ids ``0/1/2/3``.
            src_mask: Valid-token mask shaped ``(batch, 1, sequence)``.
            global_mask: Optional reference global-feature mask.
            bbox: Optional teacher-forced boxes for training/parity paths.
            bbox_mask: Optional box validity mask, reserved for parity paths.
            inference: Compatibility flag; greedy decoding is pipeline-owned.
            generator: Optional PyTorch generator for stochastic GMM sampling.
            output_hidden_states: Whether to include encoder hidden states.
            return_dict: Whether to return a ``ModelOutput``.

        Returns:
            Raw LT-Net output dataclass or tuple.

        Raises:
            ValueError: If required tensor shapes are invalid.
        """
        _ = bbox_mask
        if input_token.shape != input_obj_id.shape:
            raise ValueError("input_token and input_obj_id must have the same shape")

        effective_src_mask = (
            src_mask
            if src_mask is not None
            else input_token.ne(self.config.pad_token_id).unsqueeze(1)
        )
        if effective_src_mask.ndim == 2:
            effective_src_mask = effective_src_mask.unsqueeze(1)
        encoder_outputs = self.encoder(
            input_token,
            input_obj_id,
            segment_label,
            token_type,
            effective_src_mask,
        )
        (
            hidden_states,
            vocab_logits,
            obj_id_logits,
            token_type_logits,
            src,
            class_embeds,
        ) = encoder_outputs
        effective_global_mask = (
            global_mask if global_mask is not None else input_token.ge(2)
        )
        if inference:
            coarse_box, coarse_gmm, refine_box, refine_gmm = self.bbox_head.inference(
                hidden_states,
                effective_src_mask,
                src,
                class_embeds,
                effective_global_mask,
                generator=generator,
            )
        else:
            if bbox is None:
                bbox = input_token.new_full(
                    (input_token.size(0), input_token.size(1) - 1, 4), 2.0
                ).float()
            elif bbox.size(1) == input_token.size(1):
                bbox = bbox[:, :-1, :]
            trg_mask = effective_src_mask.new_ones([1, 1, 1])
            coarse_box, coarse_gmm, refine_box, refine_gmm = self.bbox_head(
                0,
                hidden_states,
                effective_src_mask,
                src,
                class_embeds,
                bbox,
                trg_mask,
                effective_global_mask,
                generator=generator,
            )
        output = LTNetModelOutput(
            vocab_logits=vocab_logits,
            obj_id_logits=obj_id_logits,
            token_type_logits=token_type_logits,
            coarse_box=coarse_box,
            coarse_gmm=coarse_gmm,
            refine_box=refine_box,
            refine_gmm=refine_gmm,
            hidden_states=hidden_states if output_hidden_states else None,
        )
        if return_dict:
            return output
        return output.to_tuple()

    @torch.no_grad()
    def _generate_boxes(
        self,
        input_token: Int[torch.Tensor, "batch sequence"],
        input_obj_id: Int[torch.Tensor, "batch sequence"],
        segment_label: Int[torch.Tensor, "batch sequence"],
        token_type: Int[torch.Tensor, "batch sequence"],
        src_mask: Bool[torch.Tensor, "batch 1 sequence"] | None = None,
        global_mask: Bool[torch.Tensor, "batch sequence"] | None = None,
        generator: torch.Generator | None = None,
    ) -> LTNetModelOutput:
        """Private pipeline helper for layout-level generation."""
        output = self(
            input_token=input_token,
            input_obj_id=input_obj_id,
            segment_label=segment_label,
            token_type=token_type,
            src_mask=src_mask,
            global_mask=global_mask,
            inference=True,
            generator=generator,
            return_dict=True,
        )
        return cast(LTNetModelOutput, output)

__init__

__init__(config: LTNetConfig) -> None

Initialize relation encoder and bbox head.

Source code in models/ltnet/src/ltnet/modeling_ltnet.py
61
62
63
64
65
66
def __init__(self, config: LTNetConfig) -> None:
    """Initialize relation encoder and bbox head."""
    super().__init__(config)
    self.encoder = RelEncoder(config)
    self.bbox_head = BBoxHead(config)
    self.all_tied_weights_keys = dict(self._tied_weights_keys)

forward

forward(
    input_token: Int[Tensor, "batch sequence"],
    input_obj_id: Int[Tensor, "batch sequence"],
    segment_label: Int[Tensor, "batch sequence"],
    token_type: Int[Tensor, "batch sequence"],
    src_mask: Bool[Tensor, "batch 1 sequence"]
    | None = None,
    global_mask: Bool[Tensor, "batch sequence"]
    | None = None,
    bbox: Float[Tensor, "batch sequence 4"] | None = None,
    bbox_mask: Bool[Tensor, "batch sequence"] | None = None,
    inference: bool = False,
    generator: Generator | None = None,
    output_hidden_states: bool = False,
    return_dict: bool = True,
) -> (
    LTNetModelOutput
    | tuple[
        Float[torch.Tensor, "batch sequence feature"]
        | None,
        ...,
    ]
)

Run LT-Net relation encoding and bbox prediction.

Parameters:

Name Type Description Default
input_token Int[Tensor, 'batch sequence']

Mixed object/predicate token ids.

required
input_obj_id Int[Tensor, 'batch sequence']

Stable object ids for object-token positions.

required
segment_label Int[Tensor, 'batch sequence']

Relation segment ids.

required
token_type Int[Tensor, 'batch sequence']

Token type ids 0/1/2/3.

required
src_mask Bool[Tensor, 'batch 1 sequence'] | None

Valid-token mask shaped (batch, 1, sequence).

None
global_mask Bool[Tensor, 'batch sequence'] | None

Optional reference global-feature mask.

None
bbox Float[Tensor, 'batch sequence 4'] | None

Optional teacher-forced boxes for training/parity paths.

None
bbox_mask Bool[Tensor, 'batch sequence'] | None

Optional box validity mask, reserved for parity paths.

None
inference bool

Compatibility flag; greedy decoding is pipeline-owned.

False
generator Generator | None

Optional PyTorch generator for stochastic GMM sampling.

None
output_hidden_states bool

Whether to include encoder hidden states.

False
return_dict bool

Whether to return a ModelOutput.

True

Returns:

Type Description
LTNetModelOutput | tuple[Float[Tensor, 'batch sequence feature'] | None, ...]

Raw LT-Net output dataclass or tuple.

Raises:

Type Description
ValueError

If required tensor shapes are invalid.

Source code in models/ltnet/src/ltnet/modeling_ltnet.py
 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
def forward(
    self,
    input_token: Int[torch.Tensor, "batch sequence"],
    input_obj_id: Int[torch.Tensor, "batch sequence"],
    segment_label: Int[torch.Tensor, "batch sequence"],
    token_type: Int[torch.Tensor, "batch sequence"],
    src_mask: Bool[torch.Tensor, "batch 1 sequence"] | None = None,
    global_mask: Bool[torch.Tensor, "batch sequence"] | None = None,
    bbox: Float[torch.Tensor, "batch sequence 4"] | None = None,
    bbox_mask: Bool[torch.Tensor, "batch sequence"] | None = None,
    inference: bool = False,
    generator: torch.Generator | None = None,
    output_hidden_states: bool = False,
    return_dict: bool = True,
) -> (
    LTNetModelOutput
    | tuple[Float[torch.Tensor, "batch sequence feature"] | None, ...]
):
    """Run LT-Net relation encoding and bbox prediction.

    Args:
        input_token: Mixed object/predicate token ids.
        input_obj_id: Stable object ids for object-token positions.
        segment_label: Relation segment ids.
        token_type: Token type ids ``0/1/2/3``.
        src_mask: Valid-token mask shaped ``(batch, 1, sequence)``.
        global_mask: Optional reference global-feature mask.
        bbox: Optional teacher-forced boxes for training/parity paths.
        bbox_mask: Optional box validity mask, reserved for parity paths.
        inference: Compatibility flag; greedy decoding is pipeline-owned.
        generator: Optional PyTorch generator for stochastic GMM sampling.
        output_hidden_states: Whether to include encoder hidden states.
        return_dict: Whether to return a ``ModelOutput``.

    Returns:
        Raw LT-Net output dataclass or tuple.

    Raises:
        ValueError: If required tensor shapes are invalid.
    """
    _ = bbox_mask
    if input_token.shape != input_obj_id.shape:
        raise ValueError("input_token and input_obj_id must have the same shape")

    effective_src_mask = (
        src_mask
        if src_mask is not None
        else input_token.ne(self.config.pad_token_id).unsqueeze(1)
    )
    if effective_src_mask.ndim == 2:
        effective_src_mask = effective_src_mask.unsqueeze(1)
    encoder_outputs = self.encoder(
        input_token,
        input_obj_id,
        segment_label,
        token_type,
        effective_src_mask,
    )
    (
        hidden_states,
        vocab_logits,
        obj_id_logits,
        token_type_logits,
        src,
        class_embeds,
    ) = encoder_outputs
    effective_global_mask = (
        global_mask if global_mask is not None else input_token.ge(2)
    )
    if inference:
        coarse_box, coarse_gmm, refine_box, refine_gmm = self.bbox_head.inference(
            hidden_states,
            effective_src_mask,
            src,
            class_embeds,
            effective_global_mask,
            generator=generator,
        )
    else:
        if bbox is None:
            bbox = input_token.new_full(
                (input_token.size(0), input_token.size(1) - 1, 4), 2.0
            ).float()
        elif bbox.size(1) == input_token.size(1):
            bbox = bbox[:, :-1, :]
        trg_mask = effective_src_mask.new_ones([1, 1, 1])
        coarse_box, coarse_gmm, refine_box, refine_gmm = self.bbox_head(
            0,
            hidden_states,
            effective_src_mask,
            src,
            class_embeds,
            bbox,
            trg_mask,
            effective_global_mask,
            generator=generator,
        )
    output = LTNetModelOutput(
        vocab_logits=vocab_logits,
        obj_id_logits=obj_id_logits,
        token_type_logits=token_type_logits,
        coarse_box=coarse_box,
        coarse_gmm=coarse_gmm,
        refine_box=refine_box,
        refine_gmm=refine_gmm,
        hidden_states=hidden_states if output_hidden_states else None,
    )
    if return_dict:
        return output
    return output.to_tuple()

LTNetModelOutput dataclass

Bases: ModelOutput

Raw LT-Net model outputs.

Attributes:

Name Type Description
vocab_logits Float[Tensor, 'batch sequence vocab'] | None

Mixed token vocabulary logits.

obj_id_logits Float[Tensor, 'batch sequence object_classes'] | None

Object-id classifier logits.

token_type_logits Float[Tensor, 'batch sequence token_types'] | None

Token-type classifier logits.

coarse_box Float[Tensor, 'batch sequence 4'] | None

Coarse normalized center xywh boxes.

coarse_gmm Float[Tensor, 'batch sequence gmm_params'] | None

Optional coarse GMM parameters.

refine_box Float[Tensor, 'batch sequence 4'] | None

Optional refined normalized center xywh boxes.

refine_gmm Float[Tensor, 'batch sequence gmm_params'] | None

Optional refinement GMM parameters.

hidden_states Float[Tensor, 'batch sequence hidden'] | None

Optional encoder hidden states.

Examples:

>>> import torch
>>> out = LTNetModelOutput(
...     vocab_logits=torch.zeros(1, 2, 3),
...     obj_id_logits=torch.zeros(1, 2, 4),
...     token_type_logits=torch.zeros(1, 2, 4),
...     coarse_box=torch.zeros(1, 2, 4),
... )
>>> out.coarse_box.shape
torch.Size([1, 2, 4])
Source code in models/ltnet/src/ltnet/modeling_ltnet.py
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
@dataclass
class LTNetModelOutput(ModelOutput):
    """Raw LT-Net model outputs.

    Attributes:
        vocab_logits: Mixed token vocabulary logits.
        obj_id_logits: Object-id classifier logits.
        token_type_logits: Token-type classifier logits.
        coarse_box: Coarse normalized center ``xywh`` boxes.
        coarse_gmm: Optional coarse GMM parameters.
        refine_box: Optional refined normalized center ``xywh`` boxes.
        refine_gmm: Optional refinement GMM parameters.
        hidden_states: Optional encoder hidden states.

    Examples:
        >>> import torch
        >>> out = LTNetModelOutput(
        ...     vocab_logits=torch.zeros(1, 2, 3),
        ...     obj_id_logits=torch.zeros(1, 2, 4),
        ...     token_type_logits=torch.zeros(1, 2, 4),
        ...     coarse_box=torch.zeros(1, 2, 4),
        ... )
        >>> out.coarse_box.shape
        torch.Size([1, 2, 4])
    """

    vocab_logits: Float[torch.Tensor, "batch sequence vocab"] | None = None
    obj_id_logits: Float[torch.Tensor, "batch sequence object_classes"] | None = None
    token_type_logits: Float[torch.Tensor, "batch sequence token_types"] | None = None
    coarse_box: Float[torch.Tensor, "batch sequence 4"] | None = None
    coarse_gmm: Float[torch.Tensor, "batch sequence gmm_params"] | None = None
    refine_box: Float[torch.Tensor, "batch sequence 4"] | None = None
    refine_gmm: Float[torch.Tensor, "batch sequence gmm_params"] | None = None
    hidden_states: Float[torch.Tensor, "batch sequence hidden"] | None = None

LTNetPipeline

Bases: LayoutGenerationPipeline

Compose an LT-Net model and processor for scene-graph layout inference.

Parameters:

Name Type Description Default
model LTNetForLayoutGeneration

Converted LT-Net model.

required
processor LTNetProcessor

Matching scene-graph processor/tokenizer.

required
config LTNetConfig | None

Optional root pipeline config. Defaults to model.config.

None

Examples:

>>> processor = LTNetProcessor.from_config()
>>> config = LTNetConfig(
...     vocab_size=processor.tokenizer.vocab_size,
...     hidden_size=32,
...     num_hidden_layers=1,
...     num_attention_heads=4,
... )
>>> pipe = LTNetPipeline(
...     model=LTNetForLayoutGeneration(config),
...     processor=processor,
... )
>>> pipe.config.model_type
'ltnet'
Source code in models/ltnet/src/ltnet/pipeline_ltnet.py
 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
231
232
233
234
235
236
237
class LTNetPipeline(LayoutGenerationPipeline):
    """Compose an LT-Net model and processor for scene-graph layout inference.

    Args:
        model: Converted LT-Net model.
        processor: Matching scene-graph processor/tokenizer.
        config: Optional root pipeline config. Defaults to ``model.config``.

    Examples:
        >>> processor = LTNetProcessor.from_config()
        >>> config = LTNetConfig(
        ...     vocab_size=processor.tokenizer.vocab_size,
        ...     hidden_size=32,
        ...     num_hidden_layers=1,
        ...     num_attention_heads=4,
        ... )
        >>> pipe = LTNetPipeline(
        ...     model=LTNetForLayoutGeneration(config),
        ...     processor=processor,
        ... )
        >>> pipe.config.model_type
        'ltnet'
    """

    config_class: ClassVar[type[PretrainedConfig]] = LTNetConfig
    component_specs: ClassVar[dict[str, PipelineComponentSpec]] = {
        "model": PipelineComponentSpec(
            attribute_name="model",
            loader=_load_model_component,
            marker_file="config.json",
        ),
        "processor": PipelineComponentSpec(
            attribute_name="processor",
            loader=_load_processor_component,
            marker_file="preprocessor_config.json",
            save_with_is_main_process=False,
        ),
    }

    config: LTNetConfig
    model: LTNetForLayoutGeneration
    processor: LTNetProcessor

    def __init__(
        self,
        model: LTNetForLayoutGeneration,
        processor: LTNetProcessor,
        config: LTNetConfig | None = None,
    ) -> None:
        """Initialize the pipeline with model and processor components."""
        super().__init__(config or model.config)
        self.config = config or model.config
        self.model = model
        self.processor = processor

    @classmethod
    def _from_pretrained_components(  # ty: ignore[invalid-method-override]
        cls,
        *,
        config: PretrainedConfig,
        components: Mapping[str, LTNetForLayoutGeneration | LTNetProcessor | None],
    ) -> "LTNetPipeline":
        """Build a pipeline from loaded root components."""
        return cls(
            config=cast(LTNetConfig, config),
            model=cast(LTNetForLayoutGeneration, components["model"]),
            processor=cast(LTNetProcessor, components["processor"]),
        )

    @torch.no_grad()
    def __call__(  # ty: ignore[invalid-method-override]
        self,
        *,
        batch_size: int = 1,
        seed: int | None = None,
        generator: torch.Generator | None = None,
        condition_type: ConditionType | str = ConditionType.relation,
        labels: Int[torch.Tensor, "batch elements"] | Sequence[int] | None = None,
        bbox: Float[torch.Tensor, "batch elements 4"]
        | Sequence[Sequence[float]]
        | None = None,
        mask: Bool[torch.Tensor, "batch elements"] | Sequence[bool] | None = None,
        num_elements: int | list[int] | Int[torch.Tensor, "batch"] | None = None,
        box_format: BoxFormat | str = BoxFormat.xywh,
        normalized: bool = True,
        canvas_size: tuple[int, int] | None = None,
        num_inference_steps: int | None = None,
        output_type: OutputType = "dataclass",
        return_intermediates: bool = False,
        scene_graph: SceneGraphInput | SceneGraphMapping | None = None,
        objects: Sequence[LayoutObject] | None = None,
        relations: Sequence[LayoutRelation] | None = None,
    ) -> (
        LayoutGenerationOutput
        | dict[
            str,
            Float[torch.Tensor, "..."]
            | Int[torch.Tensor, "..."]
            | Bool[torch.Tensor, "..."]
            | dict[int, str]
            | dict[str, Float[torch.Tensor, "..."] | None],
        ]
    ):
        """Generate layouts from a public scene graph.

        Args:
            batch_size: Number of layouts to generate from the same graph.
            seed: Convenience seed used only when ``generator`` is absent.
            generator: Optional PyTorch generator; takes precedence over seed.
            condition_type: Must normalize to ``relation``.
            labels: Reserved v1 interface input; LT-Net uses ``scene_graph``.
            bbox: Reserved v1 box constraint input.
            mask: Reserved v1 validity-mask input.
            num_elements: Reserved v1 element-count input.
            box_format: Output box format; LT-Net returns normalized ``xywh``.
            normalized: Whether output boxes should be normalized.
            canvas_size: Reserved denormalization canvas size.
            num_inference_steps: Reserved v1 step count.
            output_type: Return dataclass or dict.
            return_intermediates: Include raw logits/boxes in intermediates.
            scene_graph: Public relation payload.
            objects: Object-node shorthand when ``scene_graph`` is omitted.
            relations: Relation-edge shorthand when ``scene_graph`` is omitted.

        Returns:
            Layout generation output dataclass or dict.

        Raises:
            ValueError: If the condition or graph payload is unsupported.
        """
        _ = (labels, bbox, mask, num_elements, num_inference_steps)
        model_device = next(self.model.parameters()).device
        prepared_generator = self.prepare_generator(
            generator=generator, seed=seed, device=model_device
        )
        encoded = self.processor(
            scene_graph=scene_graph,
            objects=objects,
            relations=relations,
            batch_size=batch_size,
            condition_type=condition_type,
            return_tensors="pt",
        )
        model_inputs = {
            key: value.to(model_device)
            for key, value in encoded.items()
            if isinstance(value, torch.Tensor)
        }
        was_training = self.model.training
        self.model.eval()
        try:
            output = self.model._generate_boxes(
                **model_inputs,
                generator=prepared_generator,
            )
        finally:
            self.model.train(was_training)
        return self.processor.post_process_layout_generation(
            output,
            input_token=model_inputs["input_token"],
            input_obj_id=model_inputs["input_obj_id"],
            token_type=model_inputs["token_type"],
            box_format=box_format,
            normalized=normalized,
            canvas_size=canvas_size,
            output_type=output_type,
            return_intermediates=return_intermediates,
        )

__init__

__init__(
    model: LTNetForLayoutGeneration,
    processor: LTNetProcessor,
    config: LTNetConfig | None = None,
) -> None

Initialize the pipeline with model and processor components.

Source code in models/ltnet/src/ltnet/pipeline_ltnet.py
113
114
115
116
117
118
119
120
121
122
123
def __init__(
    self,
    model: LTNetForLayoutGeneration,
    processor: LTNetProcessor,
    config: LTNetConfig | None = None,
) -> None:
    """Initialize the pipeline with model and processor components."""
    super().__init__(config or model.config)
    self.config = config or model.config
    self.model = model
    self.processor = processor

__call__

__call__(
    *,
    batch_size: int = 1,
    seed: int | None = None,
    generator: Generator | None = None,
    condition_type: ConditionType
    | str = ConditionType.relation,
    labels: Int[Tensor, "batch elements"]
    | Sequence[int]
    | None = None,
    bbox: Float[Tensor, "batch elements 4"]
    | Sequence[Sequence[float]]
    | None = None,
    mask: Bool[Tensor, "batch elements"]
    | Sequence[bool]
    | None = None,
    num_elements: int
    | list[int]
    | Int[Tensor, "batch"]
    | None = None,
    box_format: BoxFormat | str = BoxFormat.xywh,
    normalized: bool = True,
    canvas_size: tuple[int, int] | None = None,
    num_inference_steps: int | None = None,
    output_type: OutputType = "dataclass",
    return_intermediates: bool = False,
    scene_graph: SceneGraphInput
    | SceneGraphMapping
    | None = None,
    objects: Sequence[LayoutObject] | None = None,
    relations: Sequence[LayoutRelation] | None = None,
) -> (
    LayoutGenerationOutput
    | dict[
        str,
        Float[torch.Tensor, "..."]
        | Int[torch.Tensor, "..."]
        | Bool[torch.Tensor, "..."]
        | dict[int, str]
        | dict[str, Float[torch.Tensor, "..."] | None],
    ]
)

Generate layouts from a public scene graph.

Parameters:

Name Type Description Default
batch_size int

Number of layouts to generate from the same graph.

1
seed int | None

Convenience seed used only when generator is absent.

None
generator Generator | None

Optional PyTorch generator; takes precedence over seed.

None
condition_type ConditionType | str

Must normalize to relation.

relation
labels Int[Tensor, 'batch elements'] | Sequence[int] | None

Reserved v1 interface input; LT-Net uses scene_graph.

None
bbox Float[Tensor, 'batch elements 4'] | Sequence[Sequence[float]] | None

Reserved v1 box constraint input.

None
mask Bool[Tensor, 'batch elements'] | Sequence[bool] | None

Reserved v1 validity-mask input.

None
num_elements int | list[int] | Int[Tensor, 'batch'] | None

Reserved v1 element-count input.

None
box_format BoxFormat | str

Output box format; LT-Net returns normalized xywh.

xywh
normalized bool

Whether output boxes should be normalized.

True
canvas_size tuple[int, int] | None

Reserved denormalization canvas size.

None
num_inference_steps int | None

Reserved v1 step count.

None
output_type OutputType

Return dataclass or dict.

'dataclass'
return_intermediates bool

Include raw logits/boxes in intermediates.

False
scene_graph SceneGraphInput | SceneGraphMapping | None

Public relation payload.

None
objects Sequence[LayoutObject] | None

Object-node shorthand when scene_graph is omitted.

None
relations Sequence[LayoutRelation] | None

Relation-edge shorthand when scene_graph is omitted.

None

Returns:

Type Description
LayoutGenerationOutput | dict[str, Float[Tensor, '...'] | Int[Tensor, '...'] | Bool[Tensor, '...'] | dict[int, str] | dict[str, Float[Tensor, '...'] | None]]

Layout generation output dataclass or dict.

Raises:

Type Description
ValueError

If the condition or graph payload is unsupported.

Source code in models/ltnet/src/ltnet/pipeline_ltnet.py
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
231
232
233
234
235
236
237
@torch.no_grad()
def __call__(  # ty: ignore[invalid-method-override]
    self,
    *,
    batch_size: int = 1,
    seed: int | None = None,
    generator: torch.Generator | None = None,
    condition_type: ConditionType | str = ConditionType.relation,
    labels: Int[torch.Tensor, "batch elements"] | Sequence[int] | None = None,
    bbox: Float[torch.Tensor, "batch elements 4"]
    | Sequence[Sequence[float]]
    | None = None,
    mask: Bool[torch.Tensor, "batch elements"] | Sequence[bool] | None = None,
    num_elements: int | list[int] | Int[torch.Tensor, "batch"] | None = None,
    box_format: BoxFormat | str = BoxFormat.xywh,
    normalized: bool = True,
    canvas_size: tuple[int, int] | None = None,
    num_inference_steps: int | None = None,
    output_type: OutputType = "dataclass",
    return_intermediates: bool = False,
    scene_graph: SceneGraphInput | SceneGraphMapping | None = None,
    objects: Sequence[LayoutObject] | None = None,
    relations: Sequence[LayoutRelation] | None = None,
) -> (
    LayoutGenerationOutput
    | dict[
        str,
        Float[torch.Tensor, "..."]
        | Int[torch.Tensor, "..."]
        | Bool[torch.Tensor, "..."]
        | dict[int, str]
        | dict[str, Float[torch.Tensor, "..."] | None],
    ]
):
    """Generate layouts from a public scene graph.

    Args:
        batch_size: Number of layouts to generate from the same graph.
        seed: Convenience seed used only when ``generator`` is absent.
        generator: Optional PyTorch generator; takes precedence over seed.
        condition_type: Must normalize to ``relation``.
        labels: Reserved v1 interface input; LT-Net uses ``scene_graph``.
        bbox: Reserved v1 box constraint input.
        mask: Reserved v1 validity-mask input.
        num_elements: Reserved v1 element-count input.
        box_format: Output box format; LT-Net returns normalized ``xywh``.
        normalized: Whether output boxes should be normalized.
        canvas_size: Reserved denormalization canvas size.
        num_inference_steps: Reserved v1 step count.
        output_type: Return dataclass or dict.
        return_intermediates: Include raw logits/boxes in intermediates.
        scene_graph: Public relation payload.
        objects: Object-node shorthand when ``scene_graph`` is omitted.
        relations: Relation-edge shorthand when ``scene_graph`` is omitted.

    Returns:
        Layout generation output dataclass or dict.

    Raises:
        ValueError: If the condition or graph payload is unsupported.
    """
    _ = (labels, bbox, mask, num_elements, num_inference_steps)
    model_device = next(self.model.parameters()).device
    prepared_generator = self.prepare_generator(
        generator=generator, seed=seed, device=model_device
    )
    encoded = self.processor(
        scene_graph=scene_graph,
        objects=objects,
        relations=relations,
        batch_size=batch_size,
        condition_type=condition_type,
        return_tensors="pt",
    )
    model_inputs = {
        key: value.to(model_device)
        for key, value in encoded.items()
        if isinstance(value, torch.Tensor)
    }
    was_training = self.model.training
    self.model.eval()
    try:
        output = self.model._generate_boxes(
            **model_inputs,
            generator=prepared_generator,
        )
    finally:
        self.model.train(was_training)
    return self.processor.post_process_layout_generation(
        output,
        input_token=model_inputs["input_token"],
        input_obj_id=model_inputs["input_obj_id"],
        token_type=model_inputs["token_type"],
        box_format=box_format,
        normalized=normalized,
        canvas_size=canvas_size,
        output_type=output_type,
        return_intermediates=return_intermediates,
    )

LTNetProcessor

Bases: ProcessorMixin

Normalize scene graphs, tokenize LT-Net inputs, and postprocess boxes.

Source code in models/ltnet/src/ltnet/processing_ltnet.py
 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
231
232
233
234
235
236
237
238
239
240
241
242
243
244
245
246
247
248
249
250
251
252
253
254
255
256
257
258
259
260
261
262
263
264
265
266
267
268
269
270
271
272
273
274
275
276
277
278
279
280
281
282
283
284
285
286
287
288
289
290
291
292
293
294
295
296
297
298
299
300
301
302
303
304
305
306
307
308
309
310
311
312
313
314
315
316
317
318
319
320
321
322
323
324
325
326
327
328
329
330
331
332
333
334
335
336
337
338
339
340
341
342
343
344
345
346
347
348
349
350
351
352
353
354
355
356
357
358
359
360
361
362
363
364
365
366
367
368
369
370
371
372
373
374
375
376
377
378
379
380
381
382
383
384
385
386
387
388
389
390
391
392
393
394
395
396
397
398
399
400
401
402
403
404
405
406
407
408
409
410
411
412
413
414
415
416
417
418
419
420
421
422
423
424
425
426
427
428
429
430
431
432
433
434
435
436
437
438
439
440
441
442
443
444
445
446
447
448
449
450
451
452
453
454
455
456
457
458
459
460
461
class LTNetProcessor(ProcessorMixin):
    """Normalize scene graphs, tokenize LT-Net inputs, and postprocess boxes."""

    attributes = ["tokenizer"]
    tokenizer_class = "LTNetRelationTokenizer"

    def __init__(
        self,
        tokenizer: LTNetRelationTokenizer,
        dataset_name: str = "coco",
        max_sequence_length: int = 128,
        id2label: Mapping[int, str] | Mapping[str, str] | None = None,
        relation_id2label: Mapping[int, str] | Mapping[str, str] | None = None,
        object_reduce: Literal["first", "last", "mean"] = "first",
    ) -> None:
        """Initialize processor label maps and tokenizer component."""
        self.tokenizer = tokenizer
        self.dataset_name = dataset_name
        self.max_sequence_length = max_sequence_length
        self.id2label = {
            int(key): str(value)
            for key, value in (id2label or DEFAULT_ID2LABEL).items()
        }
        self.relation_id2label = {
            int(key): str(value)
            for key, value in (relation_id2label or DEFAULT_RELATION_ID2LABEL).items()
        }
        self.label2id = {value.lower(): key for key, value in self.id2label.items()}
        self.relation_label2id = {
            value.lower(): key for key, value in self.relation_id2label.items()
        }
        if object_reduce not in {"first", "last", "mean"}:
            raise ValueError("object_reduce must be 'first', 'last', or 'mean'")

        self.object_reduce = object_reduce
        super().__init__(tokenizer=tokenizer)

    @classmethod
    def from_config(
        cls,
        *,
        dataset_name: str = "coco",
        max_sequence_length: int = 128,
        id2label: Mapping[int, str] | Mapping[str, str] | None = None,
        relation_id2label: Mapping[int, str] | Mapping[str, str] | None = None,
    ) -> "LTNetProcessor":
        """Construct a processor and tokenizer without external files.

        Returns:
            Processor with synthetic vocabulary derived from the label maps.

        Examples:
            >>> processor = LTNetProcessor.from_config()
            >>> processor.tokenizer.cls_token_id
            1
        """
        object_labels = {
            int(key): str(value)
            for key, value in (id2label or DEFAULT_ID2LABEL).items()
        }
        relation_labels = {
            int(key): str(value)
            for key, value in (relation_id2label or DEFAULT_RELATION_ID2LABEL).items()
        }
        tokens = ["__image__"]
        tokens.extend(value for _, value in sorted(object_labels.items()))
        tokens.extend(value for _, value in sorted(relation_labels.items()))
        tokenizer = LTNetRelationTokenizer(tokens=tokens)
        return cls(
            tokenizer=tokenizer,
            dataset_name=dataset_name,
            max_sequence_length=max_sequence_length,
            id2label=object_labels,
            relation_id2label=relation_labels,
        )

    @classmethod
    def _load_tokenizer_from_pretrained(
        cls,
        sub_processor_type: str,
        pretrained_model_name_or_path: str | PathLike[str],
        subfolder: str = "",
        **kwargs: str | int | float | bool | None,
    ) -> LTNetRelationTokenizer:
        """Load tokenizer for ``ProcessorMixin.from_pretrained``."""
        _ = sub_processor_type
        path = Path(pretrained_model_name_or_path)
        tokenizer_path = path / subfolder if subfolder else path
        token = kwargs.get("token")
        return LTNetRelationTokenizer.from_pretrained(
            tokenizer_path,
            cache_dir=cast(str | PathLike[str] | None, kwargs.get("cache_dir")),
            force_download=bool(kwargs.get("force_download", False)),
            local_files_only=bool(kwargs.get("local_files_only", False)),
            token=token if isinstance(token, str | bool) else None,
            revision=str(kwargs.get("revision", "main")),
        )

    def normalize_condition_type(
        self, condition_type: ConditionType | str
    ) -> ConditionType:
        """Normalize and validate the LT-Net public condition type."""
        condition = normalize_condition_type(condition_type)
        if condition is not ConditionType.relation:
            raise ValueError(
                "LT-Net only supports condition_type='relation' "
                "and aliases 'scene_graph', 'graph', or 'gen_r'."
            )

        return condition

    def _label_to_id(self, label: int | str) -> int:
        if isinstance(label, int):
            return label
        lowered = label.lower()
        if lowered in self.label2id:
            return self.label2id[lowered]
        raise ValueError(f"Unknown object label: {label}")

    def _relation_to_id(self, predicate: int | str) -> int:
        if isinstance(predicate, int):
            return predicate
        lowered = predicate.lower()
        if lowered in self.relation_label2id:
            return self.relation_label2id[lowered]
        raise ValueError(f"Unknown relation label: {predicate}")

    def _token_for_object(self, label_id: int) -> str:
        return self.id2label.get(label_id, str(label_id))

    def _token_for_relation(self, relation_id: int) -> str:
        return self.relation_id2label.get(relation_id, str(relation_id))

    def _normalize_scene_graph(
        self,
        scene_graph: SceneGraphInput | SceneGraphMapping | None,
        *,
        objects: Sequence[LayoutObject] | None,
        relations: Sequence[LayoutRelation] | None,
    ) -> SceneGraphInput:
        if isinstance(scene_graph, SceneGraphInput):
            return scene_graph
        if scene_graph is None:
            if objects is None:
                raise ValueError("scene_graph or objects must be provided")

            return SceneGraphInput(
                objects=tuple(objects),
                relations=tuple(relations or ()),
                id2label=self.id2label,
                relation_id2label=self.relation_id2label,
            )
        nodes = scene_graph.get("nodes", scene_graph.get("objects", ()))
        edges = scene_graph.get("edges", scene_graph.get("relations", ()))

        normalized_objects: list[LayoutObject] = []

        for node in cast(Sequence[SceneGraphItemMapping], nodes):
            item = node
            node_id = cast(int | str, item["id"])
            label = cast(int | str, item.get("label_id", item.get("label")))
            bbox = cast(tuple[float, float, float, float] | None, item.get("bbox"))
            normalized_objects.append(LayoutObject(id=node_id, label=label, bbox=bbox))

        normalized_relations: list[LayoutRelation] = []
        for edge in cast(Sequence[SceneGraphItemMapping], edges):
            item = edge
            subject = cast(int | str, item.get("source", item.get("subject")))
            predicate = cast(int | str, item.get("predicate_id", item.get("predicate")))
            target = cast(int | str, item.get("target", item.get("object")))
            normalized_relations.append(
                LayoutRelation(
                    subject=subject,
                    predicate=predicate,
                    object=target,
                    score=cast(float | None, item.get("score")),
                )
            )

        return SceneGraphInput(
            objects=tuple(normalized_objects),
            relations=tuple(normalized_relations),
            id2label=cast(dict[int, str] | None, scene_graph.get("id2label")),
            relation_id2label=cast(
                dict[int, str] | None,
                scene_graph.get("relation_id2label"),
            ),
        )

    def _serialize_graph(
        self,
        graph: SceneGraphInput,
    ) -> tuple[list[int], list[int], list[int], list[int], list[int]]:
        object_by_id = {item.id: item for item in graph.objects}
        object_ids = {item.id: idx + 1 for idx, item in enumerate(graph.objects)}
        tokens = [self.tokenizer.cls_token]
        input_obj_id = [0]
        segment_label = [0]
        token_type = [0]
        segment = 1
        for relation in graph.relations:
            subject = object_by_id[relation.subject]
            target = object_by_id[relation.object]
            subject_label = self._label_to_id(subject.label)
            target_label = self._label_to_id(target.label)
            relation_id = self._relation_to_id(relation.predicate)
            triple_tokens = [
                self._token_for_object(subject_label),
                self._token_for_relation(relation_id),
                self._token_for_object(target_label),
                self.tokenizer.sep_token,
            ]
            tokens.extend(triple_tokens)
            input_obj_id.extend([object_ids[subject.id], 0, object_ids[target.id], 0])
            segment_label.extend([segment] * 4)
            token_type.extend([1, 2, 3, 0])
            segment += 1
        if not graph.relations:
            for item in graph.objects:
                label_id = self._label_to_id(item.label)
                tokens.extend(
                    [self._token_for_object(label_id), self.tokenizer.sep_token]
                )
                input_obj_id.extend([object_ids[item.id], 0])
                segment_label.extend([segment, segment])
                token_type.extend([1, 0])
                segment += 1

        input_token = self.tokenizer.encode_scene_graph_tokens(tokens)
        length = min(len(input_token), self.max_sequence_length)
        input_token = input_token[:length]
        input_obj_id = input_obj_id[:length]
        segment_label = segment_label[:length]
        token_type = token_type[:length]

        src_mask = [1] * length
        pad_length = self.max_sequence_length - length
        input_token.extend([self.tokenizer.pad_token_id] * pad_length)
        input_obj_id.extend([0] * pad_length)
        segment_label.extend([0] * pad_length)
        token_type.extend([0] * pad_length)
        src_mask.extend([0] * pad_length)
        return input_token, input_obj_id, segment_label, token_type, src_mask

    def __call__(
        self,
        *,
        scene_graph: SceneGraphInput | SceneGraphMapping | None = None,
        objects: Sequence[LayoutObject] | None = None,
        relations: Sequence[LayoutRelation] | None = None,
        batch_size: int = 1,
        condition_type: ConditionType | str = ConditionType.relation,
        return_tensors: Literal["pt", "np"] = "pt",
        max_sequence_length: int | None = None,
    ) -> BatchEncoding:
        """Build LT-Net model tensors from public scene-graph inputs."""
        self.normalize_condition_type(condition_type)
        original_max_length = self.max_sequence_length
        try:
            if max_sequence_length is not None:
                self.max_sequence_length = max_sequence_length
            graph = self._normalize_scene_graph(
                scene_graph,
                objects=objects,
                relations=relations,
            )
            rows = [self._serialize_graph(graph) for _ in range(batch_size)]
        finally:
            self.max_sequence_length = original_max_length
        data = {
            "input_token": torch.tensor([row[0] for row in rows], dtype=torch.long),
            "input_obj_id": torch.tensor([row[1] for row in rows], dtype=torch.long),
            "segment_label": torch.tensor([row[2] for row in rows], dtype=torch.long),
            "token_type": torch.tensor([row[3] for row in rows], dtype=torch.long),
            "src_mask": torch.tensor(
                [row[4] for row in rows], dtype=torch.bool
            ).unsqueeze(1),
            "global_mask": torch.tensor([row[0] for row in rows], dtype=torch.long).ge(
                2
            ),
        }
        if return_tensors == "pt":
            return BatchEncoding(data)
        if return_tensors == "np":
            return BatchEncoding({key: value.numpy() for key, value in data.items()})
        raise ValueError("return_tensors must be 'pt' or 'np'")

    def post_process_layout_generation(
        self,
        model_outputs: LTNetModelOutput,
        *,
        input_token: Int[torch.Tensor, "batch sequence"] | None = None,
        input_obj_id: Int[torch.Tensor, "batch sequence"],
        token_type: Int[torch.Tensor, "batch sequence"],
        box_format: BoxFormat | str = BoxFormat.xywh,
        normalized: bool = True,
        canvas_size: tuple[int, int] | None = None,
        output_type: OutputType = "dataclass",
        return_intermediates: bool = False,
    ) -> (
        LayoutGenerationOutput
        | dict[
            str,
            Float[torch.Tensor, "..."]
            | Int[torch.Tensor, "..."]
            | Bool[torch.Tensor, "..."]
            | dict[int, str]
            | dict[str, Float[torch.Tensor, "..."] | None],
        ]
    ):
        """Convert raw token-level boxes into public object-level layouts."""
        _ = (canvas_size, normalize_box_format(box_format))
        if not normalized:
            raise ValueError("LT-Net outputs normalized boxes only")

        raw_box = model_outputs.refine_box
        if raw_box is None:
            raw_box = model_outputs.coarse_box
        if raw_box is None:
            raise ValueError("model_outputs must contain coarse_box or refine_box")

        batch_boxes: list[Float[torch.Tensor, "elements 4"]] = []
        batch_labels: list[Int[torch.Tensor, "elements"]] = []
        batch_masks: list[Bool[torch.Tensor, "elements"]] = []
        token_rows = (
            [None] * raw_box.size(0)
            if input_token is None
            else list(input_token.unbind(dim=0))
        )
        for row_box, row_obj_id, row_type, row_token in zip(
            raw_box, input_obj_id, token_type, token_rows, strict=True
        ):
            object_positions = row_type.eq(1) | row_type.eq(3)
            object_ids = row_obj_id[object_positions]
            boxes = row_box[object_positions].clamp(0.0, 1.0)
            labels = (
                torch.clamp(object_ids - 1, min=0).long()
                if row_token is None
                else row_token[object_positions].long()
            )
            valid = object_ids.gt(0)
            boxes = boxes[valid]
            labels = labels[valid]
            object_ids = object_ids[valid]
            reduced_boxes: list[Float[torch.Tensor, "elements 4"]] = []
            reduced_labels: list[Int[torch.Tensor, "elements"]] = []
            for object_id in object_ids.unique(sorted=True):
                positions = object_ids.eq(object_id).nonzero().flatten()
                if self.object_reduce == "mean":
                    reduced_boxes.append(boxes[positions].mean(dim=0))
                    reduced_labels.append(labels[positions[0]])
                else:
                    selected = (
                        positions[0] if self.object_reduce == "first" else positions[-1]
                    )
                    reduced_boxes.append(boxes[selected])
                    reduced_labels.append(labels[selected])
            if reduced_boxes:
                batch_boxes.append(torch.stack(reduced_boxes))
                batch_labels.append(torch.stack(reduced_labels).long())
            else:
                batch_boxes.append(boxes)
                batch_labels.append(labels)
            batch_masks.append(
                torch.ones(len(reduced_boxes), dtype=torch.bool, device=row_box.device)
            )
        max_items = max((item.size(0) for item in batch_boxes), default=0)
        padded_boxes = raw_box.new_zeros((raw_box.size(0), max_items, 4))
        padded_labels = input_obj_id.new_zeros((raw_box.size(0), max_items))
        padded_masks = torch.zeros(
            (raw_box.size(0), max_items), dtype=torch.bool, device=raw_box.device
        )
        for idx, (boxes, labels, mask) in enumerate(
            zip(batch_boxes, batch_labels, batch_masks, strict=True)
        ):
            length = boxes.size(0)
            padded_boxes[idx, :length] = boxes
            padded_labels[idx, :length] = labels
            padded_masks[idx, :length] = mask
        intermediates = None
        if return_intermediates:
            intermediates = {
                "coarse_box": model_outputs.coarse_box,
                "refine_box": model_outputs.refine_box,
                "vocab_logits": model_outputs.vocab_logits,
                "obj_id_logits": model_outputs.obj_id_logits,
                "token_type_logits": model_outputs.token_type_logits,
            }
        output = LayoutGenerationOutput(
            bbox=padded_boxes,
            labels=padded_labels,
            mask=padded_masks,
            id2label=dict(self.id2label),
            intermediates=intermediates,
        )
        if output_type == "dict":
            return dict(output.items())
        return output

    def save_pretrained(
        self,
        save_directory: str | PathLike[str],
        push_to_hub: bool = False,
        **kwargs: str | int | float | bool | None,
    ) -> tuple[str, ...]:
        """Save processor metadata and tokenizer files."""
        paths = super().save_pretrained(
            save_directory, push_to_hub=push_to_hub, **kwargs
        )
        metadata = {
            "processor_class": self.__class__.__name__,
            "dataset_name": self.dataset_name,
            "max_sequence_length": self.max_sequence_length,
            "id2label": self.id2label,
            "relation_id2label": self.relation_id2label,
            "object_reduce": self.object_reduce,
        }
        with (Path(save_directory) / "preprocessor_config.json").open("w") as f:
            json.dump(metadata, f, indent=2, sort_keys=True)
        return paths

__init__

__init__(
    tokenizer: LTNetRelationTokenizer,
    dataset_name: str = "coco",
    max_sequence_length: int = 128,
    id2label: Mapping[int, str]
    | Mapping[str, str]
    | None = None,
    relation_id2label: Mapping[int, str]
    | Mapping[str, str]
    | None = None,
    object_reduce: Literal[
        "first", "last", "mean"
    ] = "first",
) -> None

Initialize processor label maps and tokenizer component.

Source code in models/ltnet/src/ltnet/processing_ltnet.py
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
def __init__(
    self,
    tokenizer: LTNetRelationTokenizer,
    dataset_name: str = "coco",
    max_sequence_length: int = 128,
    id2label: Mapping[int, str] | Mapping[str, str] | None = None,
    relation_id2label: Mapping[int, str] | Mapping[str, str] | None = None,
    object_reduce: Literal["first", "last", "mean"] = "first",
) -> None:
    """Initialize processor label maps and tokenizer component."""
    self.tokenizer = tokenizer
    self.dataset_name = dataset_name
    self.max_sequence_length = max_sequence_length
    self.id2label = {
        int(key): str(value)
        for key, value in (id2label or DEFAULT_ID2LABEL).items()
    }
    self.relation_id2label = {
        int(key): str(value)
        for key, value in (relation_id2label or DEFAULT_RELATION_ID2LABEL).items()
    }
    self.label2id = {value.lower(): key for key, value in self.id2label.items()}
    self.relation_label2id = {
        value.lower(): key for key, value in self.relation_id2label.items()
    }
    if object_reduce not in {"first", "last", "mean"}:
        raise ValueError("object_reduce must be 'first', 'last', or 'mean'")

    self.object_reduce = object_reduce
    super().__init__(tokenizer=tokenizer)

from_config classmethod

from_config(
    *,
    dataset_name: str = "coco",
    max_sequence_length: int = 128,
    id2label: Mapping[int, str]
    | Mapping[str, str]
    | None = None,
    relation_id2label: Mapping[int, str]
    | Mapping[str, str]
    | None = None,
) -> "LTNetProcessor"

Construct a processor and tokenizer without external files.

Returns:

Type Description
'LTNetProcessor'

Processor with synthetic vocabulary derived from the label maps.

Examples:

>>> processor = LTNetProcessor.from_config()
>>> processor.tokenizer.cls_token_id
1
Source code in models/ltnet/src/ltnet/processing_ltnet.py
 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
@classmethod
def from_config(
    cls,
    *,
    dataset_name: str = "coco",
    max_sequence_length: int = 128,
    id2label: Mapping[int, str] | Mapping[str, str] | None = None,
    relation_id2label: Mapping[int, str] | Mapping[str, str] | None = None,
) -> "LTNetProcessor":
    """Construct a processor and tokenizer without external files.

    Returns:
        Processor with synthetic vocabulary derived from the label maps.

    Examples:
        >>> processor = LTNetProcessor.from_config()
        >>> processor.tokenizer.cls_token_id
        1
    """
    object_labels = {
        int(key): str(value)
        for key, value in (id2label or DEFAULT_ID2LABEL).items()
    }
    relation_labels = {
        int(key): str(value)
        for key, value in (relation_id2label or DEFAULT_RELATION_ID2LABEL).items()
    }
    tokens = ["__image__"]
    tokens.extend(value for _, value in sorted(object_labels.items()))
    tokens.extend(value for _, value in sorted(relation_labels.items()))
    tokenizer = LTNetRelationTokenizer(tokens=tokens)
    return cls(
        tokenizer=tokenizer,
        dataset_name=dataset_name,
        max_sequence_length=max_sequence_length,
        id2label=object_labels,
        relation_id2label=relation_labels,
    )

normalize_condition_type

normalize_condition_type(
    condition_type: ConditionType | str,
) -> ConditionType

Normalize and validate the LT-Net public condition type.

Source code in models/ltnet/src/ltnet/processing_ltnet.py
140
141
142
143
144
145
146
147
148
149
150
151
def normalize_condition_type(
    self, condition_type: ConditionType | str
) -> ConditionType:
    """Normalize and validate the LT-Net public condition type."""
    condition = normalize_condition_type(condition_type)
    if condition is not ConditionType.relation:
        raise ValueError(
            "LT-Net only supports condition_type='relation' "
            "and aliases 'scene_graph', 'graph', or 'gen_r'."
        )

    return condition

__call__

__call__(
    *,
    scene_graph: SceneGraphInput
    | SceneGraphMapping
    | None = None,
    objects: Sequence[LayoutObject] | None = None,
    relations: Sequence[LayoutRelation] | None = None,
    batch_size: int = 1,
    condition_type: ConditionType
    | str = ConditionType.relation,
    return_tensors: Literal["pt", "np"] = "pt",
    max_sequence_length: int | None = None,
) -> BatchEncoding

Build LT-Net model tensors from public scene-graph inputs.

Source code in models/ltnet/src/ltnet/processing_ltnet.py
286
287
288
289
290
291
292
293
294
295
296
297
298
299
300
301
302
303
304
305
306
307
308
309
310
311
312
313
314
315
316
317
318
319
320
321
322
323
324
325
326
327
def __call__(
    self,
    *,
    scene_graph: SceneGraphInput | SceneGraphMapping | None = None,
    objects: Sequence[LayoutObject] | None = None,
    relations: Sequence[LayoutRelation] | None = None,
    batch_size: int = 1,
    condition_type: ConditionType | str = ConditionType.relation,
    return_tensors: Literal["pt", "np"] = "pt",
    max_sequence_length: int | None = None,
) -> BatchEncoding:
    """Build LT-Net model tensors from public scene-graph inputs."""
    self.normalize_condition_type(condition_type)
    original_max_length = self.max_sequence_length
    try:
        if max_sequence_length is not None:
            self.max_sequence_length = max_sequence_length
        graph = self._normalize_scene_graph(
            scene_graph,
            objects=objects,
            relations=relations,
        )
        rows = [self._serialize_graph(graph) for _ in range(batch_size)]
    finally:
        self.max_sequence_length = original_max_length
    data = {
        "input_token": torch.tensor([row[0] for row in rows], dtype=torch.long),
        "input_obj_id": torch.tensor([row[1] for row in rows], dtype=torch.long),
        "segment_label": torch.tensor([row[2] for row in rows], dtype=torch.long),
        "token_type": torch.tensor([row[3] for row in rows], dtype=torch.long),
        "src_mask": torch.tensor(
            [row[4] for row in rows], dtype=torch.bool
        ).unsqueeze(1),
        "global_mask": torch.tensor([row[0] for row in rows], dtype=torch.long).ge(
            2
        ),
    }
    if return_tensors == "pt":
        return BatchEncoding(data)
    if return_tensors == "np":
        return BatchEncoding({key: value.numpy() for key, value in data.items()})
    raise ValueError("return_tensors must be 'pt' or 'np'")

post_process_layout_generation

post_process_layout_generation(
    model_outputs: LTNetModelOutput,
    *,
    input_token: Int[Tensor, "batch sequence"]
    | None = None,
    input_obj_id: Int[Tensor, "batch sequence"],
    token_type: Int[Tensor, "batch sequence"],
    box_format: BoxFormat | str = BoxFormat.xywh,
    normalized: bool = True,
    canvas_size: tuple[int, int] | None = None,
    output_type: OutputType = "dataclass",
    return_intermediates: bool = False,
) -> (
    LayoutGenerationOutput
    | dict[
        str,
        Float[torch.Tensor, "..."]
        | Int[torch.Tensor, "..."]
        | Bool[torch.Tensor, "..."]
        | dict[int, str]
        | dict[str, Float[torch.Tensor, "..."] | None],
    ]
)

Convert raw token-level boxes into public object-level layouts.

Source code in models/ltnet/src/ltnet/processing_ltnet.py
329
330
331
332
333
334
335
336
337
338
339
340
341
342
343
344
345
346
347
348
349
350
351
352
353
354
355
356
357
358
359
360
361
362
363
364
365
366
367
368
369
370
371
372
373
374
375
376
377
378
379
380
381
382
383
384
385
386
387
388
389
390
391
392
393
394
395
396
397
398
399
400
401
402
403
404
405
406
407
408
409
410
411
412
413
414
415
416
417
418
419
420
421
422
423
424
425
426
427
428
429
430
431
432
433
434
435
436
437
438
439
def post_process_layout_generation(
    self,
    model_outputs: LTNetModelOutput,
    *,
    input_token: Int[torch.Tensor, "batch sequence"] | None = None,
    input_obj_id: Int[torch.Tensor, "batch sequence"],
    token_type: Int[torch.Tensor, "batch sequence"],
    box_format: BoxFormat | str = BoxFormat.xywh,
    normalized: bool = True,
    canvas_size: tuple[int, int] | None = None,
    output_type: OutputType = "dataclass",
    return_intermediates: bool = False,
) -> (
    LayoutGenerationOutput
    | dict[
        str,
        Float[torch.Tensor, "..."]
        | Int[torch.Tensor, "..."]
        | Bool[torch.Tensor, "..."]
        | dict[int, str]
        | dict[str, Float[torch.Tensor, "..."] | None],
    ]
):
    """Convert raw token-level boxes into public object-level layouts."""
    _ = (canvas_size, normalize_box_format(box_format))
    if not normalized:
        raise ValueError("LT-Net outputs normalized boxes only")

    raw_box = model_outputs.refine_box
    if raw_box is None:
        raw_box = model_outputs.coarse_box
    if raw_box is None:
        raise ValueError("model_outputs must contain coarse_box or refine_box")

    batch_boxes: list[Float[torch.Tensor, "elements 4"]] = []
    batch_labels: list[Int[torch.Tensor, "elements"]] = []
    batch_masks: list[Bool[torch.Tensor, "elements"]] = []
    token_rows = (
        [None] * raw_box.size(0)
        if input_token is None
        else list(input_token.unbind(dim=0))
    )
    for row_box, row_obj_id, row_type, row_token in zip(
        raw_box, input_obj_id, token_type, token_rows, strict=True
    ):
        object_positions = row_type.eq(1) | row_type.eq(3)
        object_ids = row_obj_id[object_positions]
        boxes = row_box[object_positions].clamp(0.0, 1.0)
        labels = (
            torch.clamp(object_ids - 1, min=0).long()
            if row_token is None
            else row_token[object_positions].long()
        )
        valid = object_ids.gt(0)
        boxes = boxes[valid]
        labels = labels[valid]
        object_ids = object_ids[valid]
        reduced_boxes: list[Float[torch.Tensor, "elements 4"]] = []
        reduced_labels: list[Int[torch.Tensor, "elements"]] = []
        for object_id in object_ids.unique(sorted=True):
            positions = object_ids.eq(object_id).nonzero().flatten()
            if self.object_reduce == "mean":
                reduced_boxes.append(boxes[positions].mean(dim=0))
                reduced_labels.append(labels[positions[0]])
            else:
                selected = (
                    positions[0] if self.object_reduce == "first" else positions[-1]
                )
                reduced_boxes.append(boxes[selected])
                reduced_labels.append(labels[selected])
        if reduced_boxes:
            batch_boxes.append(torch.stack(reduced_boxes))
            batch_labels.append(torch.stack(reduced_labels).long())
        else:
            batch_boxes.append(boxes)
            batch_labels.append(labels)
        batch_masks.append(
            torch.ones(len(reduced_boxes), dtype=torch.bool, device=row_box.device)
        )
    max_items = max((item.size(0) for item in batch_boxes), default=0)
    padded_boxes = raw_box.new_zeros((raw_box.size(0), max_items, 4))
    padded_labels = input_obj_id.new_zeros((raw_box.size(0), max_items))
    padded_masks = torch.zeros(
        (raw_box.size(0), max_items), dtype=torch.bool, device=raw_box.device
    )
    for idx, (boxes, labels, mask) in enumerate(
        zip(batch_boxes, batch_labels, batch_masks, strict=True)
    ):
        length = boxes.size(0)
        padded_boxes[idx, :length] = boxes
        padded_labels[idx, :length] = labels
        padded_masks[idx, :length] = mask
    intermediates = None
    if return_intermediates:
        intermediates = {
            "coarse_box": model_outputs.coarse_box,
            "refine_box": model_outputs.refine_box,
            "vocab_logits": model_outputs.vocab_logits,
            "obj_id_logits": model_outputs.obj_id_logits,
            "token_type_logits": model_outputs.token_type_logits,
        }
    output = LayoutGenerationOutput(
        bbox=padded_boxes,
        labels=padded_labels,
        mask=padded_masks,
        id2label=dict(self.id2label),
        intermediates=intermediates,
    )
    if output_type == "dict":
        return dict(output.items())
    return output

save_pretrained

save_pretrained(
    save_directory: str | PathLike[str],
    push_to_hub: bool = False,
    **kwargs: str | int | float | bool | None,
) -> tuple[str, ...]

Save processor metadata and tokenizer files.

Source code in models/ltnet/src/ltnet/processing_ltnet.py
441
442
443
444
445
446
447
448
449
450
451
452
453
454
455
456
457
458
459
460
461
def save_pretrained(
    self,
    save_directory: str | PathLike[str],
    push_to_hub: bool = False,
    **kwargs: str | int | float | bool | None,
) -> tuple[str, ...]:
    """Save processor metadata and tokenizer files."""
    paths = super().save_pretrained(
        save_directory, push_to_hub=push_to_hub, **kwargs
    )
    metadata = {
        "processor_class": self.__class__.__name__,
        "dataset_name": self.dataset_name,
        "max_sequence_length": self.max_sequence_length,
        "id2label": self.id2label,
        "relation_id2label": self.relation_id2label,
        "object_reduce": self.object_reduce,
    }
    with (Path(save_directory) / "preprocessor_config.json").open("w") as f:
        json.dump(metadata, f, indent=2, sort_keys=True)
    return paths

LayoutObject dataclass

One scene-graph object node.

Parameters:

Name Type Description Default
id int | str

Stable object id within a scene graph.

required
label int | str

Dataset-local object label id or label string.

required
bbox tuple[float, float, float, float] | None

Optional normalized center xywh constraint.

None

Examples:

>>> LayoutObject(id="person-1", label="person").id
'person-1'
Source code in models/ltnet/src/ltnet/relation_schema.py
 8
 9
10
11
12
13
14
15
16
17
18
19
20
21
22
23
24
@dataclass(frozen=True)
class LayoutObject:
    """One scene-graph object node.

    Args:
        id: Stable object id within a scene graph.
        label: Dataset-local object label id or label string.
        bbox: Optional normalized center ``xywh`` constraint.

    Examples:
        >>> LayoutObject(id="person-1", label="person").id
        'person-1'
    """

    id: int | str
    label: int | str
    bbox: tuple[float, float, float, float] | None = None

LayoutRelation dataclass

One directed scene-graph edge.

Parameters:

Name Type Description Default
subject int | str

Source object id.

required
predicate int | str

Relation id or label.

required
object int | str

Target object id.

required
bbox_delta tuple[float, float, float, float] | None

Optional relation geometry.

None
score float | None

Optional edge confidence.

None

Examples:

>>> LayoutRelation("a", "left of", "b").predicate
'left of'
Source code in models/ltnet/src/ltnet/relation_schema.py
27
28
29
30
31
32
33
34
35
36
37
38
39
40
41
42
43
44
45
46
47
@dataclass(frozen=True)
class LayoutRelation:
    """One directed scene-graph edge.

    Args:
        subject: Source object id.
        predicate: Relation id or label.
        object: Target object id.
        bbox_delta: Optional relation geometry.
        score: Optional edge confidence.

    Examples:
        >>> LayoutRelation("a", "left of", "b").predicate
        'left of'
    """

    subject: int | str
    predicate: int | str
    object: int | str
    bbox_delta: tuple[float, float, float, float] | None = None
    score: float | None = None

SceneGraphInput dataclass

Normalized scene graph payload.

Parameters:

Name Type Description Default
objects tuple[LayoutObject, ...]

Scene-graph object nodes.

required
relations tuple[LayoutRelation, ...]

Scene-graph relation edges.

required
id2label dict[int, str] | None

Optional public object label mapping.

None
relation_id2label dict[int, str] | None

Optional relation label mapping.

None

Examples:

>>> graph = SceneGraphInput(
...     objects=(LayoutObject("a", "person"), LayoutObject("b", "table")),
...     relations=(LayoutRelation("a", "left of", "b"),),
... )
>>> len(graph.relations)
1
Source code in models/ltnet/src/ltnet/relation_schema.py
50
51
52
53
54
55
56
57
58
59
60
61
62
63
64
65
66
67
68
69
70
71
72
@dataclass(frozen=True)
class SceneGraphInput:
    """Normalized scene graph payload.

    Args:
        objects: Scene-graph object nodes.
        relations: Scene-graph relation edges.
        id2label: Optional public object label mapping.
        relation_id2label: Optional relation label mapping.

    Examples:
        >>> graph = SceneGraphInput(
        ...     objects=(LayoutObject("a", "person"), LayoutObject("b", "table")),
        ...     relations=(LayoutRelation("a", "left of", "b"),),
        ... )
        >>> len(graph.relations)
        1
    """

    objects: tuple[LayoutObject, ...]
    relations: tuple[LayoutRelation, ...]
    id2label: dict[int, str] | None = None
    relation_id2label: dict[int, str] | None = None

LTNetRelationTokenizer

Bases: WhitespaceTokenizerMixin, PreTrainedTokenizer

Discrete scene-graph tokenizer saved as a standard HF tokenizer.

Parameters:

Name Type Description Default
vocab_file str | None

Optional path to object_pred_id2name.json.

None
tokens list[str] | None

Optional token list used when no vocab file is supplied.

None
object_token_ids list[int] | None

Optional ids that represent object classes.

None
relation_token_ids list[int] | None

Optional ids that represent predicates.

None
pad_token str

Padding token.

'[PAD]'
cls_token str

Sequence-start token.

'[CLS]'
sep_token str

Triple separator token.

'[SEP]'
mask_token str

Mask token.

'[MASK]'
unk_token str

Unknown token.

'[MASK]'
model_max_length int

Maximum tokenizer length metadata.

DEFAULT_MODEL_MAX_LENGTH
kwargs str | int | float | bool | None

Additional tokenizer compatibility fields.

{}

Examples:

>>> tokenizer = LTNetRelationTokenizer(tokens=["__image__", "person"])
>>> tokenizer.convert_tokens_to_ids("[CLS]")
1
Source code in models/ltnet/src/ltnet/tokenization_ltnet.py
 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
class LTNetRelationTokenizer(WhitespaceTokenizerMixin, PreTrainedTokenizer):
    """Discrete scene-graph tokenizer saved as a standard HF tokenizer.

    Args:
        vocab_file: Optional path to ``object_pred_id2name.json``.
        tokens: Optional token list used when no vocab file is supplied.
        object_token_ids: Optional ids that represent object classes.
        relation_token_ids: Optional ids that represent predicates.
        pad_token: Padding token.
        cls_token: Sequence-start token.
        sep_token: Triple separator token.
        mask_token: Mask token.
        unk_token: Unknown token.
        model_max_length: Maximum tokenizer length metadata.
        kwargs: Additional tokenizer compatibility fields.

    Examples:
        >>> tokenizer = LTNetRelationTokenizer(tokens=["__image__", "person"])
        >>> tokenizer.convert_tokens_to_ids("[CLS]")
        1
    """

    vocab_files_names = {"vocab_file": "object_pred_id2name.json"}
    model_input_names = ["input_token", "input_obj_id", "segment_label", "token_type"]

    def __init__(
        self,
        vocab_file: str | None = None,
        tokens: list[str] | None = None,
        object_token_ids: list[int] | None = None,
        relation_token_ids: list[int] | None = None,
        pad_token: str = "[PAD]",
        cls_token: str = "[CLS]",
        sep_token: str = "[SEP]",
        mask_token: str = "[MASK]",
        unk_token: str = "[MASK]",
        model_max_length: int = DEFAULT_MODEL_MAX_LENGTH,
        padding_side: str = "right",
        truncation_side: str = "right",
        clean_up_tokenization_spaces: bool = False,
        added_tokens_decoder: dict[int | str, str | AddedToken] | None = None,
        name_or_path: str = "",
        **kwargs: str | int | float | bool | None,
    ) -> None:
        """Initialize vocabulary and object/relation id metadata."""
        _ = kwargs
        token2id, id2token = build_token_maps(
            vocab_file=vocab_file,
            tokens=tokens,
            base_tokens=SPECIAL_TOKENS,
            numeric_id_vocab=True,
        )
        self._token2id = token2id
        self._id2token = id2token
        self.object_token_ids = [int(item) for item in object_token_ids or []]
        self.relation_token_ids = [int(item) for item in relation_token_ids or []]
        tokenizer_kwargs: dict[str, object] = {
            "pad_token": pad_token,
            "cls_token": cls_token,
            "sep_token": sep_token,
            "mask_token": mask_token,
            "unk_token": unk_token,
            "model_max_length": model_max_length,
            "padding_side": padding_side,
            "truncation_side": truncation_side,
            "clean_up_tokenization_spaces": clean_up_tokenization_spaces,
            "name_or_path": name_or_path,
        }
        if added_tokens_decoder is not None:
            tokenizer_kwargs["added_tokens_decoder"] = added_tokens_decoder
        super().__init__(**tokenizer_kwargs)

    def encode_scene_graph_tokens(self, tokens: list[str]) -> list[int]:
        """Encode already-normalized scene-graph token strings.

        Args:
            tokens: Token strings in LT-Net order.

        Returns:
            Integer token ids.

        Raises:
            ValueError: If an unknown token appears.
        """
        ids: list[int] = []
        for token in tokens:
            token_id = self._convert_token_to_id(token)
            if token_id == self.unk_token_id and token != self.unk_token:
                raise ValueError(f"Unknown scene-graph token: {token}")

            ids.append(token_id)
        return ids

    def decode_scene_graph_tokens(self, input_token: list[int]) -> list[str]:
        """Decode integer scene-graph token ids into token strings."""
        return [self._convert_id_to_token(token_id) for token_id in input_token]

    def save_vocabulary(
        self, save_directory: str, filename_prefix: str | None = None
    ) -> tuple[str]:
        """Save ``object_pred_id2name.json`` as id-to-token metadata."""
        return save_json_vocabulary(
            save_directory=save_directory,
            filename="object_pred_id2name.json",
            data={str(key): value for key, value in sorted(self._id2token.items())},
            filename_prefix=filename_prefix,
        )

    def save_pretrained(
        self,
        save_directory: str | PathLike[str],
        legacy_format: bool | None = None,
        filename_prefix: str | None = None,
        push_to_hub: bool = False,
        **kwargs: str | int | float | bool | None,
    ) -> tuple[str, ...]:
        """Save tokenizer files plus LT-Net tokenizer metadata."""
        _ = kwargs
        paths = super().save_pretrained(
            str(save_directory),
            legacy_format=legacy_format,
            filename_prefix=filename_prefix,
            push_to_hub=push_to_hub,
        )
        metadata = {
            "object_token_ids": self.object_token_ids,
            "relation_token_ids": self.relation_token_ids,
        }
        with (Path(save_directory) / "ltnet_tokenizer_config.json").open("w") as f:
            json.dump(metadata, f, indent=2, sort_keys=True)
        return paths

    @classmethod
    def from_pretrained(
        cls,
        pretrained_model_name_or_path: str | PathLike[str],
        cache_dir: str | PathLike[str] | None = None,
        force_download: bool = False,
        local_files_only: bool = False,
        token: str | bool | None = None,
        revision: str = "main",
        object_token_ids: list[int] | None = None,
        relation_token_ids: list[int] | None = None,
        **kwargs: str | int | float | bool | None,
    ) -> "LTNetRelationTokenizer":
        """Load tokenizer and LT-Net metadata."""
        path = Path(pretrained_model_name_or_path)
        metadata_path = path / "ltnet_tokenizer_config.json"
        metadata: dict[str, object] = {}
        if metadata_path.exists():
            with metadata_path.open() as f:
                metadata = json.load(f)
        if object_token_ids is not None:
            metadata["object_token_ids"] = object_token_ids
        if relation_token_ids is not None:
            metadata["relation_token_ids"] = relation_token_ids
        metadata.update(kwargs)
        return cast(
            "LTNetRelationTokenizer",
            super().from_pretrained(
                str(pretrained_model_name_or_path),
                cache_dir=cache_dir,
                force_download=force_download,
                local_files_only=local_files_only,
                token=token,
                revision=revision,
                **metadata,
            ),
        )

__init__

__init__(
    vocab_file: str | None = None,
    tokens: list[str] | None = None,
    object_token_ids: list[int] | None = None,
    relation_token_ids: list[int] | None = None,
    pad_token: str = "[PAD]",
    cls_token: str = "[CLS]",
    sep_token: str = "[SEP]",
    mask_token: str = "[MASK]",
    unk_token: str = "[MASK]",
    model_max_length: int = DEFAULT_MODEL_MAX_LENGTH,
    padding_side: str = "right",
    truncation_side: str = "right",
    clean_up_tokenization_spaces: bool = False,
    added_tokens_decoder: dict[int | str, str | AddedToken]
    | None = None,
    name_or_path: str = "",
    **kwargs: str | int | float | bool | None,
) -> None

Initialize vocabulary and object/relation id metadata.

Source code in models/ltnet/src/ltnet/tokenization_ltnet.py
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
def __init__(
    self,
    vocab_file: str | None = None,
    tokens: list[str] | None = None,
    object_token_ids: list[int] | None = None,
    relation_token_ids: list[int] | None = None,
    pad_token: str = "[PAD]",
    cls_token: str = "[CLS]",
    sep_token: str = "[SEP]",
    mask_token: str = "[MASK]",
    unk_token: str = "[MASK]",
    model_max_length: int = DEFAULT_MODEL_MAX_LENGTH,
    padding_side: str = "right",
    truncation_side: str = "right",
    clean_up_tokenization_spaces: bool = False,
    added_tokens_decoder: dict[int | str, str | AddedToken] | None = None,
    name_or_path: str = "",
    **kwargs: str | int | float | bool | None,
) -> None:
    """Initialize vocabulary and object/relation id metadata."""
    _ = kwargs
    token2id, id2token = build_token_maps(
        vocab_file=vocab_file,
        tokens=tokens,
        base_tokens=SPECIAL_TOKENS,
        numeric_id_vocab=True,
    )
    self._token2id = token2id
    self._id2token = id2token
    self.object_token_ids = [int(item) for item in object_token_ids or []]
    self.relation_token_ids = [int(item) for item in relation_token_ids or []]
    tokenizer_kwargs: dict[str, object] = {
        "pad_token": pad_token,
        "cls_token": cls_token,
        "sep_token": sep_token,
        "mask_token": mask_token,
        "unk_token": unk_token,
        "model_max_length": model_max_length,
        "padding_side": padding_side,
        "truncation_side": truncation_side,
        "clean_up_tokenization_spaces": clean_up_tokenization_spaces,
        "name_or_path": name_or_path,
    }
    if added_tokens_decoder is not None:
        tokenizer_kwargs["added_tokens_decoder"] = added_tokens_decoder
    super().__init__(**tokenizer_kwargs)

encode_scene_graph_tokens

encode_scene_graph_tokens(tokens: list[str]) -> list[int]

Encode already-normalized scene-graph token strings.

Parameters:

Name Type Description Default
tokens list[str]

Token strings in LT-Net order.

required

Returns:

Type Description
list[int]

Integer token ids.

Raises:

Type Description
ValueError

If an unknown token appears.

Source code in models/ltnet/src/ltnet/tokenization_ltnet.py
 94
 95
 96
 97
 98
 99
100
101
102
103
104
105
106
107
108
109
110
111
112
113
def encode_scene_graph_tokens(self, tokens: list[str]) -> list[int]:
    """Encode already-normalized scene-graph token strings.

    Args:
        tokens: Token strings in LT-Net order.

    Returns:
        Integer token ids.

    Raises:
        ValueError: If an unknown token appears.
    """
    ids: list[int] = []
    for token in tokens:
        token_id = self._convert_token_to_id(token)
        if token_id == self.unk_token_id and token != self.unk_token:
            raise ValueError(f"Unknown scene-graph token: {token}")

        ids.append(token_id)
    return ids

decode_scene_graph_tokens

decode_scene_graph_tokens(
    input_token: list[int],
) -> list[str]

Decode integer scene-graph token ids into token strings.

Source code in models/ltnet/src/ltnet/tokenization_ltnet.py
115
116
117
def decode_scene_graph_tokens(self, input_token: list[int]) -> list[str]:
    """Decode integer scene-graph token ids into token strings."""
    return [self._convert_id_to_token(token_id) for token_id in input_token]

save_vocabulary

save_vocabulary(
    save_directory: str, filename_prefix: str | None = None
) -> tuple[str]

Save object_pred_id2name.json as id-to-token metadata.

Source code in models/ltnet/src/ltnet/tokenization_ltnet.py
119
120
121
122
123
124
125
126
127
128
def save_vocabulary(
    self, save_directory: str, filename_prefix: str | None = None
) -> tuple[str]:
    """Save ``object_pred_id2name.json`` as id-to-token metadata."""
    return save_json_vocabulary(
        save_directory=save_directory,
        filename="object_pred_id2name.json",
        data={str(key): value for key, value in sorted(self._id2token.items())},
        filename_prefix=filename_prefix,
    )

save_pretrained

save_pretrained(
    save_directory: str | PathLike[str],
    legacy_format: bool | None = None,
    filename_prefix: str | None = None,
    push_to_hub: bool = False,
    **kwargs: str | int | float | bool | None,
) -> tuple[str, ...]

Save tokenizer files plus LT-Net tokenizer metadata.

Source code in models/ltnet/src/ltnet/tokenization_ltnet.py
130
131
132
133
134
135
136
137
138
139
140
141
142
143
144
145
146
147
148
149
150
151
152
def save_pretrained(
    self,
    save_directory: str | PathLike[str],
    legacy_format: bool | None = None,
    filename_prefix: str | None = None,
    push_to_hub: bool = False,
    **kwargs: str | int | float | bool | None,
) -> tuple[str, ...]:
    """Save tokenizer files plus LT-Net tokenizer metadata."""
    _ = kwargs
    paths = super().save_pretrained(
        str(save_directory),
        legacy_format=legacy_format,
        filename_prefix=filename_prefix,
        push_to_hub=push_to_hub,
    )
    metadata = {
        "object_token_ids": self.object_token_ids,
        "relation_token_ids": self.relation_token_ids,
    }
    with (Path(save_directory) / "ltnet_tokenizer_config.json").open("w") as f:
        json.dump(metadata, f, indent=2, sort_keys=True)
    return paths

from_pretrained classmethod

from_pretrained(
    pretrained_model_name_or_path: str | PathLike[str],
    cache_dir: str | PathLike[str] | None = None,
    force_download: bool = False,
    local_files_only: bool = False,
    token: str | bool | None = None,
    revision: str = "main",
    object_token_ids: list[int] | None = None,
    relation_token_ids: list[int] | None = None,
    **kwargs: str | int | float | bool | None,
) -> "LTNetRelationTokenizer"

Load tokenizer and LT-Net metadata.

Source code in models/ltnet/src/ltnet/tokenization_ltnet.py
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
@classmethod
def from_pretrained(
    cls,
    pretrained_model_name_or_path: str | PathLike[str],
    cache_dir: str | PathLike[str] | None = None,
    force_download: bool = False,
    local_files_only: bool = False,
    token: str | bool | None = None,
    revision: str = "main",
    object_token_ids: list[int] | None = None,
    relation_token_ids: list[int] | None = None,
    **kwargs: str | int | float | bool | None,
) -> "LTNetRelationTokenizer":
    """Load tokenizer and LT-Net metadata."""
    path = Path(pretrained_model_name_or_path)
    metadata_path = path / "ltnet_tokenizer_config.json"
    metadata: dict[str, object] = {}
    if metadata_path.exists():
        with metadata_path.open() as f:
            metadata = json.load(f)
    if object_token_ids is not None:
        metadata["object_token_ids"] = object_token_ids
    if relation_token_ids is not None:
        metadata["relation_token_ids"] = relation_token_ids
    metadata.update(kwargs)
    return cast(
        "LTNetRelationTokenizer",
        super().from_pretrained(
            str(pretrained_model_name_or_path),
            cache_dir=cache_dir,
            force_download=force_download,
            local_files_only=local_files_only,
            token=token,
            revision=revision,
            **metadata,
        ),
    )

configuration_ltnet

Configuration for converted LT-Net checkpoints.

DecoderHeadType

Bases: StrEnum

Supported LT-Net bbox decoder head types.

Source code in models/ltnet/src/ltnet/configuration_ltnet.py
12
13
14
15
16
class DecoderHeadType(StrEnum):
    """Supported LT-Net bbox decoder head types."""

    gmm = auto()
    linear = auto()

BoxLossType

Bases: StrEnum

Supported LT-Net box objective families.

Source code in models/ltnet/src/ltnet/configuration_ltnet.py
19
20
21
22
23
class BoxLossType(StrEnum):
    """Supported LT-Net box objective families."""

    pdf = auto()
    reg = auto()

LTNetConfig

Bases: PretrainedConfig

Architecture and processor metadata for LT-Net checkpoints.

Parameters:

Name Type Description Default
dataset_name str

Dataset slug for the converted checkpoint.

'coco'
vocab_size int

Mixed special/object/predicate token vocabulary size.

206
obj_classes_size int

Object-id embedding/classifier vocabulary size.

155
hidden_size int

Transformer hidden dimension.

256
num_hidden_layers int

Number of relation encoder layers.

4
num_attention_heads int

Number of relation encoder attention heads.

4
dropout float

Dropout used by embeddings, encoder, and bbox heads.

0.1
enable_noise bool

Whether the original config enabled relation noise.

False
noise_size int

Original relation-noise size.

64
decoder_head_type DecoderHeadType | str

GMM or Linear bbox head.

gmm
decoder_box_loss BoxLossType | str

PDF or Reg objective family.

pdf
decoder_schedule_sample bool

Whether training used scheduled sampling.

False
decoder_two_path bool

Whether the original config enabled two-path decoding.

False
decoder_global_feature bool

Whether to concatenate max-pooled global features.

True
decoder_greedy bool

Whether inference uses GMM means instead of sampling.

True
xy_temperature float

GMM mixture temperature for center coordinates.

1.0
wh_temperature float

GMM mixture temperature for box size.

1.0
refine bool

Whether a refinement head is present.

False
refine_head_type DecoderHeadType | str

Refinement bbox head type.

linear
refine_box_loss BoxLossType | str

Refinement objective family.

reg
refine_x_softmax bool

Whether the decoder records XY PDF scores for refine.

True
max_sequence_length int

Processor/model sequence length.

128
id2label Mapping[int, str] | Mapping[str, str] | None

Public object label mapping.

None
relation_id2label Mapping[int, str] | Mapping[str, str] | None

Public relation label mapping.

None
model_type str | None

Ignored compatibility field from serialized configs.

None
transformers_version str | None

Ignored compatibility field.

None
kwargs str | int | float | bool | None

Additional PretrainedConfig fields.

{}

Examples:

>>> config = LTNetConfig(hidden_size=32, num_attention_heads=4)
>>> config.model_type
'ltnet'
Source code in models/ltnet/src/ltnet/configuration_ltnet.py
 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
class LTNetConfig(PretrainedConfig):
    """Architecture and processor metadata for LT-Net checkpoints.

    Args:
        dataset_name: Dataset slug for the converted checkpoint.
        vocab_size: Mixed special/object/predicate token vocabulary size.
        obj_classes_size: Object-id embedding/classifier vocabulary size.
        hidden_size: Transformer hidden dimension.
        num_hidden_layers: Number of relation encoder layers.
        num_attention_heads: Number of relation encoder attention heads.
        dropout: Dropout used by embeddings, encoder, and bbox heads.
        enable_noise: Whether the original config enabled relation noise.
        noise_size: Original relation-noise size.
        decoder_head_type: ``GMM`` or ``Linear`` bbox head.
        decoder_box_loss: ``PDF`` or ``Reg`` objective family.
        decoder_schedule_sample: Whether training used scheduled sampling.
        decoder_two_path: Whether the original config enabled two-path decoding.
        decoder_global_feature: Whether to concatenate max-pooled global features.
        decoder_greedy: Whether inference uses GMM means instead of sampling.
        xy_temperature: GMM mixture temperature for center coordinates.
        wh_temperature: GMM mixture temperature for box size.
        refine: Whether a refinement head is present.
        refine_head_type: Refinement bbox head type.
        refine_box_loss: Refinement objective family.
        refine_x_softmax: Whether the decoder records XY PDF scores for refine.
        max_sequence_length: Processor/model sequence length.
        id2label: Public object label mapping.
        relation_id2label: Public relation label mapping.
        model_type: Ignored compatibility field from serialized configs.
        transformers_version: Ignored compatibility field.
        kwargs: Additional ``PretrainedConfig`` fields.

    Examples:
        >>> config = LTNetConfig(hidden_size=32, num_attention_heads=4)
        >>> config.model_type
        'ltnet'
    """

    model_type = "ltnet"

    def __init__(
        self,
        *,
        dataset_name: str = "coco",
        vocab_size: int = 206,
        obj_classes_size: int = 155,
        hidden_size: int = 256,
        num_hidden_layers: int = 4,
        num_attention_heads: int = 4,
        dropout: float = 0.1,
        enable_noise: bool = False,
        noise_size: int = 64,
        decoder_head_type: DecoderHeadType | str = DecoderHeadType.gmm,
        decoder_box_loss: BoxLossType | str = BoxLossType.pdf,
        decoder_schedule_sample: bool = False,
        decoder_two_path: bool = False,
        decoder_global_feature: bool = True,
        decoder_greedy: bool = True,
        xy_temperature: float = 1.0,
        wh_temperature: float = 1.0,
        refine: bool = False,
        refine_head_type: DecoderHeadType | str = DecoderHeadType.linear,
        refine_box_loss: BoxLossType | str = BoxLossType.reg,
        refine_x_softmax: bool = True,
        max_sequence_length: int = 128,
        id2label: Mapping[int, str] | Mapping[str, str] | None = None,
        relation_id2label: Mapping[int, str] | Mapping[str, str] | None = None,
        bos_token_id: int = 1,
        eos_token_id: int = 2,
        pad_token_id: int = 0,
        mask_token_id: int = 3,
        model_type: str | None = None,
        transformers_version: str | None = None,
        **kwargs: str | int | float | bool | None,
    ) -> None:
        """Initialize LT-Net architecture and metadata fields."""
        _ = (model_type, transformers_version)

        self.dataset_name = dataset_name
        self.vocab_size = vocab_size
        self.obj_classes_size = obj_classes_size
        self.hidden_size = hidden_size
        self.num_hidden_layers = num_hidden_layers
        self.num_attention_heads = num_attention_heads
        self.dropout = dropout
        self.enable_noise = enable_noise
        self.noise_size = noise_size

        self.decoder_head_type = str(_normalize_head_type(decoder_head_type))
        self.decoder_box_loss = str(_normalize_box_loss(decoder_box_loss))
        self.decoder_schedule_sample = decoder_schedule_sample
        self.decoder_two_path = decoder_two_path
        self.decoder_global_feature = decoder_global_feature
        self.decoder_greedy = decoder_greedy
        self.xy_temperature = xy_temperature
        self.wh_temperature = wh_temperature
        self.refine = refine

        self.refine_head_type = str(_normalize_head_type(refine_head_type))
        self.refine_box_loss = str(_normalize_box_loss(refine_box_loss))
        self.refine_x_softmax = refine_x_softmax
        self.max_sequence_length = max_sequence_length
        self.mask_token_id = mask_token_id
        _ = kwargs

        super().__init__()
        normalized_id2label = id2label or DEFAULT_ID2LABEL
        normalized_relation_id2label = relation_id2label or DEFAULT_RELATION_ID2LABEL
        self.id2label = {
            int(key): str(value) for key, value in normalized_id2label.items()
        }
        self.relation_id2label = {
            int(key): str(value) for key, value in normalized_relation_id2label.items()
        }
        self.bos_token_id = bos_token_id
        self.eos_token_id = eos_token_id
        self.pad_token_id = pad_token_id
        self.label2id = {value: key for key, value in self.id2label.items()}

__init__

__init__(
    *,
    dataset_name: str = "coco",
    vocab_size: int = 206,
    obj_classes_size: int = 155,
    hidden_size: int = 256,
    num_hidden_layers: int = 4,
    num_attention_heads: int = 4,
    dropout: float = 0.1,
    enable_noise: bool = False,
    noise_size: int = 64,
    decoder_head_type: DecoderHeadType
    | str = DecoderHeadType.gmm,
    decoder_box_loss: BoxLossType | str = BoxLossType.pdf,
    decoder_schedule_sample: bool = False,
    decoder_two_path: bool = False,
    decoder_global_feature: bool = True,
    decoder_greedy: bool = True,
    xy_temperature: float = 1.0,
    wh_temperature: float = 1.0,
    refine: bool = False,
    refine_head_type: DecoderHeadType
    | str = DecoderHeadType.linear,
    refine_box_loss: BoxLossType | str = BoxLossType.reg,
    refine_x_softmax: bool = True,
    max_sequence_length: int = 128,
    id2label: Mapping[int, str]
    | Mapping[str, str]
    | None = None,
    relation_id2label: Mapping[int, str]
    | Mapping[str, str]
    | None = None,
    bos_token_id: int = 1,
    eos_token_id: int = 2,
    pad_token_id: int = 0,
    mask_token_id: int = 3,
    model_type: str | None = None,
    transformers_version: str | None = None,
    **kwargs: str | int | float | bool | None,
) -> None

Initialize LT-Net architecture and metadata fields.

Source code in models/ltnet/src/ltnet/configuration_ltnet.py
 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
def __init__(
    self,
    *,
    dataset_name: str = "coco",
    vocab_size: int = 206,
    obj_classes_size: int = 155,
    hidden_size: int = 256,
    num_hidden_layers: int = 4,
    num_attention_heads: int = 4,
    dropout: float = 0.1,
    enable_noise: bool = False,
    noise_size: int = 64,
    decoder_head_type: DecoderHeadType | str = DecoderHeadType.gmm,
    decoder_box_loss: BoxLossType | str = BoxLossType.pdf,
    decoder_schedule_sample: bool = False,
    decoder_two_path: bool = False,
    decoder_global_feature: bool = True,
    decoder_greedy: bool = True,
    xy_temperature: float = 1.0,
    wh_temperature: float = 1.0,
    refine: bool = False,
    refine_head_type: DecoderHeadType | str = DecoderHeadType.linear,
    refine_box_loss: BoxLossType | str = BoxLossType.reg,
    refine_x_softmax: bool = True,
    max_sequence_length: int = 128,
    id2label: Mapping[int, str] | Mapping[str, str] | None = None,
    relation_id2label: Mapping[int, str] | Mapping[str, str] | None = None,
    bos_token_id: int = 1,
    eos_token_id: int = 2,
    pad_token_id: int = 0,
    mask_token_id: int = 3,
    model_type: str | None = None,
    transformers_version: str | None = None,
    **kwargs: str | int | float | bool | None,
) -> None:
    """Initialize LT-Net architecture and metadata fields."""
    _ = (model_type, transformers_version)

    self.dataset_name = dataset_name
    self.vocab_size = vocab_size
    self.obj_classes_size = obj_classes_size
    self.hidden_size = hidden_size
    self.num_hidden_layers = num_hidden_layers
    self.num_attention_heads = num_attention_heads
    self.dropout = dropout
    self.enable_noise = enable_noise
    self.noise_size = noise_size

    self.decoder_head_type = str(_normalize_head_type(decoder_head_type))
    self.decoder_box_loss = str(_normalize_box_loss(decoder_box_loss))
    self.decoder_schedule_sample = decoder_schedule_sample
    self.decoder_two_path = decoder_two_path
    self.decoder_global_feature = decoder_global_feature
    self.decoder_greedy = decoder_greedy
    self.xy_temperature = xy_temperature
    self.wh_temperature = wh_temperature
    self.refine = refine

    self.refine_head_type = str(_normalize_head_type(refine_head_type))
    self.refine_box_loss = str(_normalize_box_loss(refine_box_loss))
    self.refine_x_softmax = refine_x_softmax
    self.max_sequence_length = max_sequence_length
    self.mask_token_id = mask_token_id
    _ = kwargs

    super().__init__()
    normalized_id2label = id2label or DEFAULT_ID2LABEL
    normalized_relation_id2label = relation_id2label or DEFAULT_RELATION_ID2LABEL
    self.id2label = {
        int(key): str(value) for key, value in normalized_id2label.items()
    }
    self.relation_id2label = {
        int(key): str(value) for key, value in normalized_relation_id2label.items()
    }
    self.bos_token_id = bos_token_id
    self.eos_token_id = eos_token_id
    self.pad_token_id = pad_token_id
    self.label2id = {value: key for key, value in self.id2label.items()}

conversion

Conversion helpers for original LT-Net checkpoints.

convert_original_checkpoint

convert_original_checkpoint(
    *,
    checkpoint_path: str | Path,
    cfg_path: str | Path,
    vocab_path: str | Path,
    output_dir: str | Path,
    dataset_name: Literal["coco", "vg_msdn"],
    push_to_hub: bool = False,
    hub_model_id: str | None = None,
    strict: bool = True,
) -> None

Convert an original LT-Net checkpoint into local HF-style files.

Parameters:

Name Type Description Default
checkpoint_path str | Path

Vendor .pth checkpoint containing state_dict or a raw state dict.

required
cfg_path str | Path

Vendor YAML config path. Stored in metadata for traceability.

required
vocab_path str | Path

JSON id-to-token vocabulary exported from object_pred_idx_to_name.pkl.

required
output_dir str | Path

Directory to write converted model/processor files.

required
dataset_name Literal['coco', 'vg_msdn']

Dataset identifier for the converted checkpoint.

required
push_to_hub bool

Reserved publish flag; implementation PRs keep it false.

False
hub_model_id str | None

Optional Hub repo id used only when publishing is enabled.

None
strict bool

Whether checkpoint keys must exactly match the converted model.

True

Raises:

Type Description
NotImplementedError

If Hub upload is requested from this helper.

Examples:

>>> convert_original_checkpoint
<function convert_original_checkpoint at ...>
Source code in models/ltnet/src/ltnet/conversion.py
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
def convert_original_checkpoint(
    *,
    checkpoint_path: str | Path,
    cfg_path: str | Path,
    vocab_path: str | Path,
    output_dir: str | Path,
    dataset_name: Literal["coco", "vg_msdn"],
    push_to_hub: bool = False,
    hub_model_id: str | None = None,
    strict: bool = True,
) -> None:
    """Convert an original LT-Net checkpoint into local HF-style files.

    Args:
        checkpoint_path: Vendor ``.pth`` checkpoint containing ``state_dict`` or
            a raw state dict.
        cfg_path: Vendor YAML config path. Stored in metadata for traceability.
        vocab_path: JSON id-to-token vocabulary exported from
            ``object_pred_idx_to_name.pkl``.
        output_dir: Directory to write converted model/processor files.
        dataset_name: Dataset identifier for the converted checkpoint.
        push_to_hub: Reserved publish flag; implementation PRs keep it false.
        hub_model_id: Optional Hub repo id used only when publishing is enabled.
        strict: Whether checkpoint keys must exactly match the converted model.

    Raises:
        NotImplementedError: If Hub upload is requested from this helper.

    Examples:
        >>> convert_original_checkpoint  # doctest: +ELLIPSIS
        <function convert_original_checkpoint at ...>
    """
    if push_to_hub or hub_model_id is not None:
        raise NotImplementedError(
            "Hub upload is intentionally not part of PR conversion"
        )

    if not strict:
        raise ValueError("LT-Net conversion requires strict=True")

    out = Path(output_dir)
    out.mkdir(parents=True, exist_ok=True)
    vocab = _load_vocab(vocab_path)
    id2label, relation_id2label, object_token_ids, relation_token_ids = _split_vocab(
        vocab
    )
    config = _load_config(
        cfg_path=cfg_path,
        dataset_name=dataset_name,
        vocab_size=len(vocab),
        id2label=id2label,
        relation_id2label=relation_id2label,
    )
    model = LTNetForLayoutGeneration(config)
    state_dict = load_original_state_dict(checkpoint_path)
    incompatible = model.load_state_dict(state_dict, strict=True)
    metadata = {
        "checkpoint_path": str(checkpoint_path),
        "cfg_path": str(cfg_path),
        "vocab_path": str(vocab_path),
        "dataset_name": dataset_name,
        "missing_keys": list(incompatible.missing_keys),
        "unexpected_keys": list(incompatible.unexpected_keys),
        "strict_vendor_key_mapping": True,
    }
    model.save_pretrained(out)
    tokenizer = LTNetRelationTokenizer(
        tokens=[vocab[idx] for idx in sorted(vocab)],
        object_token_ids=object_token_ids,
        relation_token_ids=relation_token_ids,
    )
    processor = LTNetProcessor(
        tokenizer=tokenizer,
        dataset_name=dataset_name,
        max_sequence_length=config.max_sequence_length,
        id2label=cast(dict[int, str], config.id2label),
        relation_id2label=config.relation_id2label,
    )
    processor.save_pretrained(out)
    with (out / "conversion_metadata.json").open("w") as f:
        json.dump(metadata, f, indent=2, sort_keys=True)

modeling_lt_compatible

Local LT-Net modules with original checkpoint-compatible state keys.

MultiHeadedAttention

Bases: Module

Multi-head attention matching the original LT-Net implementation.

Source code in models/ltnet/src/ltnet/modeling_lt_compatible.py
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
class MultiHeadedAttention(nn.Module):
    """Multi-head attention matching the original LT-Net implementation."""

    def __init__(self, num_heads: int, size: int, dropout: float = 0.1) -> None:
        super().__init__()
        if size % num_heads != 0:
            raise ValueError("attention size must be divisible by num_heads")

        self.head_size = size // num_heads
        self.model_size = size
        self.num_heads = num_heads
        self.k_layer = nn.Linear(size, size)
        self.v_layer = nn.Linear(size, size)
        self.q_layer = nn.Linear(size, size)
        self.output_layer = nn.Linear(size, size)
        self.softmax = nn.Softmax(dim=-1)
        self.dropout = nn.Dropout(dropout)

    def forward(
        self,
        k: Float[torch.Tensor, "batch sequence features"],
        v: Float[torch.Tensor, "batch sequence features"],
        q: Float[torch.Tensor, "batch sequence features"],
        mask: Bool[torch.Tensor, "..."] | None = None,
    ) -> Float[torch.Tensor, "batch sequence features"]:
        batch_size = k.size(0)
        num_heads = self.num_heads
        k = self.k_layer(k)
        v = self.v_layer(v)
        q = self.q_layer(q)
        k = k.view(batch_size, -1, num_heads, self.head_size).transpose(1, 2)
        v = v.view(batch_size, -1, num_heads, self.head_size).transpose(1, 2)
        q = q.view(batch_size, -1, num_heads, self.head_size).transpose(1, 2)
        q = q / math.sqrt(self.head_size)
        scores = torch.matmul(q, k.transpose(2, 3))
        if mask is not None:
            scores = scores.masked_fill(~mask.unsqueeze(1), float("-inf"))
        attention = self.dropout(self.softmax(scores))
        context = torch.matmul(attention, v)
        context = (
            context.transpose(1, 2)
            .contiguous()
            .view(batch_size, -1, num_heads * self.head_size)
        )
        return self.output_layer(context)

ContMultiHeadedAttention

Bases: Module

Continuous-valued attention used by the original bbox decoder.

Source code in models/ltnet/src/ltnet/modeling_lt_compatible.py
 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
class ContMultiHeadedAttention(nn.Module):
    """Continuous-valued attention used by the original bbox decoder."""

    def __init__(
        self, num_heads: int, size: int, size_v: int, dropout: float = 0.1
    ) -> None:
        super().__init__()
        if size % num_heads != 0:
            raise ValueError("attention size must be divisible by num_heads")

        self.head_size = size // num_heads
        self.model_size = size
        self.num_heads = num_heads
        self.k_layer = nn.Linear(size, size)
        self.v_layer = nn.Linear(size_v, size)
        self.q_layer = nn.Linear(size, size)
        self.output_layer = nn.Linear(size, size_v)
        self.softmax = nn.Softmax(dim=-1)
        self.dropout = nn.Dropout(dropout)

    def forward(
        self,
        k: Float[torch.Tensor, "batch sequence features"],
        v: Float[torch.Tensor, "batch sequence features"],
        q: Float[torch.Tensor, "batch sequence features"],
        mask: Bool[torch.Tensor, "..."] | None = None,
    ) -> Float[torch.Tensor, "batch sequence features"]:
        batch_size = k.size(0)
        num_heads = self.num_heads
        k = self.k_layer(k)
        v = self.v_layer(v)
        q = self.q_layer(q)
        k = k.view(batch_size, -1, num_heads, self.head_size).transpose(1, 2)
        v = v.view(batch_size, -1, num_heads, self.head_size).transpose(1, 2)
        q = q.view(batch_size, -1, num_heads, self.head_size).transpose(1, 2)
        q = q / math.sqrt(self.head_size)
        scores = torch.matmul(q, k.transpose(2, 3))
        if mask is not None:
            scores = scores.masked_fill(~mask.unsqueeze(1), float("-inf"))
        attention = self.dropout(self.softmax(scores))
        context = torch.matmul(attention, v)
        context = (
            context.transpose(1, 2)
            .contiguous()
            .view(batch_size, -1, num_heads * self.head_size)
        )
        return self.output_layer(context)

CustomAttention

Bases: Module

Refinement attention with optional PDF confidence reweighting.

Source code in models/ltnet/src/ltnet/modeling_lt_compatible.py
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
class CustomAttention(nn.Module):
    """Refinement attention with optional PDF confidence reweighting."""

    def __init__(
        self, num_heads: int, size: int, dropout: float = 0.1, sent_length: int = 128
    ) -> None:
        super().__init__()
        if size % num_heads != 0:
            raise ValueError("attention size must be divisible by num_heads")

        self.head_size = size // num_heads
        self.model_size = size
        self.num_heads = num_heads
        self.k_layer = nn.Linear(size // 4, num_heads * self.head_size // 4)
        self.v_layer = nn.Linear(size, num_heads * self.head_size)
        self.q_layer = nn.Linear(size // 4, num_heads * self.head_size // 4)
        self.confident_layer = nn.Sequential(
            nn.Linear(sent_length, sent_length), nn.ReLU()
        )
        self.output_layer = nn.Linear(size, size)
        self.softmax = nn.Softmax(dim=-1)
        self.dropout = nn.Dropout(dropout)

    def forward(
        self,
        k: Float[torch.Tensor, "batch sequence box_features"],
        v: Float[torch.Tensor, "batch sequence hidden"],
        q: Float[torch.Tensor, "batch sequence box_features"],
        mask: Bool[torch.Tensor, "..."] | None = None,
        xy_pdf_score: Float[torch.Tensor, "batch sequence"] | None = None,
    ) -> Float[torch.Tensor, "batch sequence hidden"]:
        batch_size = k.size(0)
        num_heads = self.num_heads
        k = self.k_layer(k)
        v = self.v_layer(v)
        q = self.q_layer(q)
        k = k.view(batch_size, -1, num_heads, self.head_size // 4).transpose(1, 2)
        v = v.view(batch_size, -1, num_heads, self.head_size).transpose(1, 2)
        q = q.view(batch_size, -1, num_heads, self.head_size // 4).transpose(1, 2)
        q = q / math.sqrt(self.head_size)
        scores = torch.matmul(q, k.transpose(2, 3))
        if mask is not None:
            scores = scores.masked_fill(~mask.unsqueeze(1), float("-inf"))
        if xy_pdf_score is None:
            attention = self.softmax(scores)
        else:
            xy_pdf_score = self.confident_layer(xy_pdf_score).view(batch_size, 1, 1, -1)
            scores_exp = scores.exp()
            new_scores = scores_exp * xy_pdf_score
            attention = new_scores / new_scores.sum(-1).unsqueeze(-1)
        attention = self.dropout(attention)
        context = torch.matmul(attention, v)
        context = (
            context.transpose(1, 2)
            .contiguous()
            .view(batch_size, -1, num_heads * self.head_size)
        )
        return self.output_layer(context)

GELU

Bases: Module

Original LT-Net GELU implementation.

Source code in models/ltnet/src/ltnet/modeling_lt_compatible.py
199
200
201
202
203
204
205
206
207
208
209
class GELU(nn.Module):
    """Original LT-Net GELU implementation."""

    def forward(
        self, x: Float[torch.Tensor, "batch sequence features"]
    ) -> Float[torch.Tensor, "batch sequence features"]:
        return (
            0.5
            * x
            * (1 + torch.tanh(math.sqrt(2 / math.pi) * (x + 0.044715 * x.pow(3))))
        )

PositionwiseFeedForward

Bases: Module

Original pre-norm feed-forward block.

Source code in models/ltnet/src/ltnet/modeling_lt_compatible.py
212
213
214
215
216
217
218
219
220
221
222
223
224
225
226
227
228
229
class PositionwiseFeedForward(nn.Module):
    """Original pre-norm feed-forward block."""

    def __init__(self, input_size: int, ff_size: int, dropout: float = 0.1) -> None:
        super().__init__()
        self.layer_norm = nn.LayerNorm(input_size, eps=1e-6)
        self.pwff_layer = nn.Sequential(
            nn.Linear(input_size, ff_size),
            GELU(),
            nn.Dropout(dropout),
            nn.Linear(ff_size, input_size),
            nn.Dropout(dropout),
        )

    def forward(
        self, x: Float[torch.Tensor, "batch sequence features"]
    ) -> Float[torch.Tensor, "batch sequence features"]:
        return self.pwff_layer(self.layer_norm(x)) + x

TransformerEncoderLayer

Bases: Module

Original relation encoder layer.

Source code in models/ltnet/src/ltnet/modeling_lt_compatible.py
232
233
234
235
236
237
238
239
240
241
242
243
244
245
246
247
248
249
250
251
252
class TransformerEncoderLayer(nn.Module):
    """Original relation encoder layer."""

    def __init__(
        self, size: int = 0, ff_size: int = 0, num_heads: int = 0, dropout: float = 0.1
    ) -> None:
        super().__init__()
        self.layer_norm = nn.LayerNorm(size, eps=1e-6)
        self.src_src_att = MultiHeadedAttention(num_heads, size, dropout=dropout)
        self.feed_forward = PositionwiseFeedForward(size, ff_size=ff_size)
        self.dropout = nn.Dropout(dropout)
        self.size = size

    def forward(
        self,
        x: Float[torch.Tensor, "batch sequence hidden"],
        mask: Bool[torch.Tensor, "..."],
    ) -> Float[torch.Tensor, "batch sequence hidden"]:
        x_norm = self.layer_norm(x)
        h = self.src_src_att(x_norm, x_norm, x_norm, mask)
        return self.feed_forward(self.dropout(h) + x)

TransformerEncoder

Bases: Module

Original relation transformer encoder.

Source code in models/ltnet/src/ltnet/modeling_lt_compatible.py
255
256
257
258
259
260
261
262
263
264
265
266
267
268
269
270
271
272
273
274
275
276
277
278
279
280
281
282
283
284
285
286
287
288
289
290
291
292
class TransformerEncoder(nn.Module):
    """Original relation transformer encoder."""

    def __init__(
        self,
        hidden_size: int = 512,
        ff_size: int = 2048,
        num_layers: int = 6,
        num_heads: int = 8,
        dropout: float = 0.1,
        emb_dropout: float = 0.1,
    ) -> None:
        super().__init__()
        self.layers = nn.ModuleList(
            [
                TransformerEncoderLayer(
                    size=hidden_size,
                    ff_size=ff_size,
                    num_heads=num_heads,
                    dropout=dropout,
                )
                for _ in range(num_layers)
            ]
        )
        self.layer_norm = nn.LayerNorm(hidden_size, eps=1e-6)
        self.emb_dropout = nn.Dropout(p=emb_dropout)
        self._output_size = hidden_size
        self._hidden_size = hidden_size

    def forward(
        self,
        embed_src: Float[torch.Tensor, "batch sequence hidden"],
        mask: Bool[torch.Tensor, "..."],
    ) -> Float[torch.Tensor, "batch sequence hidden"]:
        x = embed_src
        for layer in self.layers:
            x = layer(x, mask)
        return self.layer_norm(x)

SentenceEmbeddings

Bases: Module

Original sentence/object/token-type embeddings.

Source code in models/ltnet/src/ltnet/modeling_lt_compatible.py
295
296
297
298
299
300
301
302
303
304
305
306
307
308
309
310
311
312
313
314
315
316
317
318
319
320
321
322
323
324
325
326
327
328
329
330
331
332
333
class SentenceEmbeddings(nn.Module):
    """Original sentence/object/token-type embeddings."""

    def __init__(
        self,
        vocab_size: int = 204,
        obj_classes_size: int = 154,
        hidden_size: int = 512,
        max_rel_pair: int = 33,
        max_token_type: int = 4,
        hidden_dropout_prob: float = 0.1,
    ) -> None:
        super().__init__()
        self.word_embeddings = nn.Embedding(vocab_size, hidden_size, padding_idx=0)
        self.obj_id_embeddings = nn.Embedding(
            obj_classes_size, hidden_size, padding_idx=0
        )
        self.sentence_type = nn.Embedding(max_rel_pair, hidden_size, padding_idx=0)
        self.token_type = nn.Embedding(max_token_type, hidden_size, padding_idx=0)
        self.dropout = nn.Dropout(hidden_dropout_prob)

    def forward(
        self,
        input_token: Int[torch.Tensor, "batch sequence"],
        input_obj_id: Int[torch.Tensor, "batch sequence"],
        segment_label: Int[torch.Tensor, "batch sequence"],
        token_type: Int[torch.Tensor, "batch sequence"],
    ) -> tuple[
        Float[torch.Tensor, "batch sequence hidden"],
        Float[torch.Tensor, "batch sequence hidden"],
    ]:
        inputs_embeds = self.word_embeddings(input_token)
        embeddings = (
            inputs_embeds
            + self.sentence_type(segment_label)
            + self.obj_id_embeddings(input_obj_id)
            + self.token_type(token_type)
        )
        return self.dropout(embeddings), inputs_embeds

RelEncoder

Bases: Module

Original LT-Net relation encoder and token classifiers.

Source code in models/ltnet/src/ltnet/modeling_lt_compatible.py
336
337
338
339
340
341
342
343
344
345
346
347
348
349
350
351
352
353
354
355
356
357
358
359
360
361
362
363
364
365
366
367
368
369
370
371
372
373
374
375
376
377
378
379
380
381
382
383
384
385
386
387
class RelEncoder(nn.Module):
    """Original LT-Net relation encoder and token classifiers."""

    def __init__(self, config: LTNetConfig) -> None:
        super().__init__()
        self.input_embeddings = SentenceEmbeddings(
            config.vocab_size,
            config.obj_classes_size,
            config.hidden_size,
            max_rel_pair=33,
            hidden_dropout_prob=config.dropout,
        )
        self.encoder = TransformerEncoder(
            hidden_size=config.hidden_size,
            ff_size=config.hidden_size * 4,
            num_layers=config.num_hidden_layers,
            num_heads=config.num_attention_heads,
            dropout=config.dropout,
            emb_dropout=config.dropout,
        )
        self.hidden_size = config.hidden_size
        self.vocab_classifier = nn.Linear(config.hidden_size, config.vocab_size)
        self.obj_id_classifier = nn.Linear(config.hidden_size, config.obj_classes_size)
        self.token_type_classifier = nn.Linear(config.hidden_size, 4)

    def forward(
        self,
        input_token: Int[torch.Tensor, "batch sequence"],
        input_obj_id: Int[torch.Tensor, "batch sequence"],
        segment_label: Int[torch.Tensor, "batch sequence"],
        token_type: Int[torch.Tensor, "batch sequence"],
        src_mask: Bool[torch.Tensor, "..."],
    ) -> tuple[
        Float[torch.Tensor, "batch sequence hidden"],
        Float[torch.Tensor, "batch sequence vocab"],
        Float[torch.Tensor, "batch sequence object_classes"],
        Float[torch.Tensor, "batch sequence token_types"],
        Float[torch.Tensor, "batch sequence hidden"],
        Float[torch.Tensor, "batch sequence hidden"],
    ]:
        src, class_embeds = self.input_embeddings(
            input_token, input_obj_id, segment_label, token_type
        )
        encoder_output = self.encoder(src, src_mask)
        return (
            encoder_output,
            self.vocab_classifier(encoder_output),
            self.obj_id_classifier(encoder_output),
            self.token_type_classifier(encoder_output),
            src,
            class_embeds,
        )

CustomTransformerDecoderLayer

Bases: Module

Original bbox decoder layer.

Source code in models/ltnet/src/ltnet/modeling_lt_compatible.py
390
391
392
393
394
395
396
397
398
399
400
401
402
403
404
405
406
407
408
409
410
411
412
413
414
415
416
417
418
419
420
421
422
423
424
425
426
427
428
429
430
431
432
class CustomTransformerDecoderLayer(nn.Module):
    """Original bbox decoder layer."""

    def __init__(
        self,
        size: int = 0,
        bb_size: int = 64,
        ff_size: int = 0,
        num_heads: int = 0,
        dropout: float = 0.1,
    ) -> None:
        super().__init__()
        self.size = size
        self.trg_trg_att = ContMultiHeadedAttention(
            num_heads, bb_size, bb_size, dropout=dropout
        )
        self.src_trg_att = ContMultiHeadedAttention(
            num_heads, size, size, dropout=dropout
        )
        self.feed_forward_h1 = PositionwiseFeedForward(bb_size, ff_size=ff_size)
        self.feed_forward_h2 = PositionwiseFeedForward(size, ff_size=ff_size)
        self.x_layer_norm = nn.LayerNorm(size, eps=1e-6)
        self.spa_layer_norm = nn.LayerNorm(bb_size, eps=1e-6)
        self.dropout = nn.Dropout(dropout)

    def forward(
        self,
        spatial_x: Float[torch.Tensor, "batch box_sequence box_features"],
        semantic_x: Float[torch.Tensor, "batch decoder_sequence hidden"],
        memory: Float[torch.Tensor, "batch sequence hidden"],
        src_mask: Bool[torch.Tensor, "..."] | None = None,
        trg_mask: Bool[torch.Tensor, "..."] | None = None,
    ) -> Float[torch.Tensor, "batch decoder_sequence features"]:
        _ = src_mask
        spatial_x_norm = self.spa_layer_norm(spatial_x)
        self.x_layer_norm(semantic_x)
        h1 = self.trg_trg_att(
            spatial_x_norm, spatial_x_norm, spatial_x_norm, mask=trg_mask
        )
        h1 = self.dropout(h1) + spatial_x
        o1 = self.feed_forward_h1(h1)
        o2 = memory[:, 1:, :]
        return torch.cat((o2, o1), dim=-1)

CustomTransformerDecoder

Bases: Module

Original custom bbox transformer decoder.

Source code in models/ltnet/src/ltnet/modeling_lt_compatible.py
435
436
437
438
439
440
441
442
443
444
445
446
447
448
449
450
451
452
453
454
455
456
457
458
459
460
461
462
463
464
465
466
467
468
469
470
471
472
473
474
475
476
477
478
479
480
481
482
483
484
485
486
487
488
489
490
491
492
493
494
495
496
497
class CustomTransformerDecoder(nn.Module):
    """Original custom bbox transformer decoder."""

    def __init__(
        self,
        hidden_size: int = 768,
        hidden_bb_size: int = 64,
        ff_size: int = 2048,
        num_layers: int = 6,
        num_heads: int = 8,
        dropout: float = 0.1,
        emb_dropout: float = 0.1,
    ) -> None:
        super().__init__()
        self._hidden_size = hidden_size
        self.layers = nn.ModuleList(
            [
                CustomTransformerDecoderLayer(
                    size=hidden_size,
                    bb_size=hidden_bb_size,
                    ff_size=ff_size,
                    num_heads=num_heads,
                    dropout=dropout,
                )
                for _ in range(num_layers)
            ]
        )
        self.layer_norm = nn.LayerNorm(hidden_size + hidden_bb_size, eps=1e-6)
        self.emb_dropout = nn.Dropout(p=emb_dropout)

    def forward(
        self,
        trg_embed_0: Float[torch.Tensor, "batch box_sequence box_features"],
        trg_embed_1: Float[torch.Tensor, "batch decoder_sequence hidden"],
        encoder_output: Float[torch.Tensor, "batch sequence hidden"],
        encoder_hidden: Float[torch.Tensor, "batch sequence hidden"] | None = None,
        src_mask: Bool[torch.Tensor, "..."] | None = None,
        unroll_steps: int | None = None,
        hidden: Float[torch.Tensor, "batch decoder_sequence hidden"] | None = None,
        trg_mask: Bool[torch.Tensor, "..."] | None = None,
    ) -> Float[torch.Tensor, "batch decoder_sequence features"]:
        _ = (encoder_hidden, unroll_steps, hidden)
        if trg_mask is None:
            raise ValueError("trg_mask required for Transformer")

        trg_mask = trg_mask & self.subsequent_mask(trg_embed_0.size(1)).type_as(
            trg_mask
        ).to(trg_mask.device)
        x = torch.cat((trg_embed_1, trg_embed_0[:, : trg_embed_1.size(1)]), dim=-1)
        for layer in self.layers:
            x = layer(
                spatial_x=trg_embed_0,
                semantic_x=trg_embed_1,
                memory=encoder_output,
                src_mask=src_mask,
                trg_mask=trg_mask,
            )
        return self.layer_norm(x)

    @staticmethod
    def subsequent_mask(size: int) -> Bool[torch.Tensor, "..."]:
        mask = torch.triu(torch.ones((1, size, size), dtype=torch.uint8), diagonal=1)
        return mask == 0

DecoderLinearHead

Bases: Module

Original linear decoder box head.

Source code in models/ltnet/src/ltnet/modeling_lt_compatible.py
500
501
502
503
504
505
506
507
508
509
510
511
512
513
514
515
516
517
518
class DecoderLinearHead(nn.Module):
    """Original linear decoder box head."""

    def __init__(self, input_dim: int, box_dim: int) -> None:
        super().__init__()
        self.dense = nn.Linear(input_dim, box_dim)
        self.activation = nn.Sigmoid()

    def forward(
        self, x: Float[torch.Tensor, "batch sequence hidden"]
    ) -> tuple[
        Float[torch.Tensor, "batch sequence 2"],
        Float[torch.Tensor, "batch sequence 2"],
        None,
        None,
        None,
    ]:
        x = self.activation(self.dense(x))
        return x[:, :, 2:], x[:, :, :2], None, None, None

LinearHead

Bases: Module

Original linear refinement box head.

Source code in models/ltnet/src/ltnet/modeling_lt_compatible.py
521
522
523
524
525
526
527
528
529
530
531
532
533
534
535
536
537
538
539
540
541
542
543
544
545
546
547
548
class LinearHead(nn.Module):
    """Original linear refinement box head."""

    def __init__(
        self, input_dim: int, box_dim: int, output_dim: int, box_emb_size: int
    ) -> None:
        super().__init__()
        self.box_emb_size = box_emb_size
        self.box_embedding = nn.Linear(box_dim, self.box_emb_size)
        self.dense = nn.Linear(input_dim + self.box_emb_size, self.box_emb_size)
        self.feed_forward = nn.Linear(self.box_emb_size, output_dim)
        self.activation = nn.Sigmoid()

    def forward(
        self,
        x: Float[torch.Tensor, "batch sequence hidden"],
        box: Float[torch.Tensor, "batch sequence 4"],
    ) -> tuple[
        Float[torch.Tensor, "batch sequence 2"],
        Float[torch.Tensor, "batch sequence 2"],
        None,
        None,
        None,
    ]:
        box_embed = self.box_embedding(box)
        x = self.dense(torch.cat((x, box_embed), dim=-1))
        x = self.activation(self.feed_forward(x + box_embed))
        return x[:, :, 2:], x[:, :, :2], None, None, None

GMMHead

Bases: Module

Original GMM box head with optional generator plumbing.

Source code in models/ltnet/src/ltnet/modeling_lt_compatible.py
551
552
553
554
555
556
557
558
559
560
561
562
563
564
565
566
567
568
569
570
571
572
573
574
575
576
577
578
579
580
581
582
583
584
585
586
587
588
589
590
591
592
593
594
595
596
597
598
599
600
601
602
603
604
605
606
607
608
609
610
611
612
613
614
615
616
617
618
619
620
621
622
623
624
625
626
627
628
629
630
631
632
633
634
635
636
637
638
639
640
641
642
643
644
645
646
647
648
649
650
651
652
653
654
655
656
657
658
659
660
661
662
663
664
665
666
667
668
669
670
671
672
673
674
675
676
677
678
679
680
681
682
683
684
685
686
687
688
689
690
691
692
693
694
695
696
697
698
699
700
701
702
703
704
705
706
707
708
709
710
711
712
713
714
715
716
717
718
719
720
721
722
723
724
725
726
727
728
729
730
731
732
733
734
735
736
737
738
739
740
741
742
743
744
745
746
747
748
749
750
751
752
753
754
755
756
757
758
759
760
761
762
763
764
765
766
767
768
769
770
771
772
773
774
775
776
777
778
779
780
781
782
783
784
785
786
787
788
789
790
791
792
793
794
795
796
797
798
799
800
801
802
803
804
805
class GMMHead(nn.Module):
    """Original GMM box head with optional generator plumbing."""

    def __init__(
        self,
        hidden_size: int,
        *,
        condition: bool = False,
        x_softmax: bool = False,
        greedy: bool = False,
        config: LTNetConfig,
    ) -> None:
        super().__init__()
        self.hidden_size = hidden_size
        self.aug_size = max(1, hidden_size // 4)
        self.gmm_comp_num = 5
        self.gmm_param_num = 6
        self.xy_bivariate = nn.Linear(
            self.hidden_size, self.gmm_comp_num * self.gmm_param_num
        )
        self.condition = condition
        self.X_Sfotmax = x_softmax
        self.greedy = greedy
        self.xy_temperature = config.xy_temperature
        self.wh_temperature = config.wh_temperature
        if condition:
            self.xy_embedding = nn.Linear(2, self.aug_size)
            self.dropout = nn.Dropout(0.1)
            self.wh_bivariate = nn.Linear(
                self.hidden_size + self.aug_size,
                self.gmm_comp_num * self.gmm_param_num,
            )
        self.is_training = False

    def forward(
        self,
        x: Float[torch.Tensor, "batch sequence hidden"],
        generator: torch.Generator | None = None,
    ) -> tuple[
        Float[torch.Tensor, "batch sequence 2"],
        Float[torch.Tensor, "batch sequence 2"] | None,
        Float[torch.Tensor, "batch sequence gmm_params"],
        Float[torch.Tensor, "batch sequence gmm_params"] | None,
        Float[torch.Tensor, "batch sequence"] | None,
    ]:
        batch_size = x.size(0)
        xy_gmm = self.xy_bivariate(x)
        pi_xy, u_x, u_y, sigma_x, sigma_y, rho_xy = self.get_gmm_params(xy_gmm)
        sample_xy = self.sample_box(
            pi_xy,
            u_x,
            u_y,
            sigma_x,
            sigma_y,
            rho_xy,
            temp=self.xy_temperature,
            greedy=self.greedy,
            device=x.device,
            generator=generator,
        ).reshape(batch_size, -1, 2)
        sample_x = sample_xy[:, :, 0].unsqueeze(2).repeat(1, 1, self.gmm_comp_num)
        sample_y = sample_xy[:, :, 1].unsqueeze(2).repeat(1, 1, self.gmm_comp_num)
        xy_pdf = (
            self.batch_pdf(
                pi_xy,
                sample_x,
                sample_y,
                u_x,
                u_y,
                sigma_x,
                sigma_y,
                rho_xy,
                batch_size,
                self.gmm_comp_num,
                x.device,
            )
            if self.X_Sfotmax
            else None
        )
        if not self.condition:
            return sample_xy, None, xy_gmm, None, None
        xy_embed = self.dropout(self.xy_embedding(sample_xy))
        wh_gmm = self.wh_bivariate(torch.cat((x, xy_embed), dim=-1))
        pi_wh, u_w, u_h, sigma_w, sigma_h, rho_wh = self.get_gmm_params(wh_gmm)
        sample_wh = self.sample_box(
            pi_wh,
            u_w,
            u_h,
            sigma_w,
            sigma_h,
            rho_wh,
            temp=self.wh_temperature,
            greedy=self.greedy,
            device=x.device,
            generator=generator,
        ).reshape(batch_size, -1, 2)
        return sample_wh, sample_xy, wh_gmm, xy_gmm, xy_pdf

    def get_gmm_params(
        self, gmm_params: Float[torch.Tensor, "batch sequence gmm_params"]
    ) -> tuple[
        Float[torch.Tensor, "items components"],
        Float[torch.Tensor, "items components"],
        Float[torch.Tensor, "items components"],
        Float[torch.Tensor, "items components"],
        Float[torch.Tensor, "items components"],
        Float[torch.Tensor, "items components"],
    ]:
        pi, u_x, u_y, sigma_x, sigma_y, rho_xy = torch.split(
            gmm_params, self.gmm_comp_num, dim=2
        )
        pi = nn.Softmax(dim=-1)(pi).reshape(-1, self.gmm_comp_num).detach().cpu()
        u_x = u_x.reshape(-1, self.gmm_comp_num).detach().cpu()
        u_y = u_y.reshape(-1, self.gmm_comp_num).detach().cpu()
        sigma_x = torch.exp(sigma_x).reshape(-1, self.gmm_comp_num).detach().cpu()
        sigma_y = torch.exp(sigma_y).reshape(-1, self.gmm_comp_num).detach().cpu()
        rho_xy = (
            torch.tanh(rho_xy)
            .clamp(min=-0.95, max=0.95)
            .reshape(-1, self.gmm_comp_num)
            .detach()
            .cpu()
        )
        return pi, u_x, u_y, sigma_x, sigma_y, rho_xy

    def sample_box(
        self,
        pi: Float[torch.Tensor, "items components"],
        u_x: Float[torch.Tensor, "items components"],
        u_y: Float[torch.Tensor, "items components"],
        sigma_x: Float[torch.Tensor, "items components"],
        sigma_y: Float[torch.Tensor, "items components"],
        rho_xy: Float[torch.Tensor, "items components"],
        *,
        temp: float | None,
        greedy: bool,
        device: torch.device,
        generator: torch.Generator | None = None,
    ) -> Float[torch.Tensor, "items 2"]:
        if temp is not None:
            pi = self.adjust_temp(pi, temp)
        try:
            sample_pi = pi
            generator_device = None if generator is None else generator.device
            if generator_device is not None:
                sample_pi = pi.to(generator_device)
            pi_idx = torch.multinomial(sample_pi, 1, generator=generator).cpu()
        except RuntimeError:
            pi_idx = torch.multinomial(pi, 1)
        except Exception:
            pi_idx = pi.argmax(1).unsqueeze(-1)
        u_x = torch.gather(u_x, dim=1, index=pi_idx)
        u_y = torch.gather(u_y, dim=1, index=pi_idx)
        sigma_x = torch.gather(sigma_x, dim=1, index=pi_idx)
        sigma_y = torch.gather(sigma_y, dim=1, index=pi_idx)
        rho_xy = torch.gather(rho_xy, dim=1, index=pi_idx)
        return self.sample_bivariate_normal(
            u_x,
            u_y,
            sigma_x,
            sigma_y,
            rho_xy,
            temp,
            greedy=greedy,
            device=device,
            generator=generator,
        )

    @staticmethod
    def adjust_temp(
        pi_pdf: Float[torch.Tensor, "items components"], temperature: float
    ) -> Float[torch.Tensor, "items components"]:
        pi_pdf = torch.log(pi_pdf) / temperature
        pi_pdf -= torch.max(pi_pdf)
        pi_pdf = torch.exp(pi_pdf)
        pi_pdf /= torch.sum(pi_pdf)
        return pi_pdf

    @staticmethod
    def sample_bivariate_normal(
        u_x: Float[torch.Tensor, "items components"],
        u_y: Float[torch.Tensor, "items components"],
        sigma_x: Float[torch.Tensor, "items components"],
        sigma_y: Float[torch.Tensor, "items components"],
        rho_xy: Float[torch.Tensor, "items components"],
        temperature: float | None,
        *,
        greedy: bool,
        device: torch.device,
        generator: torch.Generator | None = None,
    ) -> Float[torch.Tensor, "items 2"]:
        if greedy:
            return torch.cat((u_x, u_y), dim=-1).to(device)

        sample_device = u_x.device
        if generator is not None:
            sample_device = generator.device

        mean = torch.cat((u_x, u_y), dim=1).to(sample_device)
        scale = math.sqrt(1.0 if temperature is None else temperature)
        sigma_x *= scale
        sigma_y *= scale

        sigma_x = sigma_x.to(sample_device)
        sigma_y = sigma_y.to(sample_device)
        rho_xy = rho_xy.to(sample_device)

        cov = torch.zeros((u_x.size(0), 2, 2), device=sample_device)
        cov[:, 0, 0] = sigma_x.flatten() * sigma_x.flatten()
        cov[:, 0, 1] = rho_xy.flatten() * sigma_x.flatten() * sigma_y.flatten()
        cov[:, 1, 0] = rho_xy.flatten() * sigma_x.flatten() * sigma_y.flatten()
        cov[:, 1, 1] = sigma_y.flatten() * sigma_y.flatten()
        det = cov[:, 0, 0] * cov[:, 1, 1] - cov[:, 0, 1] * cov[:, 1, 0]

        for idx in (det == 0).nonzero():
            cov[idx] *= 0.0
            cov[idx, 0, 0] += 1.0
            cov[idx, 1, 1] += 1.0

        noise = torch.randn(mean.shape, generator=generator, device=sample_device)
        sample = mean + torch.bmm(
            torch.linalg.cholesky(cov), noise.unsqueeze(-1)
        ).squeeze(-1)
        return sample.to(device)

    @staticmethod
    def batch_pdf(
        pi_xy: Float[torch.Tensor, "items components"],
        x: Float[torch.Tensor, "batch sequence gmm_params"],
        y: Float[torch.Tensor, "batch sequence gmm_params"],
        u_x: Float[torch.Tensor, "items components"],
        u_y: Float[torch.Tensor, "items components"],
        sigma_x: Float[torch.Tensor, "items components"],
        sigma_y: Float[torch.Tensor, "items components"],
        rho_xy: Float[torch.Tensor, "items components"],
        batch_size: int,
        gmm_comp_num: int,
        device: torch.device,
    ) -> Float[torch.Tensor, "batch sequence"]:
        u_x = u_x.reshape(batch_size, -1, gmm_comp_num).to(device)
        u_y = u_y.reshape(batch_size, -1, gmm_comp_num).to(device)
        sigma_x = sigma_x.reshape(batch_size, -1, gmm_comp_num).to(device)
        sigma_y = sigma_y.reshape(batch_size, -1, gmm_comp_num).to(device)
        pi_xy = pi_xy.reshape(batch_size, -1, gmm_comp_num).to(device)
        rho_xy = rho_xy.reshape(batch_size, -1, gmm_comp_num).to(device)
        z_x = ((x - u_x) / sigma_x) ** 2
        z_y = ((y - u_y) / sigma_y) ** 2
        z_xy = (x - u_x) * (y - u_y) / (sigma_x * sigma_y)
        z = z_x + z_y - 2 * rho_xy * z_xy
        exp = torch.exp(-z / (2 * (1 - rho_xy**2)))
        norm = torch.clamp(
            2 * math.pi * sigma_x * sigma_y * torch.sqrt(1 - rho_xy**2),
            min=1e-5,
        )
        return torch.sum(pi_xy * exp / norm, dim=2).detach()

TransformerRefineLayer

Bases: Module

Original refinement transformer layer.

Source code in models/ltnet/src/ltnet/modeling_lt_compatible.py
808
809
810
811
812
813
814
815
816
817
818
819
820
821
822
823
824
825
826
827
828
829
830
831
832
833
834
835
836
837
838
839
840
class TransformerRefineLayer(nn.Module):
    """Original refinement transformer layer."""

    def __init__(
        self,
        size: int = 0,
        ff_size: int = 0,
        num_heads: int = 0,
        dropout: float = 0.1,
        sent_length: int = 128,
    ) -> None:
        super().__init__()
        self.layer_norm = nn.LayerNorm(size, eps=1e-6)
        self.box_norm = nn.LayerNorm(size // 4, eps=1e-6)
        self.src_src_att = CustomAttention(
            num_heads, size, dropout=dropout, sent_length=sent_length
        )
        self.combine_layer = nn.Linear(size + size // 4, size)
        self.feed_forward = PositionwiseFeedForward(size, ff_size=ff_size)
        self.dropout = nn.Dropout(dropout)
        self.size = size

    def forward(
        self,
        context: Float[torch.Tensor, "batch sequence hidden"],
        box: Float[torch.Tensor, "batch sequence box_features"],
        mask: Bool[torch.Tensor, "..."],
        xy_pdf_score: Float[torch.Tensor, "batch sequence"] | None,
    ) -> Float[torch.Tensor, "batch sequence hidden"]:
        context_norm = self.layer_norm(context)
        box_norm = self.box_norm(box)
        h = self.src_src_att(box_norm, context_norm, box_norm, mask, xy_pdf_score)
        return self.feed_forward(self.dropout(h) + context_norm)

RefineEncoder

Bases: Module

Original refinement encoder.

Source code in models/ltnet/src/ltnet/modeling_lt_compatible.py
843
844
845
846
847
848
849
850
851
852
853
854
855
856
857
858
859
860
861
862
863
864
865
866
867
868
869
870
871
872
873
874
875
876
877
878
879
880
881
882
883
class RefineEncoder(nn.Module):
    """Original refinement encoder."""

    def __init__(
        self,
        hidden_size: int,
        num_heads: int,
        dropout: float,
        box_dim: int,
        sent_length: int = 128,
    ) -> None:
        super().__init__()
        self.aug_size = max(1, hidden_size // 4)
        self.box_embedding = nn.Linear(box_dim, self.aug_size)
        self.layer = TransformerRefineLayer(
            size=hidden_size,
            ff_size=hidden_size * 4,
            num_heads=num_heads,
            dropout=dropout,
            sent_length=sent_length,
        )
        self.layer_norm = nn.LayerNorm(hidden_size, eps=1e-6)
        self.emb_dropout = nn.Dropout(p=dropout)
        self.box_dim = box_dim
        self._output_size = hidden_size
        self._hidden_size = hidden_size
        self.blank_box = torch.Tensor([2.0, 2.0, 2.0, 2.0])

    def forward(
        self,
        context: Float[torch.Tensor, "batch sequence hidden"],
        input_box: Float[torch.Tensor, "batch sequence 4"],
        mask: Bool[torch.Tensor, "..."],
        xy_pdf_score: Float[torch.Tensor, "batch sequence"] | None,
    ) -> Float[torch.Tensor, "batch sequence hidden"]:
        box = input_box.clone()
        box[:, :, : self.box_dim][~mask.squeeze(1)] = self.blank_box[: self.box_dim].to(
            box.device
        )
        box_embed = self.emb_dropout(self.box_embedding(box[:, :, : self.box_dim]))
        return self.layer_norm(self.layer(context, box_embed, mask, xy_pdf_score))

PDFDecoder

Bases: Module

Original LT-Net PDF decoder.

Source code in models/ltnet/src/ltnet/modeling_lt_compatible.py
 886
 887
 888
 889
 890
 891
 892
 893
 894
 895
 896
 897
 898
 899
 900
 901
 902
 903
 904
 905
 906
 907
 908
 909
 910
 911
 912
 913
 914
 915
 916
 917
 918
 919
 920
 921
 922
 923
 924
 925
 926
 927
 928
 929
 930
 931
 932
 933
 934
 935
 936
 937
 938
 939
 940
 941
 942
 943
 944
 945
 946
 947
 948
 949
 950
 951
 952
 953
 954
 955
 956
 957
 958
 959
 960
 961
 962
 963
 964
 965
 966
 967
 968
 969
 970
 971
 972
 973
 974
 975
 976
 977
 978
 979
 980
 981
 982
 983
 984
 985
 986
 987
 988
 989
 990
 991
 992
 993
 994
 995
 996
 997
 998
 999
1000
1001
1002
1003
1004
1005
1006
1007
1008
1009
1010
1011
1012
1013
1014
1015
1016
1017
1018
1019
1020
1021
1022
1023
1024
1025
1026
1027
1028
1029
1030
1031
1032
1033
1034
1035
1036
1037
1038
1039
1040
1041
1042
1043
1044
1045
1046
1047
1048
1049
1050
1051
1052
1053
1054
1055
1056
1057
1058
class PDFDecoder(nn.Module):
    """Original LT-Net PDF decoder."""

    def __init__(
        self,
        *,
        box_dim: int = 4,
        hidden_size: int = 256,
        num_layers: int = 2,
        attn_heads: int = 2,
        dropout: float = 0.1,
        config: LTNetConfig,
    ) -> None:
        super().__init__()
        self.hidden_size = hidden_size
        self.schedule_sample = config.decoder_schedule_sample
        self.global_feature = config.decoder_global_feature
        self.aug_size = max(1, hidden_size // 4)
        self.box_embedding = nn.Linear(box_dim, self.aug_size)
        output_input = 2 * hidden_size + self.aug_size
        if not self.global_feature:
            output_input = hidden_size + self.aug_size
        self.output_Layer = nn.Linear(output_input, hidden_size)
        self.latent_transformer = nn.Linear(hidden_size, hidden_size - self.aug_size)
        self.decoder = CustomTransformerDecoder(
            hidden_size=hidden_size,
            hidden_bb_size=self.aug_size,
            ff_size=hidden_size * 4,
            num_layers=num_layers,
            num_heads=attn_heads,
            dropout=dropout,
            emb_dropout=dropout,
        )
        if config.decoder_head_type.upper() == "GMM":
            self.box_predictor = GMMHead(
                hidden_size,
                condition=True,
                x_softmax=config.refine_x_softmax,
                greedy=config.decoder_greedy,
                config=config,
            )
        else:
            self.box_predictor = DecoderLinearHead(hidden_size, 4)

    def random_sample(
        self,
        output_box: Float[torch.Tensor, "batch box_sequence 4"],
        pred_box: Float[torch.Tensor, "batch box_sequence 4"],
        sample_num: int,
    ) -> Float[torch.Tensor, "batch box_sequence 4"]:
        length = torch.arange(output_box.size(1))
        index = torch.Tensor(random.sample(list(enumerate(length)), sample_num))[
            :, 0
        ].long()
        mask = torch.zeros(
            output_box.size(), dtype=torch.bool, device=output_box.device
        )
        mask[:, index] = 1
        output_box[mask] = pred_box[mask]
        return output_box

    def forward(
        self,
        output_box: Float[torch.Tensor, "batch box_sequence 4"],
        output_context: Float[torch.Tensor, "batch context_sequence hidden"],
        encoder_output: Float[torch.Tensor, "batch sequence hidden"],
        src_mask: Bool[torch.Tensor, "..."],
        trg_mask: Bool[torch.Tensor, "..."],
        src: Float[torch.Tensor, "batch sequence hidden"],
        class_embeds: Float[torch.Tensor, "batch sequence hidden"],
        epoch: int = 0,
        is_train: bool = True,
        global_mask: Bool[torch.Tensor, "..."] | None = None,
        generator: torch.Generator | None = None,
    ) -> tuple[
        Float[torch.Tensor, "batch sequence hidden"],
        Float[torch.Tensor, "batch sequence 2"],
        Float[torch.Tensor, "batch sequence 2"],
        Float[torch.Tensor, "batch sequence gmm_params"] | None,
        Float[torch.Tensor, "batch sequence gmm_params"] | None,
        Float[torch.Tensor, "batch sequence"] | None,
    ]:
        _ = (src, class_embeds)
        output_box_c = output_box.clone()
        if global_mask is None:
            global_mask = torch.ones(
                encoder_output.shape[:2], dtype=torch.bool, device=encoder_output.device
            )
        if is_train:
            global_feature = encoder_output.clone()
            global_feature[~global_mask] = float("-inf")
            global_feature = torch.max(global_feature, dim=1).values
            global_feature = global_feature.unsqueeze(1).repeat(
                1, encoder_output.size(1), 1
            )
            pair_count = min(
                output_box_c[:, 2::2, :].size(1), output_box_c[:, 1::2, :].size(1)
            )
            output_box_c[:, 2 : 2 + 2 * pair_count : 2, :] = output_box_c[
                :, 1 : 1 + 2 * pair_count : 2, :
            ]
        else:
            global_feature = output_context.clone()
            global_feature[~global_mask[:, : output_context.size(1)]] = float("-inf")
            global_feature = torch.max(global_feature, dim=1).values
            global_feature = global_feature.unsqueeze(1).repeat(
                1, encoder_output.size(1), 1
            )
            if (output_box_c.size(1) - 1) % 2 == 0 and output_box_c.size(1) > 1:
                pair_count = min(
                    output_box_c[:, 2::2, :].size(1),
                    output_box_c[:, 1::2, :].size(1),
                )
                output_box_c[:, 2 : 2 + 2 * pair_count : 2, :] = output_box_c[
                    :, 1 : 1 + 2 * pair_count : 2, :
                ]
            elif (output_box_c.size(1) - 1) % 2 != 0 and output_box_c.size(1) > 1:
                output_box_c[:, 2::2, :] = output_box_c[:, 1:-1:2, :]
        output_box_embed = self.box_embedding(output_box_c)
        decoder_output = self.decoder(
            trg_embed_0=output_box_embed,
            trg_embed_1=encoder_output[:, 1:, :],
            encoder_output=encoder_output,
            encoder_hidden=None,
            src_mask=src_mask,
            unroll_steps=output_box_embed.size(1),
            hidden=None,
            trg_mask=trg_mask,
        )
        trg_input = torch.cat((encoder_output[:, :-1, :], output_box_embed), dim=-1)
        decoder_output = torch.cat(
            (trg_input[:, 0, :].unsqueeze(1), decoder_output), dim=1
        )
        if self.global_feature:
            decoder_output = torch.cat((decoder_output, global_feature), dim=-1)
        box_predictor_input = self.output_Layer(decoder_output)
        sample_wh, sample_xy, wh_gmm, xy_gmm, xy_pdf = self.box_predictor(
            box_predictor_input, generator=generator
        )
        if is_train and self.schedule_sample:
            pred_box = torch.cat((sample_xy, sample_wh), dim=-1)
            pred_box = pred_box[:, :-1]
            pred_box[:, 2::2, :] = pred_box[:, 1::2, :]
            sample_pred_num = int(
                pred_box.size(1) * (1.0 - ((epoch + 1) / 50.0) ** 1.2)
            )
            if sample_pred_num >= pred_box.size(1) / 3.0:
                sample_pred_num = int(pred_box.size(1) / 3.0)
            mix_output_box = self.random_sample(output_box_c, pred_box, sample_pred_num)
            mix_output_box_embed = self.box_embedding(mix_output_box)
            decoder_output = self.decoder(
                trg_embed_0=mix_output_box_embed,
                trg_embed_1=encoder_output[:, 1:, :],
                encoder_output=encoder_output,
                encoder_hidden=None,
                src_mask=src_mask,
                unroll_steps=mix_output_box_embed.size(1),
                hidden=None,
                trg_mask=trg_mask,
            )
            trg_input = torch.cat(
                (encoder_output[:, :-1, :], mix_output_box_embed), dim=-1
            )
            decoder_output = torch.cat(
                (trg_input[:, 0, :].unsqueeze(1), decoder_output), dim=1
            )
            if self.global_feature:
                decoder_output = torch.cat((decoder_output, global_feature), dim=-1)
            box_predictor_input = self.output_Layer(decoder_output)
            sample_wh, sample_xy, wh_gmm, xy_gmm, xy_pdf = self.box_predictor(
                box_predictor_input, generator=generator
            )
        return box_predictor_input, sample_wh, sample_xy, wh_gmm, xy_gmm, xy_pdf

BBoxHead

Bases: Module

Original LT-Net bbox head.

Source code in models/ltnet/src/ltnet/modeling_lt_compatible.py
1061
1062
1063
1064
1065
1066
1067
1068
1069
1070
1071
1072
1073
1074
1075
1076
1077
1078
1079
1080
1081
1082
1083
1084
1085
1086
1087
1088
1089
1090
1091
1092
1093
1094
1095
1096
1097
1098
1099
1100
1101
1102
1103
1104
1105
1106
1107
1108
1109
1110
1111
1112
1113
1114
1115
1116
1117
1118
1119
1120
1121
1122
1123
1124
1125
1126
1127
1128
1129
1130
1131
1132
1133
1134
1135
1136
1137
1138
1139
1140
1141
1142
1143
1144
1145
1146
1147
1148
1149
1150
1151
1152
1153
1154
1155
1156
1157
1158
1159
1160
1161
1162
1163
1164
1165
1166
1167
1168
1169
1170
1171
1172
1173
1174
1175
1176
1177
1178
1179
1180
1181
1182
1183
1184
1185
1186
1187
1188
1189
1190
1191
1192
1193
1194
1195
1196
1197
1198
1199
1200
1201
1202
1203
1204
1205
1206
1207
1208
1209
1210
1211
1212
1213
1214
1215
1216
1217
1218
1219
1220
1221
1222
1223
1224
1225
1226
1227
1228
1229
1230
1231
1232
1233
1234
1235
class BBoxHead(nn.Module):
    """Original LT-Net bbox head."""

    def __init__(self, config: LTNetConfig) -> None:
        super().__init__()
        self.pad_index = 0
        self.bos_index = 1
        self.eos_index = 2
        self.box_dim = 4
        self.cfg = _cfg(config)
        self.Decoder = PDFDecoder(
            hidden_size=config.hidden_size,
            num_layers=2,
            attn_heads=2,
            dropout=config.dropout,
            config=config,
        )
        self.refine_module = config.refine
        if config.refine:
            self.refine_encoder = RefineEncoder(
                hidden_size=config.hidden_size,
                num_heads=1,
                dropout=config.dropout,
                box_dim=self.box_dim,
                sent_length=max(1, config.max_sequence_length // 2),
            )
            if config.refine_head_type.title() == "Linear":
                self.refine_box_head = LinearHead(
                    config.hidden_size,
                    self.box_dim,
                    4,
                    max(1, config.hidden_size // 4),
                )
            elif config.refine_head_type.upper() == "GMM":
                self.refine_box_head = GMMHead(
                    config.hidden_size,
                    condition=True,
                    x_softmax=False,
                    greedy=False,
                    config=config,
                )

    def forward(
        self,
        epoch: int,
        encoder_output: Float[torch.Tensor, "batch sequence hidden"],
        mask: Bool[torch.Tensor, "..."],
        src: Float[torch.Tensor, "batch sequence hidden"],
        class_embeds: Float[torch.Tensor, "batch sequence hidden"],
        output_box: Float[torch.Tensor, "batch box_sequence 4"],
        trg_mask: Bool[torch.Tensor, "..."],
        global_mask: Bool[torch.Tensor, "..."],
        generator: torch.Generator | None = None,
    ) -> tuple[
        Float[torch.Tensor, "batch sequence 4"],
        Float[torch.Tensor, "batch sequence gmm_params"] | None,
        Float[torch.Tensor, "batch sequence 4"] | None,
        Float[torch.Tensor, "batch sequence gmm_params"] | None,
    ]:
        (
            decoder_output,
            coarse_wh,
            coarse_xy,
            coarse_wh_gmm,
            coarse_xy_gmm,
            xy_pdf_score,
        ) = self.Decoder(
            output_box,
            encoder_output[:, :-1, :],
            encoder_output,
            mask,
            trg_mask,
            src,
            class_embeds,
            epoch,
            global_mask=global_mask,
            generator=generator,
        )
        coarse_box = torch.cat((coarse_xy, coarse_wh), dim=-1)
        coarse_gmm = (
            torch.cat((coarse_xy_gmm, coarse_wh_gmm), dim=-1)
            if coarse_xy_gmm is not None and coarse_wh_gmm is not None
            else None
        )
        if not self.refine_module:
            return coarse_box, coarse_gmm, None, None
        if xy_pdf_score is not None:
            refine_context = self.refine_encoder(
                decoder_output[:, 1::2],
                coarse_box[:, 1::2],
                mask[:, :, 1::2],
                xy_pdf_score.detach()[:, 1::2],
            )
        else:
            refine_context = self.refine_encoder(
                decoder_output[:, 1::2], coarse_box[:, 1::2], mask[:, :, 1::2], None
            )
        refine_wh, refine_xy, refine_wh_gmm, refine_xy_gmm, _ = self.refine_box_head(
            refine_context, coarse_box[:, 1::2]
        )
        refine_box = torch.cat((refine_xy, refine_wh), dim=-1)
        all_refine_box = torch.zeros(coarse_box.size(), device=coarse_box.device)
        all_refine_box[:, 1::2] += refine_box
        refine_gmm = (
            torch.cat((refine_xy_gmm, refine_wh_gmm), dim=-1)
            if refine_xy_gmm is not None and refine_wh_gmm is not None
            else None
        )
        return coarse_box, coarse_gmm, all_refine_box, refine_gmm

    def inference(
        self,
        encoder_output: Float[torch.Tensor, "batch sequence hidden"],
        mask: Bool[torch.Tensor, "..."],
        src: Float[torch.Tensor, "batch sequence hidden"],
        class_embeds: Float[torch.Tensor, "batch sequence hidden"],
        global_mask: Bool[torch.Tensor, "..."],
        generator: torch.Generator | None = None,
    ) -> tuple[
        Float[torch.Tensor, "batch sequence 4"],
        Float[torch.Tensor, "batch sequence gmm_params"] | None,
        Float[torch.Tensor, "batch sequence 4"] | None,
        Float[torch.Tensor, "batch sequence gmm_params"] | None,
    ]:
        (
            decoder_output,
            coarse_wh,
            coarse_xy,
            coarse_wh_gmm,
            coarse_xy_gmm,
            xy_pdf_score,
        ) = greedy_pdf(
            src_mask=mask,
            bos_index=self.bos_index,
            eos_index=self.eos_index,
            max_output_length=128,
            decoder=self.Decoder,
            encoder_output=encoder_output,
            encoder_hidden=None,
            class_embeds=class_embeds,
            src=src,
            global_mask=global_mask,
            generator=generator,
        )
        coarse_box = torch.cat((coarse_xy, coarse_wh), dim=-1)
        coarse_gmm = (
            torch.cat((coarse_xy_gmm, coarse_wh_gmm), dim=-1)
            if coarse_xy_gmm is not None and coarse_wh_gmm is not None
            else None
        )
        if not self.refine_module:
            return coarse_box, coarse_gmm, None, None
        if xy_pdf_score is not None:
            refine_context = self.refine_encoder(
                decoder_output[:, 1::2],
                coarse_box[:, 1::2],
                mask[:, :, 1::2],
                xy_pdf_score.detach()[:, 1::2],
            )
        else:
            refine_context = self.refine_encoder(
                decoder_output[:, 1::2], coarse_box[:, 1::2], mask[:, :, 1::2], None
            )
        refine_wh, refine_xy, refine_wh_gmm, refine_xy_gmm, _ = self.refine_box_head(
            refine_context, coarse_box[:, 1::2]
        )
        refine_box = torch.cat((refine_xy, refine_wh), dim=-1)
        all_refine_box = torch.zeros(coarse_box.size(), device=coarse_box.device)
        all_refine_box[:, 1::2] += refine_box
        refine_gmm = (
            torch.cat((refine_xy_gmm, refine_wh_gmm), dim=-1)
            if refine_xy_gmm is not None and refine_wh_gmm is not None
            else None
        )
        return coarse_box, coarse_gmm, all_refine_box, refine_gmm

modeling_ltnet

PyTorch model wrapper for LT-Net.

LTNetModelOutput dataclass

Bases: ModelOutput

Raw LT-Net model outputs.

Attributes:

Name Type Description
vocab_logits Float[Tensor, 'batch sequence vocab'] | None

Mixed token vocabulary logits.

obj_id_logits Float[Tensor, 'batch sequence object_classes'] | None

Object-id classifier logits.

token_type_logits Float[Tensor, 'batch sequence token_types'] | None

Token-type classifier logits.

coarse_box Float[Tensor, 'batch sequence 4'] | None

Coarse normalized center xywh boxes.

coarse_gmm Float[Tensor, 'batch sequence gmm_params'] | None

Optional coarse GMM parameters.

refine_box Float[Tensor, 'batch sequence 4'] | None

Optional refined normalized center xywh boxes.

refine_gmm Float[Tensor, 'batch sequence gmm_params'] | None

Optional refinement GMM parameters.

hidden_states Float[Tensor, 'batch sequence hidden'] | None

Optional encoder hidden states.

Examples:

>>> import torch
>>> out = LTNetModelOutput(
...     vocab_logits=torch.zeros(1, 2, 3),
...     obj_id_logits=torch.zeros(1, 2, 4),
...     token_type_logits=torch.zeros(1, 2, 4),
...     coarse_box=torch.zeros(1, 2, 4),
... )
>>> out.coarse_box.shape
torch.Size([1, 2, 4])
Source code in models/ltnet/src/ltnet/modeling_ltnet.py
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
@dataclass
class LTNetModelOutput(ModelOutput):
    """Raw LT-Net model outputs.

    Attributes:
        vocab_logits: Mixed token vocabulary logits.
        obj_id_logits: Object-id classifier logits.
        token_type_logits: Token-type classifier logits.
        coarse_box: Coarse normalized center ``xywh`` boxes.
        coarse_gmm: Optional coarse GMM parameters.
        refine_box: Optional refined normalized center ``xywh`` boxes.
        refine_gmm: Optional refinement GMM parameters.
        hidden_states: Optional encoder hidden states.

    Examples:
        >>> import torch
        >>> out = LTNetModelOutput(
        ...     vocab_logits=torch.zeros(1, 2, 3),
        ...     obj_id_logits=torch.zeros(1, 2, 4),
        ...     token_type_logits=torch.zeros(1, 2, 4),
        ...     coarse_box=torch.zeros(1, 2, 4),
        ... )
        >>> out.coarse_box.shape
        torch.Size([1, 2, 4])
    """

    vocab_logits: Float[torch.Tensor, "batch sequence vocab"] | None = None
    obj_id_logits: Float[torch.Tensor, "batch sequence object_classes"] | None = None
    token_type_logits: Float[torch.Tensor, "batch sequence token_types"] | None = None
    coarse_box: Float[torch.Tensor, "batch sequence 4"] | None = None
    coarse_gmm: Float[torch.Tensor, "batch sequence gmm_params"] | None = None
    refine_box: Float[torch.Tensor, "batch sequence 4"] | None = None
    refine_gmm: Float[torch.Tensor, "batch sequence gmm_params"] | None = None
    hidden_states: Float[torch.Tensor, "batch sequence hidden"] | None = None

LTNetForLayoutGeneration

Bases: PreTrainedModel

Transformers PreTrainedModel for LT-Net relation-to-layout inference.

Source code in models/ltnet/src/ltnet/modeling_ltnet.py
 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
class LTNetForLayoutGeneration(PreTrainedModel):
    """Transformers ``PreTrainedModel`` for LT-Net relation-to-layout inference."""

    config_class = LTNetConfig
    base_model_prefix = "ltnet"
    main_input_name = "input_token"
    _tied_weights_keys: dict[str, str] = {}

    def __init__(self, config: LTNetConfig) -> None:
        """Initialize relation encoder and bbox head."""
        super().__init__(config)
        self.encoder = RelEncoder(config)
        self.bbox_head = BBoxHead(config)
        self.all_tied_weights_keys = dict(self._tied_weights_keys)

    def forward(
        self,
        input_token: Int[torch.Tensor, "batch sequence"],
        input_obj_id: Int[torch.Tensor, "batch sequence"],
        segment_label: Int[torch.Tensor, "batch sequence"],
        token_type: Int[torch.Tensor, "batch sequence"],
        src_mask: Bool[torch.Tensor, "batch 1 sequence"] | None = None,
        global_mask: Bool[torch.Tensor, "batch sequence"] | None = None,
        bbox: Float[torch.Tensor, "batch sequence 4"] | None = None,
        bbox_mask: Bool[torch.Tensor, "batch sequence"] | None = None,
        inference: bool = False,
        generator: torch.Generator | None = None,
        output_hidden_states: bool = False,
        return_dict: bool = True,
    ) -> (
        LTNetModelOutput
        | tuple[Float[torch.Tensor, "batch sequence feature"] | None, ...]
    ):
        """Run LT-Net relation encoding and bbox prediction.

        Args:
            input_token: Mixed object/predicate token ids.
            input_obj_id: Stable object ids for object-token positions.
            segment_label: Relation segment ids.
            token_type: Token type ids ``0/1/2/3``.
            src_mask: Valid-token mask shaped ``(batch, 1, sequence)``.
            global_mask: Optional reference global-feature mask.
            bbox: Optional teacher-forced boxes for training/parity paths.
            bbox_mask: Optional box validity mask, reserved for parity paths.
            inference: Compatibility flag; greedy decoding is pipeline-owned.
            generator: Optional PyTorch generator for stochastic GMM sampling.
            output_hidden_states: Whether to include encoder hidden states.
            return_dict: Whether to return a ``ModelOutput``.

        Returns:
            Raw LT-Net output dataclass or tuple.

        Raises:
            ValueError: If required tensor shapes are invalid.
        """
        _ = bbox_mask
        if input_token.shape != input_obj_id.shape:
            raise ValueError("input_token and input_obj_id must have the same shape")

        effective_src_mask = (
            src_mask
            if src_mask is not None
            else input_token.ne(self.config.pad_token_id).unsqueeze(1)
        )
        if effective_src_mask.ndim == 2:
            effective_src_mask = effective_src_mask.unsqueeze(1)
        encoder_outputs = self.encoder(
            input_token,
            input_obj_id,
            segment_label,
            token_type,
            effective_src_mask,
        )
        (
            hidden_states,
            vocab_logits,
            obj_id_logits,
            token_type_logits,
            src,
            class_embeds,
        ) = encoder_outputs
        effective_global_mask = (
            global_mask if global_mask is not None else input_token.ge(2)
        )
        if inference:
            coarse_box, coarse_gmm, refine_box, refine_gmm = self.bbox_head.inference(
                hidden_states,
                effective_src_mask,
                src,
                class_embeds,
                effective_global_mask,
                generator=generator,
            )
        else:
            if bbox is None:
                bbox = input_token.new_full(
                    (input_token.size(0), input_token.size(1) - 1, 4), 2.0
                ).float()
            elif bbox.size(1) == input_token.size(1):
                bbox = bbox[:, :-1, :]
            trg_mask = effective_src_mask.new_ones([1, 1, 1])
            coarse_box, coarse_gmm, refine_box, refine_gmm = self.bbox_head(
                0,
                hidden_states,
                effective_src_mask,
                src,
                class_embeds,
                bbox,
                trg_mask,
                effective_global_mask,
                generator=generator,
            )
        output = LTNetModelOutput(
            vocab_logits=vocab_logits,
            obj_id_logits=obj_id_logits,
            token_type_logits=token_type_logits,
            coarse_box=coarse_box,
            coarse_gmm=coarse_gmm,
            refine_box=refine_box,
            refine_gmm=refine_gmm,
            hidden_states=hidden_states if output_hidden_states else None,
        )
        if return_dict:
            return output
        return output.to_tuple()

    @torch.no_grad()
    def _generate_boxes(
        self,
        input_token: Int[torch.Tensor, "batch sequence"],
        input_obj_id: Int[torch.Tensor, "batch sequence"],
        segment_label: Int[torch.Tensor, "batch sequence"],
        token_type: Int[torch.Tensor, "batch sequence"],
        src_mask: Bool[torch.Tensor, "batch 1 sequence"] | None = None,
        global_mask: Bool[torch.Tensor, "batch sequence"] | None = None,
        generator: torch.Generator | None = None,
    ) -> LTNetModelOutput:
        """Private pipeline helper for layout-level generation."""
        output = self(
            input_token=input_token,
            input_obj_id=input_obj_id,
            segment_label=segment_label,
            token_type=token_type,
            src_mask=src_mask,
            global_mask=global_mask,
            inference=True,
            generator=generator,
            return_dict=True,
        )
        return cast(LTNetModelOutput, output)

__init__

__init__(config: LTNetConfig) -> None

Initialize relation encoder and bbox head.

Source code in models/ltnet/src/ltnet/modeling_ltnet.py
61
62
63
64
65
66
def __init__(self, config: LTNetConfig) -> None:
    """Initialize relation encoder and bbox head."""
    super().__init__(config)
    self.encoder = RelEncoder(config)
    self.bbox_head = BBoxHead(config)
    self.all_tied_weights_keys = dict(self._tied_weights_keys)

forward

forward(
    input_token: Int[Tensor, "batch sequence"],
    input_obj_id: Int[Tensor, "batch sequence"],
    segment_label: Int[Tensor, "batch sequence"],
    token_type: Int[Tensor, "batch sequence"],
    src_mask: Bool[Tensor, "batch 1 sequence"]
    | None = None,
    global_mask: Bool[Tensor, "batch sequence"]
    | None = None,
    bbox: Float[Tensor, "batch sequence 4"] | None = None,
    bbox_mask: Bool[Tensor, "batch sequence"] | None = None,
    inference: bool = False,
    generator: Generator | None = None,
    output_hidden_states: bool = False,
    return_dict: bool = True,
) -> (
    LTNetModelOutput
    | tuple[
        Float[torch.Tensor, "batch sequence feature"]
        | None,
        ...,
    ]
)

Run LT-Net relation encoding and bbox prediction.

Parameters:

Name Type Description Default
input_token Int[Tensor, 'batch sequence']

Mixed object/predicate token ids.

required
input_obj_id Int[Tensor, 'batch sequence']

Stable object ids for object-token positions.

required
segment_label Int[Tensor, 'batch sequence']

Relation segment ids.

required
token_type Int[Tensor, 'batch sequence']

Token type ids 0/1/2/3.

required
src_mask Bool[Tensor, 'batch 1 sequence'] | None

Valid-token mask shaped (batch, 1, sequence).

None
global_mask Bool[Tensor, 'batch sequence'] | None

Optional reference global-feature mask.

None
bbox Float[Tensor, 'batch sequence 4'] | None

Optional teacher-forced boxes for training/parity paths.

None
bbox_mask Bool[Tensor, 'batch sequence'] | None

Optional box validity mask, reserved for parity paths.

None
inference bool

Compatibility flag; greedy decoding is pipeline-owned.

False
generator Generator | None

Optional PyTorch generator for stochastic GMM sampling.

None
output_hidden_states bool

Whether to include encoder hidden states.

False
return_dict bool

Whether to return a ModelOutput.

True

Returns:

Type Description
LTNetModelOutput | tuple[Float[Tensor, 'batch sequence feature'] | None, ...]

Raw LT-Net output dataclass or tuple.

Raises:

Type Description
ValueError

If required tensor shapes are invalid.

Source code in models/ltnet/src/ltnet/modeling_ltnet.py
 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
def forward(
    self,
    input_token: Int[torch.Tensor, "batch sequence"],
    input_obj_id: Int[torch.Tensor, "batch sequence"],
    segment_label: Int[torch.Tensor, "batch sequence"],
    token_type: Int[torch.Tensor, "batch sequence"],
    src_mask: Bool[torch.Tensor, "batch 1 sequence"] | None = None,
    global_mask: Bool[torch.Tensor, "batch sequence"] | None = None,
    bbox: Float[torch.Tensor, "batch sequence 4"] | None = None,
    bbox_mask: Bool[torch.Tensor, "batch sequence"] | None = None,
    inference: bool = False,
    generator: torch.Generator | None = None,
    output_hidden_states: bool = False,
    return_dict: bool = True,
) -> (
    LTNetModelOutput
    | tuple[Float[torch.Tensor, "batch sequence feature"] | None, ...]
):
    """Run LT-Net relation encoding and bbox prediction.

    Args:
        input_token: Mixed object/predicate token ids.
        input_obj_id: Stable object ids for object-token positions.
        segment_label: Relation segment ids.
        token_type: Token type ids ``0/1/2/3``.
        src_mask: Valid-token mask shaped ``(batch, 1, sequence)``.
        global_mask: Optional reference global-feature mask.
        bbox: Optional teacher-forced boxes for training/parity paths.
        bbox_mask: Optional box validity mask, reserved for parity paths.
        inference: Compatibility flag; greedy decoding is pipeline-owned.
        generator: Optional PyTorch generator for stochastic GMM sampling.
        output_hidden_states: Whether to include encoder hidden states.
        return_dict: Whether to return a ``ModelOutput``.

    Returns:
        Raw LT-Net output dataclass or tuple.

    Raises:
        ValueError: If required tensor shapes are invalid.
    """
    _ = bbox_mask
    if input_token.shape != input_obj_id.shape:
        raise ValueError("input_token and input_obj_id must have the same shape")

    effective_src_mask = (
        src_mask
        if src_mask is not None
        else input_token.ne(self.config.pad_token_id).unsqueeze(1)
    )
    if effective_src_mask.ndim == 2:
        effective_src_mask = effective_src_mask.unsqueeze(1)
    encoder_outputs = self.encoder(
        input_token,
        input_obj_id,
        segment_label,
        token_type,
        effective_src_mask,
    )
    (
        hidden_states,
        vocab_logits,
        obj_id_logits,
        token_type_logits,
        src,
        class_embeds,
    ) = encoder_outputs
    effective_global_mask = (
        global_mask if global_mask is not None else input_token.ge(2)
    )
    if inference:
        coarse_box, coarse_gmm, refine_box, refine_gmm = self.bbox_head.inference(
            hidden_states,
            effective_src_mask,
            src,
            class_embeds,
            effective_global_mask,
            generator=generator,
        )
    else:
        if bbox is None:
            bbox = input_token.new_full(
                (input_token.size(0), input_token.size(1) - 1, 4), 2.0
            ).float()
        elif bbox.size(1) == input_token.size(1):
            bbox = bbox[:, :-1, :]
        trg_mask = effective_src_mask.new_ones([1, 1, 1])
        coarse_box, coarse_gmm, refine_box, refine_gmm = self.bbox_head(
            0,
            hidden_states,
            effective_src_mask,
            src,
            class_embeds,
            bbox,
            trg_mask,
            effective_global_mask,
            generator=generator,
        )
    output = LTNetModelOutput(
        vocab_logits=vocab_logits,
        obj_id_logits=obj_id_logits,
        token_type_logits=token_type_logits,
        coarse_box=coarse_box,
        coarse_gmm=coarse_gmm,
        refine_box=refine_box,
        refine_gmm=refine_gmm,
        hidden_states=hidden_states if output_hidden_states else None,
    )
    if return_dict:
        return output
    return output.to_tuple()

pipeline_ltnet

Pipeline wrapper for LT-Net relation-to-layout generation.

LTNetPipeline

Bases: LayoutGenerationPipeline

Compose an LT-Net model and processor for scene-graph layout inference.

Parameters:

Name Type Description Default
model LTNetForLayoutGeneration

Converted LT-Net model.

required
processor LTNetProcessor

Matching scene-graph processor/tokenizer.

required
config LTNetConfig | None

Optional root pipeline config. Defaults to model.config.

None

Examples:

>>> processor = LTNetProcessor.from_config()
>>> config = LTNetConfig(
...     vocab_size=processor.tokenizer.vocab_size,
...     hidden_size=32,
...     num_hidden_layers=1,
...     num_attention_heads=4,
... )
>>> pipe = LTNetPipeline(
...     model=LTNetForLayoutGeneration(config),
...     processor=processor,
... )
>>> pipe.config.model_type
'ltnet'
Source code in models/ltnet/src/ltnet/pipeline_ltnet.py
 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
231
232
233
234
235
236
237
class LTNetPipeline(LayoutGenerationPipeline):
    """Compose an LT-Net model and processor for scene-graph layout inference.

    Args:
        model: Converted LT-Net model.
        processor: Matching scene-graph processor/tokenizer.
        config: Optional root pipeline config. Defaults to ``model.config``.

    Examples:
        >>> processor = LTNetProcessor.from_config()
        >>> config = LTNetConfig(
        ...     vocab_size=processor.tokenizer.vocab_size,
        ...     hidden_size=32,
        ...     num_hidden_layers=1,
        ...     num_attention_heads=4,
        ... )
        >>> pipe = LTNetPipeline(
        ...     model=LTNetForLayoutGeneration(config),
        ...     processor=processor,
        ... )
        >>> pipe.config.model_type
        'ltnet'
    """

    config_class: ClassVar[type[PretrainedConfig]] = LTNetConfig
    component_specs: ClassVar[dict[str, PipelineComponentSpec]] = {
        "model": PipelineComponentSpec(
            attribute_name="model",
            loader=_load_model_component,
            marker_file="config.json",
        ),
        "processor": PipelineComponentSpec(
            attribute_name="processor",
            loader=_load_processor_component,
            marker_file="preprocessor_config.json",
            save_with_is_main_process=False,
        ),
    }

    config: LTNetConfig
    model: LTNetForLayoutGeneration
    processor: LTNetProcessor

    def __init__(
        self,
        model: LTNetForLayoutGeneration,
        processor: LTNetProcessor,
        config: LTNetConfig | None = None,
    ) -> None:
        """Initialize the pipeline with model and processor components."""
        super().__init__(config or model.config)
        self.config = config or model.config
        self.model = model
        self.processor = processor

    @classmethod
    def _from_pretrained_components(  # ty: ignore[invalid-method-override]
        cls,
        *,
        config: PretrainedConfig,
        components: Mapping[str, LTNetForLayoutGeneration | LTNetProcessor | None],
    ) -> "LTNetPipeline":
        """Build a pipeline from loaded root components."""
        return cls(
            config=cast(LTNetConfig, config),
            model=cast(LTNetForLayoutGeneration, components["model"]),
            processor=cast(LTNetProcessor, components["processor"]),
        )

    @torch.no_grad()
    def __call__(  # ty: ignore[invalid-method-override]
        self,
        *,
        batch_size: int = 1,
        seed: int | None = None,
        generator: torch.Generator | None = None,
        condition_type: ConditionType | str = ConditionType.relation,
        labels: Int[torch.Tensor, "batch elements"] | Sequence[int] | None = None,
        bbox: Float[torch.Tensor, "batch elements 4"]
        | Sequence[Sequence[float]]
        | None = None,
        mask: Bool[torch.Tensor, "batch elements"] | Sequence[bool] | None = None,
        num_elements: int | list[int] | Int[torch.Tensor, "batch"] | None = None,
        box_format: BoxFormat | str = BoxFormat.xywh,
        normalized: bool = True,
        canvas_size: tuple[int, int] | None = None,
        num_inference_steps: int | None = None,
        output_type: OutputType = "dataclass",
        return_intermediates: bool = False,
        scene_graph: SceneGraphInput | SceneGraphMapping | None = None,
        objects: Sequence[LayoutObject] | None = None,
        relations: Sequence[LayoutRelation] | None = None,
    ) -> (
        LayoutGenerationOutput
        | dict[
            str,
            Float[torch.Tensor, "..."]
            | Int[torch.Tensor, "..."]
            | Bool[torch.Tensor, "..."]
            | dict[int, str]
            | dict[str, Float[torch.Tensor, "..."] | None],
        ]
    ):
        """Generate layouts from a public scene graph.

        Args:
            batch_size: Number of layouts to generate from the same graph.
            seed: Convenience seed used only when ``generator`` is absent.
            generator: Optional PyTorch generator; takes precedence over seed.
            condition_type: Must normalize to ``relation``.
            labels: Reserved v1 interface input; LT-Net uses ``scene_graph``.
            bbox: Reserved v1 box constraint input.
            mask: Reserved v1 validity-mask input.
            num_elements: Reserved v1 element-count input.
            box_format: Output box format; LT-Net returns normalized ``xywh``.
            normalized: Whether output boxes should be normalized.
            canvas_size: Reserved denormalization canvas size.
            num_inference_steps: Reserved v1 step count.
            output_type: Return dataclass or dict.
            return_intermediates: Include raw logits/boxes in intermediates.
            scene_graph: Public relation payload.
            objects: Object-node shorthand when ``scene_graph`` is omitted.
            relations: Relation-edge shorthand when ``scene_graph`` is omitted.

        Returns:
            Layout generation output dataclass or dict.

        Raises:
            ValueError: If the condition or graph payload is unsupported.
        """
        _ = (labels, bbox, mask, num_elements, num_inference_steps)
        model_device = next(self.model.parameters()).device
        prepared_generator = self.prepare_generator(
            generator=generator, seed=seed, device=model_device
        )
        encoded = self.processor(
            scene_graph=scene_graph,
            objects=objects,
            relations=relations,
            batch_size=batch_size,
            condition_type=condition_type,
            return_tensors="pt",
        )
        model_inputs = {
            key: value.to(model_device)
            for key, value in encoded.items()
            if isinstance(value, torch.Tensor)
        }
        was_training = self.model.training
        self.model.eval()
        try:
            output = self.model._generate_boxes(
                **model_inputs,
                generator=prepared_generator,
            )
        finally:
            self.model.train(was_training)
        return self.processor.post_process_layout_generation(
            output,
            input_token=model_inputs["input_token"],
            input_obj_id=model_inputs["input_obj_id"],
            token_type=model_inputs["token_type"],
            box_format=box_format,
            normalized=normalized,
            canvas_size=canvas_size,
            output_type=output_type,
            return_intermediates=return_intermediates,
        )

__init__

__init__(
    model: LTNetForLayoutGeneration,
    processor: LTNetProcessor,
    config: LTNetConfig | None = None,
) -> None

Initialize the pipeline with model and processor components.

Source code in models/ltnet/src/ltnet/pipeline_ltnet.py
113
114
115
116
117
118
119
120
121
122
123
def __init__(
    self,
    model: LTNetForLayoutGeneration,
    processor: LTNetProcessor,
    config: LTNetConfig | None = None,
) -> None:
    """Initialize the pipeline with model and processor components."""
    super().__init__(config or model.config)
    self.config = config or model.config
    self.model = model
    self.processor = processor

__call__

__call__(
    *,
    batch_size: int = 1,
    seed: int | None = None,
    generator: Generator | None = None,
    condition_type: ConditionType
    | str = ConditionType.relation,
    labels: Int[Tensor, "batch elements"]
    | Sequence[int]
    | None = None,
    bbox: Float[Tensor, "batch elements 4"]
    | Sequence[Sequence[float]]
    | None = None,
    mask: Bool[Tensor, "batch elements"]
    | Sequence[bool]
    | None = None,
    num_elements: int
    | list[int]
    | Int[Tensor, "batch"]
    | None = None,
    box_format: BoxFormat | str = BoxFormat.xywh,
    normalized: bool = True,
    canvas_size: tuple[int, int] | None = None,
    num_inference_steps: int | None = None,
    output_type: OutputType = "dataclass",
    return_intermediates: bool = False,
    scene_graph: SceneGraphInput
    | SceneGraphMapping
    | None = None,
    objects: Sequence[LayoutObject] | None = None,
    relations: Sequence[LayoutRelation] | None = None,
) -> (
    LayoutGenerationOutput
    | dict[
        str,
        Float[torch.Tensor, "..."]
        | Int[torch.Tensor, "..."]
        | Bool[torch.Tensor, "..."]
        | dict[int, str]
        | dict[str, Float[torch.Tensor, "..."] | None],
    ]
)

Generate layouts from a public scene graph.

Parameters:

Name Type Description Default
batch_size int

Number of layouts to generate from the same graph.

1
seed int | None

Convenience seed used only when generator is absent.

None
generator Generator | None

Optional PyTorch generator; takes precedence over seed.

None
condition_type ConditionType | str

Must normalize to relation.

relation
labels Int[Tensor, 'batch elements'] | Sequence[int] | None

Reserved v1 interface input; LT-Net uses scene_graph.

None
bbox Float[Tensor, 'batch elements 4'] | Sequence[Sequence[float]] | None

Reserved v1 box constraint input.

None
mask Bool[Tensor, 'batch elements'] | Sequence[bool] | None

Reserved v1 validity-mask input.

None
num_elements int | list[int] | Int[Tensor, 'batch'] | None

Reserved v1 element-count input.

None
box_format BoxFormat | str

Output box format; LT-Net returns normalized xywh.

xywh
normalized bool

Whether output boxes should be normalized.

True
canvas_size tuple[int, int] | None

Reserved denormalization canvas size.

None
num_inference_steps int | None

Reserved v1 step count.

None
output_type OutputType

Return dataclass or dict.

'dataclass'
return_intermediates bool

Include raw logits/boxes in intermediates.

False
scene_graph SceneGraphInput | SceneGraphMapping | None

Public relation payload.

None
objects Sequence[LayoutObject] | None

Object-node shorthand when scene_graph is omitted.

None
relations Sequence[LayoutRelation] | None

Relation-edge shorthand when scene_graph is omitted.

None

Returns:

Type Description
LayoutGenerationOutput | dict[str, Float[Tensor, '...'] | Int[Tensor, '...'] | Bool[Tensor, '...'] | dict[int, str] | dict[str, Float[Tensor, '...'] | None]]

Layout generation output dataclass or dict.

Raises:

Type Description
ValueError

If the condition or graph payload is unsupported.

Source code in models/ltnet/src/ltnet/pipeline_ltnet.py
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
231
232
233
234
235
236
237
@torch.no_grad()
def __call__(  # ty: ignore[invalid-method-override]
    self,
    *,
    batch_size: int = 1,
    seed: int | None = None,
    generator: torch.Generator | None = None,
    condition_type: ConditionType | str = ConditionType.relation,
    labels: Int[torch.Tensor, "batch elements"] | Sequence[int] | None = None,
    bbox: Float[torch.Tensor, "batch elements 4"]
    | Sequence[Sequence[float]]
    | None = None,
    mask: Bool[torch.Tensor, "batch elements"] | Sequence[bool] | None = None,
    num_elements: int | list[int] | Int[torch.Tensor, "batch"] | None = None,
    box_format: BoxFormat | str = BoxFormat.xywh,
    normalized: bool = True,
    canvas_size: tuple[int, int] | None = None,
    num_inference_steps: int | None = None,
    output_type: OutputType = "dataclass",
    return_intermediates: bool = False,
    scene_graph: SceneGraphInput | SceneGraphMapping | None = None,
    objects: Sequence[LayoutObject] | None = None,
    relations: Sequence[LayoutRelation] | None = None,
) -> (
    LayoutGenerationOutput
    | dict[
        str,
        Float[torch.Tensor, "..."]
        | Int[torch.Tensor, "..."]
        | Bool[torch.Tensor, "..."]
        | dict[int, str]
        | dict[str, Float[torch.Tensor, "..."] | None],
    ]
):
    """Generate layouts from a public scene graph.

    Args:
        batch_size: Number of layouts to generate from the same graph.
        seed: Convenience seed used only when ``generator`` is absent.
        generator: Optional PyTorch generator; takes precedence over seed.
        condition_type: Must normalize to ``relation``.
        labels: Reserved v1 interface input; LT-Net uses ``scene_graph``.
        bbox: Reserved v1 box constraint input.
        mask: Reserved v1 validity-mask input.
        num_elements: Reserved v1 element-count input.
        box_format: Output box format; LT-Net returns normalized ``xywh``.
        normalized: Whether output boxes should be normalized.
        canvas_size: Reserved denormalization canvas size.
        num_inference_steps: Reserved v1 step count.
        output_type: Return dataclass or dict.
        return_intermediates: Include raw logits/boxes in intermediates.
        scene_graph: Public relation payload.
        objects: Object-node shorthand when ``scene_graph`` is omitted.
        relations: Relation-edge shorthand when ``scene_graph`` is omitted.

    Returns:
        Layout generation output dataclass or dict.

    Raises:
        ValueError: If the condition or graph payload is unsupported.
    """
    _ = (labels, bbox, mask, num_elements, num_inference_steps)
    model_device = next(self.model.parameters()).device
    prepared_generator = self.prepare_generator(
        generator=generator, seed=seed, device=model_device
    )
    encoded = self.processor(
        scene_graph=scene_graph,
        objects=objects,
        relations=relations,
        batch_size=batch_size,
        condition_type=condition_type,
        return_tensors="pt",
    )
    model_inputs = {
        key: value.to(model_device)
        for key, value in encoded.items()
        if isinstance(value, torch.Tensor)
    }
    was_training = self.model.training
    self.model.eval()
    try:
        output = self.model._generate_boxes(
            **model_inputs,
            generator=prepared_generator,
        )
    finally:
        self.model.train(was_training)
    return self.processor.post_process_layout_generation(
        output,
        input_token=model_inputs["input_token"],
        input_obj_id=model_inputs["input_obj_id"],
        token_type=model_inputs["token_type"],
        box_format=box_format,
        normalized=normalized,
        canvas_size=canvas_size,
        output_type=output_type,
        return_intermediates=return_intermediates,
    )

processing_ltnet

Processor for LT-Net scene graphs and layout outputs.

LTNetProcessor

Bases: ProcessorMixin

Normalize scene graphs, tokenize LT-Net inputs, and postprocess boxes.

Source code in models/ltnet/src/ltnet/processing_ltnet.py
 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
231
232
233
234
235
236
237
238
239
240
241
242
243
244
245
246
247
248
249
250
251
252
253
254
255
256
257
258
259
260
261
262
263
264
265
266
267
268
269
270
271
272
273
274
275
276
277
278
279
280
281
282
283
284
285
286
287
288
289
290
291
292
293
294
295
296
297
298
299
300
301
302
303
304
305
306
307
308
309
310
311
312
313
314
315
316
317
318
319
320
321
322
323
324
325
326
327
328
329
330
331
332
333
334
335
336
337
338
339
340
341
342
343
344
345
346
347
348
349
350
351
352
353
354
355
356
357
358
359
360
361
362
363
364
365
366
367
368
369
370
371
372
373
374
375
376
377
378
379
380
381
382
383
384
385
386
387
388
389
390
391
392
393
394
395
396
397
398
399
400
401
402
403
404
405
406
407
408
409
410
411
412
413
414
415
416
417
418
419
420
421
422
423
424
425
426
427
428
429
430
431
432
433
434
435
436
437
438
439
440
441
442
443
444
445
446
447
448
449
450
451
452
453
454
455
456
457
458
459
460
461
class LTNetProcessor(ProcessorMixin):
    """Normalize scene graphs, tokenize LT-Net inputs, and postprocess boxes."""

    attributes = ["tokenizer"]
    tokenizer_class = "LTNetRelationTokenizer"

    def __init__(
        self,
        tokenizer: LTNetRelationTokenizer,
        dataset_name: str = "coco",
        max_sequence_length: int = 128,
        id2label: Mapping[int, str] | Mapping[str, str] | None = None,
        relation_id2label: Mapping[int, str] | Mapping[str, str] | None = None,
        object_reduce: Literal["first", "last", "mean"] = "first",
    ) -> None:
        """Initialize processor label maps and tokenizer component."""
        self.tokenizer = tokenizer
        self.dataset_name = dataset_name
        self.max_sequence_length = max_sequence_length
        self.id2label = {
            int(key): str(value)
            for key, value in (id2label or DEFAULT_ID2LABEL).items()
        }
        self.relation_id2label = {
            int(key): str(value)
            for key, value in (relation_id2label or DEFAULT_RELATION_ID2LABEL).items()
        }
        self.label2id = {value.lower(): key for key, value in self.id2label.items()}
        self.relation_label2id = {
            value.lower(): key for key, value in self.relation_id2label.items()
        }
        if object_reduce not in {"first", "last", "mean"}:
            raise ValueError("object_reduce must be 'first', 'last', or 'mean'")

        self.object_reduce = object_reduce
        super().__init__(tokenizer=tokenizer)

    @classmethod
    def from_config(
        cls,
        *,
        dataset_name: str = "coco",
        max_sequence_length: int = 128,
        id2label: Mapping[int, str] | Mapping[str, str] | None = None,
        relation_id2label: Mapping[int, str] | Mapping[str, str] | None = None,
    ) -> "LTNetProcessor":
        """Construct a processor and tokenizer without external files.

        Returns:
            Processor with synthetic vocabulary derived from the label maps.

        Examples:
            >>> processor = LTNetProcessor.from_config()
            >>> processor.tokenizer.cls_token_id
            1
        """
        object_labels = {
            int(key): str(value)
            for key, value in (id2label or DEFAULT_ID2LABEL).items()
        }
        relation_labels = {
            int(key): str(value)
            for key, value in (relation_id2label or DEFAULT_RELATION_ID2LABEL).items()
        }
        tokens = ["__image__"]
        tokens.extend(value for _, value in sorted(object_labels.items()))
        tokens.extend(value for _, value in sorted(relation_labels.items()))
        tokenizer = LTNetRelationTokenizer(tokens=tokens)
        return cls(
            tokenizer=tokenizer,
            dataset_name=dataset_name,
            max_sequence_length=max_sequence_length,
            id2label=object_labels,
            relation_id2label=relation_labels,
        )

    @classmethod
    def _load_tokenizer_from_pretrained(
        cls,
        sub_processor_type: str,
        pretrained_model_name_or_path: str | PathLike[str],
        subfolder: str = "",
        **kwargs: str | int | float | bool | None,
    ) -> LTNetRelationTokenizer:
        """Load tokenizer for ``ProcessorMixin.from_pretrained``."""
        _ = sub_processor_type
        path = Path(pretrained_model_name_or_path)
        tokenizer_path = path / subfolder if subfolder else path
        token = kwargs.get("token")
        return LTNetRelationTokenizer.from_pretrained(
            tokenizer_path,
            cache_dir=cast(str | PathLike[str] | None, kwargs.get("cache_dir")),
            force_download=bool(kwargs.get("force_download", False)),
            local_files_only=bool(kwargs.get("local_files_only", False)),
            token=token if isinstance(token, str | bool) else None,
            revision=str(kwargs.get("revision", "main")),
        )

    def normalize_condition_type(
        self, condition_type: ConditionType | str
    ) -> ConditionType:
        """Normalize and validate the LT-Net public condition type."""
        condition = normalize_condition_type(condition_type)
        if condition is not ConditionType.relation:
            raise ValueError(
                "LT-Net only supports condition_type='relation' "
                "and aliases 'scene_graph', 'graph', or 'gen_r'."
            )

        return condition

    def _label_to_id(self, label: int | str) -> int:
        if isinstance(label, int):
            return label
        lowered = label.lower()
        if lowered in self.label2id:
            return self.label2id[lowered]
        raise ValueError(f"Unknown object label: {label}")

    def _relation_to_id(self, predicate: int | str) -> int:
        if isinstance(predicate, int):
            return predicate
        lowered = predicate.lower()
        if lowered in self.relation_label2id:
            return self.relation_label2id[lowered]
        raise ValueError(f"Unknown relation label: {predicate}")

    def _token_for_object(self, label_id: int) -> str:
        return self.id2label.get(label_id, str(label_id))

    def _token_for_relation(self, relation_id: int) -> str:
        return self.relation_id2label.get(relation_id, str(relation_id))

    def _normalize_scene_graph(
        self,
        scene_graph: SceneGraphInput | SceneGraphMapping | None,
        *,
        objects: Sequence[LayoutObject] | None,
        relations: Sequence[LayoutRelation] | None,
    ) -> SceneGraphInput:
        if isinstance(scene_graph, SceneGraphInput):
            return scene_graph
        if scene_graph is None:
            if objects is None:
                raise ValueError("scene_graph or objects must be provided")

            return SceneGraphInput(
                objects=tuple(objects),
                relations=tuple(relations or ()),
                id2label=self.id2label,
                relation_id2label=self.relation_id2label,
            )
        nodes = scene_graph.get("nodes", scene_graph.get("objects", ()))
        edges = scene_graph.get("edges", scene_graph.get("relations", ()))

        normalized_objects: list[LayoutObject] = []

        for node in cast(Sequence[SceneGraphItemMapping], nodes):
            item = node
            node_id = cast(int | str, item["id"])
            label = cast(int | str, item.get("label_id", item.get("label")))
            bbox = cast(tuple[float, float, float, float] | None, item.get("bbox"))
            normalized_objects.append(LayoutObject(id=node_id, label=label, bbox=bbox))

        normalized_relations: list[LayoutRelation] = []
        for edge in cast(Sequence[SceneGraphItemMapping], edges):
            item = edge
            subject = cast(int | str, item.get("source", item.get("subject")))
            predicate = cast(int | str, item.get("predicate_id", item.get("predicate")))
            target = cast(int | str, item.get("target", item.get("object")))
            normalized_relations.append(
                LayoutRelation(
                    subject=subject,
                    predicate=predicate,
                    object=target,
                    score=cast(float | None, item.get("score")),
                )
            )

        return SceneGraphInput(
            objects=tuple(normalized_objects),
            relations=tuple(normalized_relations),
            id2label=cast(dict[int, str] | None, scene_graph.get("id2label")),
            relation_id2label=cast(
                dict[int, str] | None,
                scene_graph.get("relation_id2label"),
            ),
        )

    def _serialize_graph(
        self,
        graph: SceneGraphInput,
    ) -> tuple[list[int], list[int], list[int], list[int], list[int]]:
        object_by_id = {item.id: item for item in graph.objects}
        object_ids = {item.id: idx + 1 for idx, item in enumerate(graph.objects)}
        tokens = [self.tokenizer.cls_token]
        input_obj_id = [0]
        segment_label = [0]
        token_type = [0]
        segment = 1
        for relation in graph.relations:
            subject = object_by_id[relation.subject]
            target = object_by_id[relation.object]
            subject_label = self._label_to_id(subject.label)
            target_label = self._label_to_id(target.label)
            relation_id = self._relation_to_id(relation.predicate)
            triple_tokens = [
                self._token_for_object(subject_label),
                self._token_for_relation(relation_id),
                self._token_for_object(target_label),
                self.tokenizer.sep_token,
            ]
            tokens.extend(triple_tokens)
            input_obj_id.extend([object_ids[subject.id], 0, object_ids[target.id], 0])
            segment_label.extend([segment] * 4)
            token_type.extend([1, 2, 3, 0])
            segment += 1
        if not graph.relations:
            for item in graph.objects:
                label_id = self._label_to_id(item.label)
                tokens.extend(
                    [self._token_for_object(label_id), self.tokenizer.sep_token]
                )
                input_obj_id.extend([object_ids[item.id], 0])
                segment_label.extend([segment, segment])
                token_type.extend([1, 0])
                segment += 1

        input_token = self.tokenizer.encode_scene_graph_tokens(tokens)
        length = min(len(input_token), self.max_sequence_length)
        input_token = input_token[:length]
        input_obj_id = input_obj_id[:length]
        segment_label = segment_label[:length]
        token_type = token_type[:length]

        src_mask = [1] * length
        pad_length = self.max_sequence_length - length
        input_token.extend([self.tokenizer.pad_token_id] * pad_length)
        input_obj_id.extend([0] * pad_length)
        segment_label.extend([0] * pad_length)
        token_type.extend([0] * pad_length)
        src_mask.extend([0] * pad_length)
        return input_token, input_obj_id, segment_label, token_type, src_mask

    def __call__(
        self,
        *,
        scene_graph: SceneGraphInput | SceneGraphMapping | None = None,
        objects: Sequence[LayoutObject] | None = None,
        relations: Sequence[LayoutRelation] | None = None,
        batch_size: int = 1,
        condition_type: ConditionType | str = ConditionType.relation,
        return_tensors: Literal["pt", "np"] = "pt",
        max_sequence_length: int | None = None,
    ) -> BatchEncoding:
        """Build LT-Net model tensors from public scene-graph inputs."""
        self.normalize_condition_type(condition_type)
        original_max_length = self.max_sequence_length
        try:
            if max_sequence_length is not None:
                self.max_sequence_length = max_sequence_length
            graph = self._normalize_scene_graph(
                scene_graph,
                objects=objects,
                relations=relations,
            )
            rows = [self._serialize_graph(graph) for _ in range(batch_size)]
        finally:
            self.max_sequence_length = original_max_length
        data = {
            "input_token": torch.tensor([row[0] for row in rows], dtype=torch.long),
            "input_obj_id": torch.tensor([row[1] for row in rows], dtype=torch.long),
            "segment_label": torch.tensor([row[2] for row in rows], dtype=torch.long),
            "token_type": torch.tensor([row[3] for row in rows], dtype=torch.long),
            "src_mask": torch.tensor(
                [row[4] for row in rows], dtype=torch.bool
            ).unsqueeze(1),
            "global_mask": torch.tensor([row[0] for row in rows], dtype=torch.long).ge(
                2
            ),
        }
        if return_tensors == "pt":
            return BatchEncoding(data)
        if return_tensors == "np":
            return BatchEncoding({key: value.numpy() for key, value in data.items()})
        raise ValueError("return_tensors must be 'pt' or 'np'")

    def post_process_layout_generation(
        self,
        model_outputs: LTNetModelOutput,
        *,
        input_token: Int[torch.Tensor, "batch sequence"] | None = None,
        input_obj_id: Int[torch.Tensor, "batch sequence"],
        token_type: Int[torch.Tensor, "batch sequence"],
        box_format: BoxFormat | str = BoxFormat.xywh,
        normalized: bool = True,
        canvas_size: tuple[int, int] | None = None,
        output_type: OutputType = "dataclass",
        return_intermediates: bool = False,
    ) -> (
        LayoutGenerationOutput
        | dict[
            str,
            Float[torch.Tensor, "..."]
            | Int[torch.Tensor, "..."]
            | Bool[torch.Tensor, "..."]
            | dict[int, str]
            | dict[str, Float[torch.Tensor, "..."] | None],
        ]
    ):
        """Convert raw token-level boxes into public object-level layouts."""
        _ = (canvas_size, normalize_box_format(box_format))
        if not normalized:
            raise ValueError("LT-Net outputs normalized boxes only")

        raw_box = model_outputs.refine_box
        if raw_box is None:
            raw_box = model_outputs.coarse_box
        if raw_box is None:
            raise ValueError("model_outputs must contain coarse_box or refine_box")

        batch_boxes: list[Float[torch.Tensor, "elements 4"]] = []
        batch_labels: list[Int[torch.Tensor, "elements"]] = []
        batch_masks: list[Bool[torch.Tensor, "elements"]] = []
        token_rows = (
            [None] * raw_box.size(0)
            if input_token is None
            else list(input_token.unbind(dim=0))
        )
        for row_box, row_obj_id, row_type, row_token in zip(
            raw_box, input_obj_id, token_type, token_rows, strict=True
        ):
            object_positions = row_type.eq(1) | row_type.eq(3)
            object_ids = row_obj_id[object_positions]
            boxes = row_box[object_positions].clamp(0.0, 1.0)
            labels = (
                torch.clamp(object_ids - 1, min=0).long()
                if row_token is None
                else row_token[object_positions].long()
            )
            valid = object_ids.gt(0)
            boxes = boxes[valid]
            labels = labels[valid]
            object_ids = object_ids[valid]
            reduced_boxes: list[Float[torch.Tensor, "elements 4"]] = []
            reduced_labels: list[Int[torch.Tensor, "elements"]] = []
            for object_id in object_ids.unique(sorted=True):
                positions = object_ids.eq(object_id).nonzero().flatten()
                if self.object_reduce == "mean":
                    reduced_boxes.append(boxes[positions].mean(dim=0))
                    reduced_labels.append(labels[positions[0]])
                else:
                    selected = (
                        positions[0] if self.object_reduce == "first" else positions[-1]
                    )
                    reduced_boxes.append(boxes[selected])
                    reduced_labels.append(labels[selected])
            if reduced_boxes:
                batch_boxes.append(torch.stack(reduced_boxes))
                batch_labels.append(torch.stack(reduced_labels).long())
            else:
                batch_boxes.append(boxes)
                batch_labels.append(labels)
            batch_masks.append(
                torch.ones(len(reduced_boxes), dtype=torch.bool, device=row_box.device)
            )
        max_items = max((item.size(0) for item in batch_boxes), default=0)
        padded_boxes = raw_box.new_zeros((raw_box.size(0), max_items, 4))
        padded_labels = input_obj_id.new_zeros((raw_box.size(0), max_items))
        padded_masks = torch.zeros(
            (raw_box.size(0), max_items), dtype=torch.bool, device=raw_box.device
        )
        for idx, (boxes, labels, mask) in enumerate(
            zip(batch_boxes, batch_labels, batch_masks, strict=True)
        ):
            length = boxes.size(0)
            padded_boxes[idx, :length] = boxes
            padded_labels[idx, :length] = labels
            padded_masks[idx, :length] = mask
        intermediates = None
        if return_intermediates:
            intermediates = {
                "coarse_box": model_outputs.coarse_box,
                "refine_box": model_outputs.refine_box,
                "vocab_logits": model_outputs.vocab_logits,
                "obj_id_logits": model_outputs.obj_id_logits,
                "token_type_logits": model_outputs.token_type_logits,
            }
        output = LayoutGenerationOutput(
            bbox=padded_boxes,
            labels=padded_labels,
            mask=padded_masks,
            id2label=dict(self.id2label),
            intermediates=intermediates,
        )
        if output_type == "dict":
            return dict(output.items())
        return output

    def save_pretrained(
        self,
        save_directory: str | PathLike[str],
        push_to_hub: bool = False,
        **kwargs: str | int | float | bool | None,
    ) -> tuple[str, ...]:
        """Save processor metadata and tokenizer files."""
        paths = super().save_pretrained(
            save_directory, push_to_hub=push_to_hub, **kwargs
        )
        metadata = {
            "processor_class": self.__class__.__name__,
            "dataset_name": self.dataset_name,
            "max_sequence_length": self.max_sequence_length,
            "id2label": self.id2label,
            "relation_id2label": self.relation_id2label,
            "object_reduce": self.object_reduce,
        }
        with (Path(save_directory) / "preprocessor_config.json").open("w") as f:
            json.dump(metadata, f, indent=2, sort_keys=True)
        return paths

__init__

__init__(
    tokenizer: LTNetRelationTokenizer,
    dataset_name: str = "coco",
    max_sequence_length: int = 128,
    id2label: Mapping[int, str]
    | Mapping[str, str]
    | None = None,
    relation_id2label: Mapping[int, str]
    | Mapping[str, str]
    | None = None,
    object_reduce: Literal[
        "first", "last", "mean"
    ] = "first",
) -> None

Initialize processor label maps and tokenizer component.

Source code in models/ltnet/src/ltnet/processing_ltnet.py
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
def __init__(
    self,
    tokenizer: LTNetRelationTokenizer,
    dataset_name: str = "coco",
    max_sequence_length: int = 128,
    id2label: Mapping[int, str] | Mapping[str, str] | None = None,
    relation_id2label: Mapping[int, str] | Mapping[str, str] | None = None,
    object_reduce: Literal["first", "last", "mean"] = "first",
) -> None:
    """Initialize processor label maps and tokenizer component."""
    self.tokenizer = tokenizer
    self.dataset_name = dataset_name
    self.max_sequence_length = max_sequence_length
    self.id2label = {
        int(key): str(value)
        for key, value in (id2label or DEFAULT_ID2LABEL).items()
    }
    self.relation_id2label = {
        int(key): str(value)
        for key, value in (relation_id2label or DEFAULT_RELATION_ID2LABEL).items()
    }
    self.label2id = {value.lower(): key for key, value in self.id2label.items()}
    self.relation_label2id = {
        value.lower(): key for key, value in self.relation_id2label.items()
    }
    if object_reduce not in {"first", "last", "mean"}:
        raise ValueError("object_reduce must be 'first', 'last', or 'mean'")

    self.object_reduce = object_reduce
    super().__init__(tokenizer=tokenizer)

from_config classmethod

from_config(
    *,
    dataset_name: str = "coco",
    max_sequence_length: int = 128,
    id2label: Mapping[int, str]
    | Mapping[str, str]
    | None = None,
    relation_id2label: Mapping[int, str]
    | Mapping[str, str]
    | None = None,
) -> "LTNetProcessor"

Construct a processor and tokenizer without external files.

Returns:

Type Description
'LTNetProcessor'

Processor with synthetic vocabulary derived from the label maps.

Examples:

>>> processor = LTNetProcessor.from_config()
>>> processor.tokenizer.cls_token_id
1
Source code in models/ltnet/src/ltnet/processing_ltnet.py
 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
@classmethod
def from_config(
    cls,
    *,
    dataset_name: str = "coco",
    max_sequence_length: int = 128,
    id2label: Mapping[int, str] | Mapping[str, str] | None = None,
    relation_id2label: Mapping[int, str] | Mapping[str, str] | None = None,
) -> "LTNetProcessor":
    """Construct a processor and tokenizer without external files.

    Returns:
        Processor with synthetic vocabulary derived from the label maps.

    Examples:
        >>> processor = LTNetProcessor.from_config()
        >>> processor.tokenizer.cls_token_id
        1
    """
    object_labels = {
        int(key): str(value)
        for key, value in (id2label or DEFAULT_ID2LABEL).items()
    }
    relation_labels = {
        int(key): str(value)
        for key, value in (relation_id2label or DEFAULT_RELATION_ID2LABEL).items()
    }
    tokens = ["__image__"]
    tokens.extend(value for _, value in sorted(object_labels.items()))
    tokens.extend(value for _, value in sorted(relation_labels.items()))
    tokenizer = LTNetRelationTokenizer(tokens=tokens)
    return cls(
        tokenizer=tokenizer,
        dataset_name=dataset_name,
        max_sequence_length=max_sequence_length,
        id2label=object_labels,
        relation_id2label=relation_labels,
    )

normalize_condition_type

normalize_condition_type(
    condition_type: ConditionType | str,
) -> ConditionType

Normalize and validate the LT-Net public condition type.

Source code in models/ltnet/src/ltnet/processing_ltnet.py
140
141
142
143
144
145
146
147
148
149
150
151
def normalize_condition_type(
    self, condition_type: ConditionType | str
) -> ConditionType:
    """Normalize and validate the LT-Net public condition type."""
    condition = normalize_condition_type(condition_type)
    if condition is not ConditionType.relation:
        raise ValueError(
            "LT-Net only supports condition_type='relation' "
            "and aliases 'scene_graph', 'graph', or 'gen_r'."
        )

    return condition

__call__

__call__(
    *,
    scene_graph: SceneGraphInput
    | SceneGraphMapping
    | None = None,
    objects: Sequence[LayoutObject] | None = None,
    relations: Sequence[LayoutRelation] | None = None,
    batch_size: int = 1,
    condition_type: ConditionType
    | str = ConditionType.relation,
    return_tensors: Literal["pt", "np"] = "pt",
    max_sequence_length: int | None = None,
) -> BatchEncoding

Build LT-Net model tensors from public scene-graph inputs.

Source code in models/ltnet/src/ltnet/processing_ltnet.py
286
287
288
289
290
291
292
293
294
295
296
297
298
299
300
301
302
303
304
305
306
307
308
309
310
311
312
313
314
315
316
317
318
319
320
321
322
323
324
325
326
327
def __call__(
    self,
    *,
    scene_graph: SceneGraphInput | SceneGraphMapping | None = None,
    objects: Sequence[LayoutObject] | None = None,
    relations: Sequence[LayoutRelation] | None = None,
    batch_size: int = 1,
    condition_type: ConditionType | str = ConditionType.relation,
    return_tensors: Literal["pt", "np"] = "pt",
    max_sequence_length: int | None = None,
) -> BatchEncoding:
    """Build LT-Net model tensors from public scene-graph inputs."""
    self.normalize_condition_type(condition_type)
    original_max_length = self.max_sequence_length
    try:
        if max_sequence_length is not None:
            self.max_sequence_length = max_sequence_length
        graph = self._normalize_scene_graph(
            scene_graph,
            objects=objects,
            relations=relations,
        )
        rows = [self._serialize_graph(graph) for _ in range(batch_size)]
    finally:
        self.max_sequence_length = original_max_length
    data = {
        "input_token": torch.tensor([row[0] for row in rows], dtype=torch.long),
        "input_obj_id": torch.tensor([row[1] for row in rows], dtype=torch.long),
        "segment_label": torch.tensor([row[2] for row in rows], dtype=torch.long),
        "token_type": torch.tensor([row[3] for row in rows], dtype=torch.long),
        "src_mask": torch.tensor(
            [row[4] for row in rows], dtype=torch.bool
        ).unsqueeze(1),
        "global_mask": torch.tensor([row[0] for row in rows], dtype=torch.long).ge(
            2
        ),
    }
    if return_tensors == "pt":
        return BatchEncoding(data)
    if return_tensors == "np":
        return BatchEncoding({key: value.numpy() for key, value in data.items()})
    raise ValueError("return_tensors must be 'pt' or 'np'")

post_process_layout_generation

post_process_layout_generation(
    model_outputs: LTNetModelOutput,
    *,
    input_token: Int[Tensor, "batch sequence"]
    | None = None,
    input_obj_id: Int[Tensor, "batch sequence"],
    token_type: Int[Tensor, "batch sequence"],
    box_format: BoxFormat | str = BoxFormat.xywh,
    normalized: bool = True,
    canvas_size: tuple[int, int] | None = None,
    output_type: OutputType = "dataclass",
    return_intermediates: bool = False,
) -> (
    LayoutGenerationOutput
    | dict[
        str,
        Float[torch.Tensor, "..."]
        | Int[torch.Tensor, "..."]
        | Bool[torch.Tensor, "..."]
        | dict[int, str]
        | dict[str, Float[torch.Tensor, "..."] | None],
    ]
)

Convert raw token-level boxes into public object-level layouts.

Source code in models/ltnet/src/ltnet/processing_ltnet.py
329
330
331
332
333
334
335
336
337
338
339
340
341
342
343
344
345
346
347
348
349
350
351
352
353
354
355
356
357
358
359
360
361
362
363
364
365
366
367
368
369
370
371
372
373
374
375
376
377
378
379
380
381
382
383
384
385
386
387
388
389
390
391
392
393
394
395
396
397
398
399
400
401
402
403
404
405
406
407
408
409
410
411
412
413
414
415
416
417
418
419
420
421
422
423
424
425
426
427
428
429
430
431
432
433
434
435
436
437
438
439
def post_process_layout_generation(
    self,
    model_outputs: LTNetModelOutput,
    *,
    input_token: Int[torch.Tensor, "batch sequence"] | None = None,
    input_obj_id: Int[torch.Tensor, "batch sequence"],
    token_type: Int[torch.Tensor, "batch sequence"],
    box_format: BoxFormat | str = BoxFormat.xywh,
    normalized: bool = True,
    canvas_size: tuple[int, int] | None = None,
    output_type: OutputType = "dataclass",
    return_intermediates: bool = False,
) -> (
    LayoutGenerationOutput
    | dict[
        str,
        Float[torch.Tensor, "..."]
        | Int[torch.Tensor, "..."]
        | Bool[torch.Tensor, "..."]
        | dict[int, str]
        | dict[str, Float[torch.Tensor, "..."] | None],
    ]
):
    """Convert raw token-level boxes into public object-level layouts."""
    _ = (canvas_size, normalize_box_format(box_format))
    if not normalized:
        raise ValueError("LT-Net outputs normalized boxes only")

    raw_box = model_outputs.refine_box
    if raw_box is None:
        raw_box = model_outputs.coarse_box
    if raw_box is None:
        raise ValueError("model_outputs must contain coarse_box or refine_box")

    batch_boxes: list[Float[torch.Tensor, "elements 4"]] = []
    batch_labels: list[Int[torch.Tensor, "elements"]] = []
    batch_masks: list[Bool[torch.Tensor, "elements"]] = []
    token_rows = (
        [None] * raw_box.size(0)
        if input_token is None
        else list(input_token.unbind(dim=0))
    )
    for row_box, row_obj_id, row_type, row_token in zip(
        raw_box, input_obj_id, token_type, token_rows, strict=True
    ):
        object_positions = row_type.eq(1) | row_type.eq(3)
        object_ids = row_obj_id[object_positions]
        boxes = row_box[object_positions].clamp(0.0, 1.0)
        labels = (
            torch.clamp(object_ids - 1, min=0).long()
            if row_token is None
            else row_token[object_positions].long()
        )
        valid = object_ids.gt(0)
        boxes = boxes[valid]
        labels = labels[valid]
        object_ids = object_ids[valid]
        reduced_boxes: list[Float[torch.Tensor, "elements 4"]] = []
        reduced_labels: list[Int[torch.Tensor, "elements"]] = []
        for object_id in object_ids.unique(sorted=True):
            positions = object_ids.eq(object_id).nonzero().flatten()
            if self.object_reduce == "mean":
                reduced_boxes.append(boxes[positions].mean(dim=0))
                reduced_labels.append(labels[positions[0]])
            else:
                selected = (
                    positions[0] if self.object_reduce == "first" else positions[-1]
                )
                reduced_boxes.append(boxes[selected])
                reduced_labels.append(labels[selected])
        if reduced_boxes:
            batch_boxes.append(torch.stack(reduced_boxes))
            batch_labels.append(torch.stack(reduced_labels).long())
        else:
            batch_boxes.append(boxes)
            batch_labels.append(labels)
        batch_masks.append(
            torch.ones(len(reduced_boxes), dtype=torch.bool, device=row_box.device)
        )
    max_items = max((item.size(0) for item in batch_boxes), default=0)
    padded_boxes = raw_box.new_zeros((raw_box.size(0), max_items, 4))
    padded_labels = input_obj_id.new_zeros((raw_box.size(0), max_items))
    padded_masks = torch.zeros(
        (raw_box.size(0), max_items), dtype=torch.bool, device=raw_box.device
    )
    for idx, (boxes, labels, mask) in enumerate(
        zip(batch_boxes, batch_labels, batch_masks, strict=True)
    ):
        length = boxes.size(0)
        padded_boxes[idx, :length] = boxes
        padded_labels[idx, :length] = labels
        padded_masks[idx, :length] = mask
    intermediates = None
    if return_intermediates:
        intermediates = {
            "coarse_box": model_outputs.coarse_box,
            "refine_box": model_outputs.refine_box,
            "vocab_logits": model_outputs.vocab_logits,
            "obj_id_logits": model_outputs.obj_id_logits,
            "token_type_logits": model_outputs.token_type_logits,
        }
    output = LayoutGenerationOutput(
        bbox=padded_boxes,
        labels=padded_labels,
        mask=padded_masks,
        id2label=dict(self.id2label),
        intermediates=intermediates,
    )
    if output_type == "dict":
        return dict(output.items())
    return output

save_pretrained

save_pretrained(
    save_directory: str | PathLike[str],
    push_to_hub: bool = False,
    **kwargs: str | int | float | bool | None,
) -> tuple[str, ...]

Save processor metadata and tokenizer files.

Source code in models/ltnet/src/ltnet/processing_ltnet.py
441
442
443
444
445
446
447
448
449
450
451
452
453
454
455
456
457
458
459
460
461
def save_pretrained(
    self,
    save_directory: str | PathLike[str],
    push_to_hub: bool = False,
    **kwargs: str | int | float | bool | None,
) -> tuple[str, ...]:
    """Save processor metadata and tokenizer files."""
    paths = super().save_pretrained(
        save_directory, push_to_hub=push_to_hub, **kwargs
    )
    metadata = {
        "processor_class": self.__class__.__name__,
        "dataset_name": self.dataset_name,
        "max_sequence_length": self.max_sequence_length,
        "id2label": self.id2label,
        "relation_id2label": self.relation_id2label,
        "object_reduce": self.object_reduce,
    }
    with (Path(save_directory) / "preprocessor_config.json").open("w") as f:
        json.dump(metadata, f, indent=2, sort_keys=True)
    return paths

relation_schema

Public scene-graph dataclasses for LT-Net processors.

LayoutObject dataclass

One scene-graph object node.

Parameters:

Name Type Description Default
id int | str

Stable object id within a scene graph.

required
label int | str

Dataset-local object label id or label string.

required
bbox tuple[float, float, float, float] | None

Optional normalized center xywh constraint.

None

Examples:

>>> LayoutObject(id="person-1", label="person").id
'person-1'
Source code in models/ltnet/src/ltnet/relation_schema.py
 8
 9
10
11
12
13
14
15
16
17
18
19
20
21
22
23
24
@dataclass(frozen=True)
class LayoutObject:
    """One scene-graph object node.

    Args:
        id: Stable object id within a scene graph.
        label: Dataset-local object label id or label string.
        bbox: Optional normalized center ``xywh`` constraint.

    Examples:
        >>> LayoutObject(id="person-1", label="person").id
        'person-1'
    """

    id: int | str
    label: int | str
    bbox: tuple[float, float, float, float] | None = None

LayoutRelation dataclass

One directed scene-graph edge.

Parameters:

Name Type Description Default
subject int | str

Source object id.

required
predicate int | str

Relation id or label.

required
object int | str

Target object id.

required
bbox_delta tuple[float, float, float, float] | None

Optional relation geometry.

None
score float | None

Optional edge confidence.

None

Examples:

>>> LayoutRelation("a", "left of", "b").predicate
'left of'
Source code in models/ltnet/src/ltnet/relation_schema.py
27
28
29
30
31
32
33
34
35
36
37
38
39
40
41
42
43
44
45
46
47
@dataclass(frozen=True)
class LayoutRelation:
    """One directed scene-graph edge.

    Args:
        subject: Source object id.
        predicate: Relation id or label.
        object: Target object id.
        bbox_delta: Optional relation geometry.
        score: Optional edge confidence.

    Examples:
        >>> LayoutRelation("a", "left of", "b").predicate
        'left of'
    """

    subject: int | str
    predicate: int | str
    object: int | str
    bbox_delta: tuple[float, float, float, float] | None = None
    score: float | None = None

SceneGraphInput dataclass

Normalized scene graph payload.

Parameters:

Name Type Description Default
objects tuple[LayoutObject, ...]

Scene-graph object nodes.

required
relations tuple[LayoutRelation, ...]

Scene-graph relation edges.

required
id2label dict[int, str] | None

Optional public object label mapping.

None
relation_id2label dict[int, str] | None

Optional relation label mapping.

None

Examples:

>>> graph = SceneGraphInput(
...     objects=(LayoutObject("a", "person"), LayoutObject("b", "table")),
...     relations=(LayoutRelation("a", "left of", "b"),),
... )
>>> len(graph.relations)
1
Source code in models/ltnet/src/ltnet/relation_schema.py
50
51
52
53
54
55
56
57
58
59
60
61
62
63
64
65
66
67
68
69
70
71
72
@dataclass(frozen=True)
class SceneGraphInput:
    """Normalized scene graph payload.

    Args:
        objects: Scene-graph object nodes.
        relations: Scene-graph relation edges.
        id2label: Optional public object label mapping.
        relation_id2label: Optional relation label mapping.

    Examples:
        >>> graph = SceneGraphInput(
        ...     objects=(LayoutObject("a", "person"), LayoutObject("b", "table")),
        ...     relations=(LayoutRelation("a", "left of", "b"),),
        ... )
        >>> len(graph.relations)
        1
    """

    objects: tuple[LayoutObject, ...]
    relations: tuple[LayoutRelation, ...]
    id2label: dict[int, str] | None = None
    relation_id2label: dict[int, str] | None = None

tokenization_ltnet

Tokenizer for LT-Net scene-graph token vocabularies.

LTNetRelationTokenizer

Bases: WhitespaceTokenizerMixin, PreTrainedTokenizer

Discrete scene-graph tokenizer saved as a standard HF tokenizer.

Parameters:

Name Type Description Default
vocab_file str | None

Optional path to object_pred_id2name.json.

None
tokens list[str] | None

Optional token list used when no vocab file is supplied.

None
object_token_ids list[int] | None

Optional ids that represent object classes.

None
relation_token_ids list[int] | None

Optional ids that represent predicates.

None
pad_token str

Padding token.

'[PAD]'
cls_token str

Sequence-start token.

'[CLS]'
sep_token str

Triple separator token.

'[SEP]'
mask_token str

Mask token.

'[MASK]'
unk_token str

Unknown token.

'[MASK]'
model_max_length int

Maximum tokenizer length metadata.

DEFAULT_MODEL_MAX_LENGTH
kwargs str | int | float | bool | None

Additional tokenizer compatibility fields.

{}

Examples:

>>> tokenizer = LTNetRelationTokenizer(tokens=["__image__", "person"])
>>> tokenizer.convert_tokens_to_ids("[CLS]")
1
Source code in models/ltnet/src/ltnet/tokenization_ltnet.py
 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
class LTNetRelationTokenizer(WhitespaceTokenizerMixin, PreTrainedTokenizer):
    """Discrete scene-graph tokenizer saved as a standard HF tokenizer.

    Args:
        vocab_file: Optional path to ``object_pred_id2name.json``.
        tokens: Optional token list used when no vocab file is supplied.
        object_token_ids: Optional ids that represent object classes.
        relation_token_ids: Optional ids that represent predicates.
        pad_token: Padding token.
        cls_token: Sequence-start token.
        sep_token: Triple separator token.
        mask_token: Mask token.
        unk_token: Unknown token.
        model_max_length: Maximum tokenizer length metadata.
        kwargs: Additional tokenizer compatibility fields.

    Examples:
        >>> tokenizer = LTNetRelationTokenizer(tokens=["__image__", "person"])
        >>> tokenizer.convert_tokens_to_ids("[CLS]")
        1
    """

    vocab_files_names = {"vocab_file": "object_pred_id2name.json"}
    model_input_names = ["input_token", "input_obj_id", "segment_label", "token_type"]

    def __init__(
        self,
        vocab_file: str | None = None,
        tokens: list[str] | None = None,
        object_token_ids: list[int] | None = None,
        relation_token_ids: list[int] | None = None,
        pad_token: str = "[PAD]",
        cls_token: str = "[CLS]",
        sep_token: str = "[SEP]",
        mask_token: str = "[MASK]",
        unk_token: str = "[MASK]",
        model_max_length: int = DEFAULT_MODEL_MAX_LENGTH,
        padding_side: str = "right",
        truncation_side: str = "right",
        clean_up_tokenization_spaces: bool = False,
        added_tokens_decoder: dict[int | str, str | AddedToken] | None = None,
        name_or_path: str = "",
        **kwargs: str | int | float | bool | None,
    ) -> None:
        """Initialize vocabulary and object/relation id metadata."""
        _ = kwargs
        token2id, id2token = build_token_maps(
            vocab_file=vocab_file,
            tokens=tokens,
            base_tokens=SPECIAL_TOKENS,
            numeric_id_vocab=True,
        )
        self._token2id = token2id
        self._id2token = id2token
        self.object_token_ids = [int(item) for item in object_token_ids or []]
        self.relation_token_ids = [int(item) for item in relation_token_ids or []]
        tokenizer_kwargs: dict[str, object] = {
            "pad_token": pad_token,
            "cls_token": cls_token,
            "sep_token": sep_token,
            "mask_token": mask_token,
            "unk_token": unk_token,
            "model_max_length": model_max_length,
            "padding_side": padding_side,
            "truncation_side": truncation_side,
            "clean_up_tokenization_spaces": clean_up_tokenization_spaces,
            "name_or_path": name_or_path,
        }
        if added_tokens_decoder is not None:
            tokenizer_kwargs["added_tokens_decoder"] = added_tokens_decoder
        super().__init__(**tokenizer_kwargs)

    def encode_scene_graph_tokens(self, tokens: list[str]) -> list[int]:
        """Encode already-normalized scene-graph token strings.

        Args:
            tokens: Token strings in LT-Net order.

        Returns:
            Integer token ids.

        Raises:
            ValueError: If an unknown token appears.
        """
        ids: list[int] = []
        for token in tokens:
            token_id = self._convert_token_to_id(token)
            if token_id == self.unk_token_id and token != self.unk_token:
                raise ValueError(f"Unknown scene-graph token: {token}")

            ids.append(token_id)
        return ids

    def decode_scene_graph_tokens(self, input_token: list[int]) -> list[str]:
        """Decode integer scene-graph token ids into token strings."""
        return [self._convert_id_to_token(token_id) for token_id in input_token]

    def save_vocabulary(
        self, save_directory: str, filename_prefix: str | None = None
    ) -> tuple[str]:
        """Save ``object_pred_id2name.json`` as id-to-token metadata."""
        return save_json_vocabulary(
            save_directory=save_directory,
            filename="object_pred_id2name.json",
            data={str(key): value for key, value in sorted(self._id2token.items())},
            filename_prefix=filename_prefix,
        )

    def save_pretrained(
        self,
        save_directory: str | PathLike[str],
        legacy_format: bool | None = None,
        filename_prefix: str | None = None,
        push_to_hub: bool = False,
        **kwargs: str | int | float | bool | None,
    ) -> tuple[str, ...]:
        """Save tokenizer files plus LT-Net tokenizer metadata."""
        _ = kwargs
        paths = super().save_pretrained(
            str(save_directory),
            legacy_format=legacy_format,
            filename_prefix=filename_prefix,
            push_to_hub=push_to_hub,
        )
        metadata = {
            "object_token_ids": self.object_token_ids,
            "relation_token_ids": self.relation_token_ids,
        }
        with (Path(save_directory) / "ltnet_tokenizer_config.json").open("w") as f:
            json.dump(metadata, f, indent=2, sort_keys=True)
        return paths

    @classmethod
    def from_pretrained(
        cls,
        pretrained_model_name_or_path: str | PathLike[str],
        cache_dir: str | PathLike[str] | None = None,
        force_download: bool = False,
        local_files_only: bool = False,
        token: str | bool | None = None,
        revision: str = "main",
        object_token_ids: list[int] | None = None,
        relation_token_ids: list[int] | None = None,
        **kwargs: str | int | float | bool | None,
    ) -> "LTNetRelationTokenizer":
        """Load tokenizer and LT-Net metadata."""
        path = Path(pretrained_model_name_or_path)
        metadata_path = path / "ltnet_tokenizer_config.json"
        metadata: dict[str, object] = {}
        if metadata_path.exists():
            with metadata_path.open() as f:
                metadata = json.load(f)
        if object_token_ids is not None:
            metadata["object_token_ids"] = object_token_ids
        if relation_token_ids is not None:
            metadata["relation_token_ids"] = relation_token_ids
        metadata.update(kwargs)
        return cast(
            "LTNetRelationTokenizer",
            super().from_pretrained(
                str(pretrained_model_name_or_path),
                cache_dir=cache_dir,
                force_download=force_download,
                local_files_only=local_files_only,
                token=token,
                revision=revision,
                **metadata,
            ),
        )

__init__

__init__(
    vocab_file: str | None = None,
    tokens: list[str] | None = None,
    object_token_ids: list[int] | None = None,
    relation_token_ids: list[int] | None = None,
    pad_token: str = "[PAD]",
    cls_token: str = "[CLS]",
    sep_token: str = "[SEP]",
    mask_token: str = "[MASK]",
    unk_token: str = "[MASK]",
    model_max_length: int = DEFAULT_MODEL_MAX_LENGTH,
    padding_side: str = "right",
    truncation_side: str = "right",
    clean_up_tokenization_spaces: bool = False,
    added_tokens_decoder: dict[int | str, str | AddedToken]
    | None = None,
    name_or_path: str = "",
    **kwargs: str | int | float | bool | None,
) -> None

Initialize vocabulary and object/relation id metadata.

Source code in models/ltnet/src/ltnet/tokenization_ltnet.py
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
def __init__(
    self,
    vocab_file: str | None = None,
    tokens: list[str] | None = None,
    object_token_ids: list[int] | None = None,
    relation_token_ids: list[int] | None = None,
    pad_token: str = "[PAD]",
    cls_token: str = "[CLS]",
    sep_token: str = "[SEP]",
    mask_token: str = "[MASK]",
    unk_token: str = "[MASK]",
    model_max_length: int = DEFAULT_MODEL_MAX_LENGTH,
    padding_side: str = "right",
    truncation_side: str = "right",
    clean_up_tokenization_spaces: bool = False,
    added_tokens_decoder: dict[int | str, str | AddedToken] | None = None,
    name_or_path: str = "",
    **kwargs: str | int | float | bool | None,
) -> None:
    """Initialize vocabulary and object/relation id metadata."""
    _ = kwargs
    token2id, id2token = build_token_maps(
        vocab_file=vocab_file,
        tokens=tokens,
        base_tokens=SPECIAL_TOKENS,
        numeric_id_vocab=True,
    )
    self._token2id = token2id
    self._id2token = id2token
    self.object_token_ids = [int(item) for item in object_token_ids or []]
    self.relation_token_ids = [int(item) for item in relation_token_ids or []]
    tokenizer_kwargs: dict[str, object] = {
        "pad_token": pad_token,
        "cls_token": cls_token,
        "sep_token": sep_token,
        "mask_token": mask_token,
        "unk_token": unk_token,
        "model_max_length": model_max_length,
        "padding_side": padding_side,
        "truncation_side": truncation_side,
        "clean_up_tokenization_spaces": clean_up_tokenization_spaces,
        "name_or_path": name_or_path,
    }
    if added_tokens_decoder is not None:
        tokenizer_kwargs["added_tokens_decoder"] = added_tokens_decoder
    super().__init__(**tokenizer_kwargs)

encode_scene_graph_tokens

encode_scene_graph_tokens(tokens: list[str]) -> list[int]

Encode already-normalized scene-graph token strings.

Parameters:

Name Type Description Default
tokens list[str]

Token strings in LT-Net order.

required

Returns:

Type Description
list[int]

Integer token ids.

Raises:

Type Description
ValueError

If an unknown token appears.

Source code in models/ltnet/src/ltnet/tokenization_ltnet.py
 94
 95
 96
 97
 98
 99
100
101
102
103
104
105
106
107
108
109
110
111
112
113
def encode_scene_graph_tokens(self, tokens: list[str]) -> list[int]:
    """Encode already-normalized scene-graph token strings.

    Args:
        tokens: Token strings in LT-Net order.

    Returns:
        Integer token ids.

    Raises:
        ValueError: If an unknown token appears.
    """
    ids: list[int] = []
    for token in tokens:
        token_id = self._convert_token_to_id(token)
        if token_id == self.unk_token_id and token != self.unk_token:
            raise ValueError(f"Unknown scene-graph token: {token}")

        ids.append(token_id)
    return ids

decode_scene_graph_tokens

decode_scene_graph_tokens(
    input_token: list[int],
) -> list[str]

Decode integer scene-graph token ids into token strings.

Source code in models/ltnet/src/ltnet/tokenization_ltnet.py
115
116
117
def decode_scene_graph_tokens(self, input_token: list[int]) -> list[str]:
    """Decode integer scene-graph token ids into token strings."""
    return [self._convert_id_to_token(token_id) for token_id in input_token]

save_vocabulary

save_vocabulary(
    save_directory: str, filename_prefix: str | None = None
) -> tuple[str]

Save object_pred_id2name.json as id-to-token metadata.

Source code in models/ltnet/src/ltnet/tokenization_ltnet.py
119
120
121
122
123
124
125
126
127
128
def save_vocabulary(
    self, save_directory: str, filename_prefix: str | None = None
) -> tuple[str]:
    """Save ``object_pred_id2name.json`` as id-to-token metadata."""
    return save_json_vocabulary(
        save_directory=save_directory,
        filename="object_pred_id2name.json",
        data={str(key): value for key, value in sorted(self._id2token.items())},
        filename_prefix=filename_prefix,
    )

save_pretrained

save_pretrained(
    save_directory: str | PathLike[str],
    legacy_format: bool | None = None,
    filename_prefix: str | None = None,
    push_to_hub: bool = False,
    **kwargs: str | int | float | bool | None,
) -> tuple[str, ...]

Save tokenizer files plus LT-Net tokenizer metadata.

Source code in models/ltnet/src/ltnet/tokenization_ltnet.py
130
131
132
133
134
135
136
137
138
139
140
141
142
143
144
145
146
147
148
149
150
151
152
def save_pretrained(
    self,
    save_directory: str | PathLike[str],
    legacy_format: bool | None = None,
    filename_prefix: str | None = None,
    push_to_hub: bool = False,
    **kwargs: str | int | float | bool | None,
) -> tuple[str, ...]:
    """Save tokenizer files plus LT-Net tokenizer metadata."""
    _ = kwargs
    paths = super().save_pretrained(
        str(save_directory),
        legacy_format=legacy_format,
        filename_prefix=filename_prefix,
        push_to_hub=push_to_hub,
    )
    metadata = {
        "object_token_ids": self.object_token_ids,
        "relation_token_ids": self.relation_token_ids,
    }
    with (Path(save_directory) / "ltnet_tokenizer_config.json").open("w") as f:
        json.dump(metadata, f, indent=2, sort_keys=True)
    return paths

from_pretrained classmethod

from_pretrained(
    pretrained_model_name_or_path: str | PathLike[str],
    cache_dir: str | PathLike[str] | None = None,
    force_download: bool = False,
    local_files_only: bool = False,
    token: str | bool | None = None,
    revision: str = "main",
    object_token_ids: list[int] | None = None,
    relation_token_ids: list[int] | None = None,
    **kwargs: str | int | float | bool | None,
) -> "LTNetRelationTokenizer"

Load tokenizer and LT-Net metadata.

Source code in models/ltnet/src/ltnet/tokenization_ltnet.py
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
@classmethod
def from_pretrained(
    cls,
    pretrained_model_name_or_path: str | PathLike[str],
    cache_dir: str | PathLike[str] | None = None,
    force_download: bool = False,
    local_files_only: bool = False,
    token: str | bool | None = None,
    revision: str = "main",
    object_token_ids: list[int] | None = None,
    relation_token_ids: list[int] | None = None,
    **kwargs: str | int | float | bool | None,
) -> "LTNetRelationTokenizer":
    """Load tokenizer and LT-Net metadata."""
    path = Path(pretrained_model_name_or_path)
    metadata_path = path / "ltnet_tokenizer_config.json"
    metadata: dict[str, object] = {}
    if metadata_path.exists():
        with metadata_path.open() as f:
            metadata = json.load(f)
    if object_token_ids is not None:
        metadata["object_token_ids"] = object_token_ids
    if relation_token_ids is not None:
        metadata["relation_token_ids"] = relation_token_ids
    metadata.update(kwargs)
    return cast(
        "LTNetRelationTokenizer",
        super().from_pretrained(
            str(pretrained_model_name_or_path),
            cache_dir=cache_dir,
            force_download=force_download,
            local_files_only=local_files_only,
            token=token,
            revision=revision,
            **metadata,
        ),
    )

vendor_state_dict

Utilities for loading LT-Net checkpoint state dictionaries.

load_original_state_dict

load_original_state_dict(
    checkpoint_path: str | Path,
) -> dict[str, Shaped[torch.Tensor, "..."]]

Load a vendor checkpoint and return model weights only.

Parameters:

Name Type Description Default
checkpoint_path str | Path

Original .pth checkpoint path.

required

Returns:

Type Description
dict[str, Shaped[Tensor, '...']]

Raw or checkpoint["state_dict"] tensor mapping with any

dict[str, Shaped[Tensor, '...']]

module. DataParallel prefix stripped.

Examples:

>>> import tempfile
>>> import torch
>>> with tempfile.NamedTemporaryFile(suffix=".pth") as handle:
...     torch.save({"state_dict": {"module.weight": torch.ones(1)}}, handle.name)
...     sorted(load_original_state_dict(handle.name))
['weight']
Source code in models/ltnet/src/ltnet/vendor_state_dict.py
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
def load_original_state_dict(
    checkpoint_path: str | Path,
) -> dict[str, Shaped[torch.Tensor, "..."]]:
    """Load a vendor checkpoint and return model weights only.

    Args:
        checkpoint_path: Original ``.pth`` checkpoint path.

    Returns:
        Raw or ``checkpoint["state_dict"]`` tensor mapping with any
        ``module.`` DataParallel prefix stripped.

    Examples:
        >>> import tempfile
        >>> import torch
        >>> with tempfile.NamedTemporaryFile(suffix=".pth") as handle:
        ...     torch.save({"state_dict": {"module.weight": torch.ones(1)}}, handle.name)
        ...     sorted(load_original_state_dict(handle.name))
        ['weight']
    """
    checkpoint = torch.load(checkpoint_path, map_location="cpu")
    state = (
        checkpoint.get("state_dict", checkpoint)
        if isinstance(checkpoint, Mapping)
        else checkpoint
    )
    tensor_state = cast(dict[str, Shaped[torch.Tensor, "..."]], state)
    if any(key.startswith("module.") for key in tensor_state):
        return {
            key.removeprefix("module."): value for key, value in tensor_state.items()
        }
    return dict(tensor_state)

load_strict_mapped_state_dict

load_strict_mapped_state_dict(
    model: Module,
    state_dict: Mapping[str, Shaped[Tensor, "..."]],
) -> None

Load a mapped state dict and fail on any key mismatch.

Parameters:

Name Type Description Default
model Module

Target converted model.

required
state_dict Mapping[str, Shaped[Tensor, '...']]

Converted tensor mapping.

required

Raises:

Type Description
RuntimeError

If keys are missing or unexpected.

Source code in models/ltnet/src/ltnet/vendor_state_dict.py
47
48
49
50
51
52
53
54
55
56
57
58
59
60
61
62
63
64
65
66
def load_strict_mapped_state_dict(
    model: torch.nn.Module,
    state_dict: Mapping[str, Shaped[torch.Tensor, "..."]],
) -> None:
    """Load a mapped state dict and fail on any key mismatch.

    Args:
        model: Target converted model.
        state_dict: Converted tensor mapping.

    Raises:
        RuntimeError: If keys are missing or unexpected.
    """
    incompatible = model.load_state_dict(dict(state_dict), strict=False)
    if incompatible.missing_keys or incompatible.unexpected_keys:
        raise RuntimeError(
            "State dict mismatch: "
            f"missing={incompatible.missing_keys}, "
            f"unexpected={incompatible.unexpected_keys}"
        )