Skip to content

Parse then place

Parse-Then-Place Transformers-style conversion package.

ParseThenPlaceConfig

Bases: PretrainedConfig

Stores Parse-Then-Place dataset and generation defaults.

Source code in models/parse-then-place/src/parse_then_place/configuration_parse_then_place.py
18
19
20
21
22
23
24
25
26
27
28
29
30
31
32
33
34
35
36
37
38
39
40
41
42
43
44
45
46
47
48
49
50
51
52
53
54
55
56
57
58
59
60
61
62
63
64
65
66
67
68
69
70
71
72
73
74
75
76
77
78
79
80
81
82
83
84
85
86
87
88
89
90
91
92
93
94
95
96
97
98
class ParseThenPlaceConfig(PretrainedConfig):
    """Stores Parse-Then-Place dataset and generation defaults."""

    model_type = "parse-then-place"

    def __init__(
        self,
        dataset_name: str = "rico",
        stage2_mode: Stage2Mode | str = Stage2Mode.finetune,
        parser_model_name: str = "google/t5-v1_1-base",
        parser_generation_max_length: int = 600,
        placement_generation_max_length: int = 500,
        temperature: float = 0.7,
        num_return_sequences: int = 5,
        canvas_size: tuple[int, int] | list[int] | None = None,
        id2label: dict[int | str, str] | None = None,
        parser_subfolder: str = "semantic_parser",
        placement_subfolder: str = "placement",
        pad_token_id: int = 0,
        eos_token_id: int = 1,
        decoder_start_token_id: int = 0,
        is_encoder_decoder: bool = True,
        transformers_version: str | None = None,
        architectures: list[str] | None = None,
        output_hidden_states: bool | None = False,
        return_dict: bool | None = True,
        dtype: str | None = None,
        torch_dtype: str | None = None,
        chunk_size_feed_forward: int = 0,
        problem_type: Literal[
            "regression", "single_label_classification", "multi_label_classification"
        ]
        | None = None,
        name_or_path: str = "",
        _commit_hash: str | None = None,
        attn_implementation: str | None = None,
        **kwargs: str | int | float | bool | None,
    ) -> None:
        """Initialize the composite checkpoint configuration."""
        dataset = normalize_dataset_name(dataset_name)
        mode = normalize_stage2_mode(stage2_mode)

        self.dataset_name = str(dataset)
        self.stage2_mode = str(mode)
        self.parser_model_name = parser_model_name
        self.parser_generation_max_length = parser_generation_max_length
        self.placement_generation_max_length = placement_generation_max_length
        self.temperature = temperature
        self.num_return_sequences = num_return_sequences
        self.canvas_size = tuple(canvas_size or canvas_size_for_dataset(dataset))
        self.parser_subfolder = parser_subfolder
        self.placement_subfolder = placement_subfolder

        label_map = id2label or id2label_for_dataset(dataset)
        normalized_id2label = {int(key): str(value) for key, value in label_map.items()}
        label2id = {value: key for key, value in normalized_id2label.items()}
        _ = kwargs.pop("id2label", None)
        _ = kwargs.pop("label2id", None)

        super().__init__(
            transformers_version=transformers_version,
            architectures=architectures,
            output_hidden_states=output_hidden_states,
            return_dict=return_dict,
            dtype=dtype or torch_dtype,
            chunk_size_feed_forward=chunk_size_feed_forward,
            is_encoder_decoder=is_encoder_decoder,
            id2label=normalized_id2label,
            label2id=label2id,
            problem_type=problem_type,
        )
        # Transformers v5 keeps model-specific token fields on the subclass;
        # only common configuration fields belong in the base call.
        self.pad_token_id = pad_token_id
        self.eos_token_id = eos_token_id
        self.decoder_start_token_id = decoder_start_token_id
        self.name_or_path = name_or_path
        self._commit_hash = _commit_hash
        self._attn_implementation = attn_implementation
        for key, value in kwargs.items():
            setattr(self, key, value)

__init__

__init__(
    dataset_name: str = "rico",
    stage2_mode: Stage2Mode | str = Stage2Mode.finetune,
    parser_model_name: str = "google/t5-v1_1-base",
    parser_generation_max_length: int = 600,
    placement_generation_max_length: int = 500,
    temperature: float = 0.7,
    num_return_sequences: int = 5,
    canvas_size: tuple[int, int] | list[int] | None = None,
    id2label: dict[int | str, str] | None = None,
    parser_subfolder: str = "semantic_parser",
    placement_subfolder: str = "placement",
    pad_token_id: int = 0,
    eos_token_id: int = 1,
    decoder_start_token_id: int = 0,
    is_encoder_decoder: bool = True,
    transformers_version: str | None = None,
    architectures: list[str] | None = None,
    output_hidden_states: bool | None = False,
    return_dict: bool | None = True,
    dtype: str | None = None,
    torch_dtype: str | None = None,
    chunk_size_feed_forward: int = 0,
    problem_type: Literal[
        "regression",
        "single_label_classification",
        "multi_label_classification",
    ]
    | None = None,
    name_or_path: str = "",
    _commit_hash: str | None = None,
    attn_implementation: str | None = None,
    **kwargs: str | int | float | bool | None,
) -> None

Initialize the composite checkpoint configuration.

Source code in models/parse-then-place/src/parse_then_place/configuration_parse_then_place.py
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
def __init__(
    self,
    dataset_name: str = "rico",
    stage2_mode: Stage2Mode | str = Stage2Mode.finetune,
    parser_model_name: str = "google/t5-v1_1-base",
    parser_generation_max_length: int = 600,
    placement_generation_max_length: int = 500,
    temperature: float = 0.7,
    num_return_sequences: int = 5,
    canvas_size: tuple[int, int] | list[int] | None = None,
    id2label: dict[int | str, str] | None = None,
    parser_subfolder: str = "semantic_parser",
    placement_subfolder: str = "placement",
    pad_token_id: int = 0,
    eos_token_id: int = 1,
    decoder_start_token_id: int = 0,
    is_encoder_decoder: bool = True,
    transformers_version: str | None = None,
    architectures: list[str] | None = None,
    output_hidden_states: bool | None = False,
    return_dict: bool | None = True,
    dtype: str | None = None,
    torch_dtype: str | None = None,
    chunk_size_feed_forward: int = 0,
    problem_type: Literal[
        "regression", "single_label_classification", "multi_label_classification"
    ]
    | None = None,
    name_or_path: str = "",
    _commit_hash: str | None = None,
    attn_implementation: str | None = None,
    **kwargs: str | int | float | bool | None,
) -> None:
    """Initialize the composite checkpoint configuration."""
    dataset = normalize_dataset_name(dataset_name)
    mode = normalize_stage2_mode(stage2_mode)

    self.dataset_name = str(dataset)
    self.stage2_mode = str(mode)
    self.parser_model_name = parser_model_name
    self.parser_generation_max_length = parser_generation_max_length
    self.placement_generation_max_length = placement_generation_max_length
    self.temperature = temperature
    self.num_return_sequences = num_return_sequences
    self.canvas_size = tuple(canvas_size or canvas_size_for_dataset(dataset))
    self.parser_subfolder = parser_subfolder
    self.placement_subfolder = placement_subfolder

    label_map = id2label or id2label_for_dataset(dataset)
    normalized_id2label = {int(key): str(value) for key, value in label_map.items()}
    label2id = {value: key for key, value in normalized_id2label.items()}
    _ = kwargs.pop("id2label", None)
    _ = kwargs.pop("label2id", None)

    super().__init__(
        transformers_version=transformers_version,
        architectures=architectures,
        output_hidden_states=output_hidden_states,
        return_dict=return_dict,
        dtype=dtype or torch_dtype,
        chunk_size_feed_forward=chunk_size_feed_forward,
        is_encoder_decoder=is_encoder_decoder,
        id2label=normalized_id2label,
        label2id=label2id,
        problem_type=problem_type,
    )
    # Transformers v5 keeps model-specific token fields on the subclass;
    # only common configuration fields belong in the base call.
    self.pad_token_id = pad_token_id
    self.eos_token_id = eos_token_id
    self.decoder_start_token_id = decoder_start_token_id
    self.name_or_path = name_or_path
    self._commit_hash = _commit_hash
    self._attn_implementation = attn_implementation
    for key, value in kwargs.items():
        setattr(self, key, value)

ParseThenPlaceDatasetName

Bases: StrEnum

Datasets supported by the original Parse-Then-Place release.

Source code in models/parse-then-place/src/parse_then_place/labels.py
11
12
13
14
15
class ParseThenPlaceDatasetName(StrEnum):
    """Datasets supported by the original Parse-Then-Place release."""

    rico = auto()
    web = auto()

Stage2Mode

Bases: StrEnum

Released stage-2 checkpoint modes.

Source code in models/parse-then-place/src/parse_then_place/labels.py
18
19
20
21
22
class Stage2Mode(StrEnum):
    """Released stage-2 checkpoint modes."""

    pretrain = auto()
    finetune = auto()

ParseThenPlacePipeline

Bases: LayoutGenerationPipeline

Compose standard seq2seq parser and placement models.

Source code in models/parse-then-place/src/parse_then_place/pipeline_parse_then_place.py
 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
class ParseThenPlacePipeline(LayoutGenerationPipeline):
    """Compose standard seq2seq parser and placement models."""

    config_class: ClassVar[type[PretrainedConfig]] = ParseThenPlaceConfig
    component_specs: ClassVar[dict[str, PipelineComponentSpec]] = {
        "parser": PipelineComponentSpec(
            attribute_name="parser",
            loader=_load_seq2seq_component,
            config_subfolder_attribute="parser_subfolder",
            required=False,
        ),
        "placement": PipelineComponentSpec(
            attribute_name="placement",
            loader=_load_seq2seq_component,
            config_subfolder_attribute="placement_subfolder",
        ),
        "processor": PipelineComponentSpec(
            attribute_name="processor",
            loader=_load_processor_component,
            marker_file="processor_config.json",
            save_with_is_main_process=False,
        ),
    }

    config: ParseThenPlaceConfig
    parser: PreTrainedModel | None
    placement: PreTrainedModel | None
    processor: ParseThenPlaceProcessor

    def __init__(
        self,
        config: ParseThenPlaceConfig,
        processor: ParseThenPlaceProcessor,
        *,
        parser: PreTrainedModel | None = None,
        placement: PreTrainedModel | None = None,
    ) -> None:
        """Initialize the composite pipeline."""
        super().__init__(config)
        self.config = config
        self.processor = processor
        self.parser = parser
        self.placement = placement

    @classmethod
    def from_pretrained(
        cls,
        pretrained_model_name_or_path: str | Path,
        *,
        parser: PreTrainedModel | None = None,
        placement: PreTrainedModel | None = None,
        processor: ParseThenPlaceProcessor | None = None,
        local_files_only: bool = False,
        config: ParseThenPlaceConfig | PretrainedConfig | None = None,
    ) -> ParseThenPlacePipeline:  # ty: ignore[invalid-method-override]
        """Load a composite pipeline from a root directory."""
        components: dict[str, PreTrainedModel | ParseThenPlaceProcessor] = {}
        if parser is not None:
            components["parser"] = parser
        if placement is not None:
            components["placement"] = placement
        if processor is None:
            loaded = super().from_pretrained(
                pretrained_model_name_or_path,
                local_files_only=local_files_only,
                config=config,
                components=components,
            )
            return cast(ParseThenPlacePipeline, loaded)
        components["processor"] = processor
        loaded = super().from_pretrained(
            pretrained_model_name_or_path,
            local_files_only=local_files_only,
            config=config,
            components=components,
        )
        return cast(ParseThenPlacePipeline, loaded)

    @classmethod
    def _from_pretrained_components(
        cls,
        *,
        config: PretrainedConfig,
        components: Mapping[str, PreTrainedModel | ParseThenPlaceProcessor | None],
    ) -> ParseThenPlacePipeline:
        """Build a pipeline from loaded config and components."""
        return cls(
            config=cast(ParseThenPlaceConfig, config),
            processor=cast(ParseThenPlaceProcessor, components["processor"]),
            parser=cast(PreTrainedModel | None, components.get("parser")),
            placement=cast(PreTrainedModel | None, components["placement"]),
        )

    @torch.no_grad()
    def parse(
        self,
        input_ids: Int[torch.Tensor, "batch tokens"],
        attention_mask: Bool[torch.Tensor, "batch tokens"] | None = None,
        *,
        generation_max_length: int | None = None,
        **generate_kwargs: str | int | float | bool | torch.Generator | None,
    ) -> Int[torch.Tensor, "batch tokens"]:
        """Generate logical-form token ids with the parser stage."""
        if self.parser is None:
            raise ValueError("Parser stage is not loaded")

        generated = cast(_GenerationModel, self.parser).generate(
            input_ids=input_ids,
            attention_mask=attention_mask,
            max_length=generation_max_length
            or self.config.parser_generation_max_length,
            **generate_kwargs,
        )
        return generated

    @torch.no_grad()
    def place(
        self,
        input_ids: Int[torch.Tensor, "batch tokens"],
        attention_mask: Bool[torch.Tensor, "batch tokens"] | None = None,
        *,
        generation_max_length: int | None = None,
        num_return_sequences: int | None = None,
        temperature: float | None = None,
        do_sample: bool = True,
        generator: torch.Generator | None = None,
        **generate_kwargs: str | float | bool | torch.Generator | None,
    ) -> Int[torch.Tensor, "batch tokens"]:
        """Generate layout token ids with the placement stage."""
        if self.placement is None:
            raise ValueError("Placement stage is not loaded")

        if generator is not None:
            generate_kwargs["generator"] = generator
        generated = cast(_GenerationModel, self.placement).generate(
            input_ids=input_ids,
            attention_mask=attention_mask,
            max_length=generation_max_length
            or self.config.placement_generation_max_length,
            num_return_sequences=num_return_sequences
            or self.config.num_return_sequences,
            temperature=temperature or self.config.temperature,
            do_sample=do_sample,
            **generate_kwargs,
        )
        return generated

    def __call__(
        self,
        *,
        prompt: str | Sequence[str] | None = None,
        batch_size: int = 1,
        seed: int | None = None,
        generator: torch.Generator | None = None,
        condition_type: ConditionType | str = ConditionType.text,
        labels: Int[torch.Tensor, "batch elements"]
        | list[ArrayLikeInput]
        | None = None,
        bbox: Float[torch.Tensor, "batch elements 4"]
        | list[ArrayLikeInput]
        | None = None,
        mask: Bool[torch.Tensor, "batch elements"] | list[ArrayLikeInput] | 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,
        num_return_sequences: int | None = None,
        temperature: float | None = None,
        output_candidate: Literal["first", "all", "best"] = "first",
        output_type: Literal["dataclass", "dict"] = "dataclass",
        return_intermediates: bool = False,
        layout_text: str | list[str] | list[list[str]] | None = None,
    ) -> LayoutGenerationOutput | ParseThenPlaceOutputDict:  # ty: ignore[invalid-method-override]
        """Generate a layout from natural-language text."""
        _ = (
            batch_size,
            labels,
            bbox,
            mask,
            num_elements,
            box_format,
            normalized,
            canvas_size,
            num_inference_steps,
        )
        condition = normalize_condition_type(condition_type)
        if condition is not ConditionType.text:
            raise NotImplementedError(
                "Parse-Then-Place only supports condition_type='text'"
            )

        if layout_text is not None:
            layout_items = (
                [layout_text] if isinstance(layout_text, str) else layout_text
            )
            return self.processor.layout_text_to_output(
                layout_items,
                output_candidate=output_candidate,
                output_type=output_type,
                return_intermediates=return_intermediates,
            )
        if prompt is None:
            raise ValueError("prompt is required for Parse-Then-Place generation")

        generation_generator = self.prepare_generator(
            generator=generator,
            seed=seed,
        )
        prompts = [prompt] if isinstance(prompt, str) else list(prompt)
        parser_inputs = self.processor(prompts)
        if "input_ids" not in parser_inputs:
            raise ValueError("processor requires parser_tokenizer for model inference")

        parser_ids = self.parse(
            parser_inputs["input_ids"],
            attention_mask=parser_inputs.get("attention_mask"),
            generation_max_length=self.config.parser_generation_max_length,
        )
        value_maps = cast(list[dict[str, str] | None], parser_inputs.get("value_maps"))
        logical_forms = self.processor.postprocess_ir(
            parser_ids,
            value_maps=value_maps,
        )
        placement_inputs = self.processor.ir_to_placement_inputs(logical_forms)
        placement_encoded = self.processor.encode_placement_inputs(placement_inputs)
        if "input_ids" not in placement_encoded:
            raise ValueError(
                "processor requires placement_tokenizer for model inference"
            )

        return_sequences = num_return_sequences or self.config.num_return_sequences
        placement_ids = self.place(
            placement_encoded["input_ids"],
            attention_mask=placement_encoded.get("attention_mask"),
            num_return_sequences=return_sequences,
            temperature=temperature,
            generator=generation_generator,
        )
        grouped = self.processor.decode_layout_sequences(
            placement_ids,
            batch_size=len(prompts),
            num_return_sequences=return_sequences,
        )
        output = self.processor.layout_text_to_output(
            grouped,
            output_candidate=output_candidate,
            output_type="dataclass",
            return_intermediates=True,
        )
        if isinstance(output, LayoutGenerationOutput):
            intermediates = (
                dict(output.intermediates)
                if isinstance(output.intermediates, dict)
                else {}
            )
            if return_intermediates:
                intermediates.update(
                    {
                        "prompt": prompts,
                        "logical_forms": logical_forms,
                        "placement_inputs": placement_inputs,
                    }
                )
            output.intermediates = intermediates if return_intermediates else None
        if output_type == "dict":
            return dict(output)
        return output

    generate = __call__

__init__

__init__(
    config: ParseThenPlaceConfig,
    processor: ParseThenPlaceProcessor,
    *,
    parser: PreTrainedModel | None = None,
    placement: PreTrainedModel | None = None,
) -> None

Initialize the composite pipeline.

Source code in models/parse-then-place/src/parse_then_place/pipeline_parse_then_place.py
111
112
113
114
115
116
117
118
119
120
121
122
123
124
def __init__(
    self,
    config: ParseThenPlaceConfig,
    processor: ParseThenPlaceProcessor,
    *,
    parser: PreTrainedModel | None = None,
    placement: PreTrainedModel | None = None,
) -> None:
    """Initialize the composite pipeline."""
    super().__init__(config)
    self.config = config
    self.processor = processor
    self.parser = parser
    self.placement = placement

from_pretrained classmethod

from_pretrained(
    pretrained_model_name_or_path: str | Path,
    *,
    parser: PreTrainedModel | None = None,
    placement: PreTrainedModel | None = None,
    processor: ParseThenPlaceProcessor | None = None,
    local_files_only: bool = False,
    config: ParseThenPlaceConfig
    | PretrainedConfig
    | None = None,
) -> ParseThenPlacePipeline

Load a composite pipeline from a root directory.

Source code in models/parse-then-place/src/parse_then_place/pipeline_parse_then_place.py
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
@classmethod
def from_pretrained(
    cls,
    pretrained_model_name_or_path: str | Path,
    *,
    parser: PreTrainedModel | None = None,
    placement: PreTrainedModel | None = None,
    processor: ParseThenPlaceProcessor | None = None,
    local_files_only: bool = False,
    config: ParseThenPlaceConfig | PretrainedConfig | None = None,
) -> ParseThenPlacePipeline:  # ty: ignore[invalid-method-override]
    """Load a composite pipeline from a root directory."""
    components: dict[str, PreTrainedModel | ParseThenPlaceProcessor] = {}
    if parser is not None:
        components["parser"] = parser
    if placement is not None:
        components["placement"] = placement
    if processor is None:
        loaded = super().from_pretrained(
            pretrained_model_name_or_path,
            local_files_only=local_files_only,
            config=config,
            components=components,
        )
        return cast(ParseThenPlacePipeline, loaded)
    components["processor"] = processor
    loaded = super().from_pretrained(
        pretrained_model_name_or_path,
        local_files_only=local_files_only,
        config=config,
        components=components,
    )
    return cast(ParseThenPlacePipeline, loaded)

parse

parse(
    input_ids: Int[Tensor, "batch tokens"],
    attention_mask: Bool[Tensor, "batch tokens"]
    | None = None,
    *,
    generation_max_length: int | None = None,
    **generate_kwargs: str
    | int
    | float
    | bool
    | Generator
    | None,
) -> Int[torch.Tensor, "batch tokens"]

Generate logical-form token ids with the parser stage.

Source code in models/parse-then-place/src/parse_then_place/pipeline_parse_then_place.py
175
176
177
178
179
180
181
182
183
184
185
186
187
188
189
190
191
192
193
194
195
@torch.no_grad()
def parse(
    self,
    input_ids: Int[torch.Tensor, "batch tokens"],
    attention_mask: Bool[torch.Tensor, "batch tokens"] | None = None,
    *,
    generation_max_length: int | None = None,
    **generate_kwargs: str | int | float | bool | torch.Generator | None,
) -> Int[torch.Tensor, "batch tokens"]:
    """Generate logical-form token ids with the parser stage."""
    if self.parser is None:
        raise ValueError("Parser stage is not loaded")

    generated = cast(_GenerationModel, self.parser).generate(
        input_ids=input_ids,
        attention_mask=attention_mask,
        max_length=generation_max_length
        or self.config.parser_generation_max_length,
        **generate_kwargs,
    )
    return generated

place

place(
    input_ids: Int[Tensor, "batch tokens"],
    attention_mask: Bool[Tensor, "batch tokens"]
    | None = None,
    *,
    generation_max_length: int | None = None,
    num_return_sequences: int | None = None,
    temperature: float | None = None,
    do_sample: bool = True,
    generator: Generator | None = None,
    **generate_kwargs: str
    | float
    | bool
    | Generator
    | None,
) -> Int[torch.Tensor, "batch tokens"]

Generate layout token ids with the placement stage.

Source code in models/parse-then-place/src/parse_then_place/pipeline_parse_then_place.py
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
@torch.no_grad()
def place(
    self,
    input_ids: Int[torch.Tensor, "batch tokens"],
    attention_mask: Bool[torch.Tensor, "batch tokens"] | None = None,
    *,
    generation_max_length: int | None = None,
    num_return_sequences: int | None = None,
    temperature: float | None = None,
    do_sample: bool = True,
    generator: torch.Generator | None = None,
    **generate_kwargs: str | float | bool | torch.Generator | None,
) -> Int[torch.Tensor, "batch tokens"]:
    """Generate layout token ids with the placement stage."""
    if self.placement is None:
        raise ValueError("Placement stage is not loaded")

    if generator is not None:
        generate_kwargs["generator"] = generator
    generated = cast(_GenerationModel, self.placement).generate(
        input_ids=input_ids,
        attention_mask=attention_mask,
        max_length=generation_max_length
        or self.config.placement_generation_max_length,
        num_return_sequences=num_return_sequences
        or self.config.num_return_sequences,
        temperature=temperature or self.config.temperature,
        do_sample=do_sample,
        **generate_kwargs,
    )
    return generated

__call__

__call__(
    *,
    prompt: str | Sequence[str] | None = None,
    batch_size: int = 1,
    seed: int | None = None,
    generator: Generator | None = None,
    condition_type: ConditionType
    | str = ConditionType.text,
    labels: Int[Tensor, "batch elements"]
    | list[ArrayLikeInput]
    | None = None,
    bbox: Float[Tensor, "batch elements 4"]
    | list[ArrayLikeInput]
    | None = None,
    mask: Bool[Tensor, "batch elements"]
    | list[ArrayLikeInput]
    | 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,
    num_return_sequences: int | None = None,
    temperature: float | None = None,
    output_candidate: Literal[
        "first", "all", "best"
    ] = "first",
    output_type: Literal["dataclass", "dict"] = "dataclass",
    return_intermediates: bool = False,
    layout_text: str
    | list[str]
    | list[list[str]]
    | None = None,
) -> LayoutGenerationOutput | ParseThenPlaceOutputDict

Generate a layout from natural-language text.

Source code in models/parse-then-place/src/parse_then_place/pipeline_parse_then_place.py
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
def __call__(
    self,
    *,
    prompt: str | Sequence[str] | None = None,
    batch_size: int = 1,
    seed: int | None = None,
    generator: torch.Generator | None = None,
    condition_type: ConditionType | str = ConditionType.text,
    labels: Int[torch.Tensor, "batch elements"]
    | list[ArrayLikeInput]
    | None = None,
    bbox: Float[torch.Tensor, "batch elements 4"]
    | list[ArrayLikeInput]
    | None = None,
    mask: Bool[torch.Tensor, "batch elements"] | list[ArrayLikeInput] | 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,
    num_return_sequences: int | None = None,
    temperature: float | None = None,
    output_candidate: Literal["first", "all", "best"] = "first",
    output_type: Literal["dataclass", "dict"] = "dataclass",
    return_intermediates: bool = False,
    layout_text: str | list[str] | list[list[str]] | None = None,
) -> LayoutGenerationOutput | ParseThenPlaceOutputDict:  # ty: ignore[invalid-method-override]
    """Generate a layout from natural-language text."""
    _ = (
        batch_size,
        labels,
        bbox,
        mask,
        num_elements,
        box_format,
        normalized,
        canvas_size,
        num_inference_steps,
    )
    condition = normalize_condition_type(condition_type)
    if condition is not ConditionType.text:
        raise NotImplementedError(
            "Parse-Then-Place only supports condition_type='text'"
        )

    if layout_text is not None:
        layout_items = (
            [layout_text] if isinstance(layout_text, str) else layout_text
        )
        return self.processor.layout_text_to_output(
            layout_items,
            output_candidate=output_candidate,
            output_type=output_type,
            return_intermediates=return_intermediates,
        )
    if prompt is None:
        raise ValueError("prompt is required for Parse-Then-Place generation")

    generation_generator = self.prepare_generator(
        generator=generator,
        seed=seed,
    )
    prompts = [prompt] if isinstance(prompt, str) else list(prompt)
    parser_inputs = self.processor(prompts)
    if "input_ids" not in parser_inputs:
        raise ValueError("processor requires parser_tokenizer for model inference")

    parser_ids = self.parse(
        parser_inputs["input_ids"],
        attention_mask=parser_inputs.get("attention_mask"),
        generation_max_length=self.config.parser_generation_max_length,
    )
    value_maps = cast(list[dict[str, str] | None], parser_inputs.get("value_maps"))
    logical_forms = self.processor.postprocess_ir(
        parser_ids,
        value_maps=value_maps,
    )
    placement_inputs = self.processor.ir_to_placement_inputs(logical_forms)
    placement_encoded = self.processor.encode_placement_inputs(placement_inputs)
    if "input_ids" not in placement_encoded:
        raise ValueError(
            "processor requires placement_tokenizer for model inference"
        )

    return_sequences = num_return_sequences or self.config.num_return_sequences
    placement_ids = self.place(
        placement_encoded["input_ids"],
        attention_mask=placement_encoded.get("attention_mask"),
        num_return_sequences=return_sequences,
        temperature=temperature,
        generator=generation_generator,
    )
    grouped = self.processor.decode_layout_sequences(
        placement_ids,
        batch_size=len(prompts),
        num_return_sequences=return_sequences,
    )
    output = self.processor.layout_text_to_output(
        grouped,
        output_candidate=output_candidate,
        output_type="dataclass",
        return_intermediates=True,
    )
    if isinstance(output, LayoutGenerationOutput):
        intermediates = (
            dict(output.intermediates)
            if isinstance(output.intermediates, dict)
            else {}
        )
        if return_intermediates:
            intermediates.update(
                {
                    "prompt": prompts,
                    "logical_forms": logical_forms,
                    "placement_inputs": placement_inputs,
                }
            )
        output.intermediates = intermediates if return_intermediates else None
    if output_type == "dict":
        return dict(output)
    return output

ParseThenPlaceProcessor

Bases: ProcessorMixin

Build stage inputs and parse placement-model output text.

Source code in models/parse-then-place/src/parse_then_place/processing_parse_then_place.py
 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
class ParseThenPlaceProcessor(ProcessorMixin):
    """Build stage inputs and parse placement-model output text."""

    attributes = ["parser_tokenizer", "placement_tokenizer"]
    parser_tokenizer_class = "AutoTokenizer"
    placement_tokenizer_class = "T5Tokenizer"

    def __init__(
        self,
        parser_tokenizer: PreTrainedTokenizerBase | None = None,
        placement_tokenizer: PreTrainedTokenizerBase | None = None,
        dataset_name: ParseThenPlaceDatasetName | str = ParseThenPlaceDatasetName.rico,
        canvas_size: tuple[int, int] | None = None,
        id2label: dict[int, str] | None = None,
    ) -> None:
        """Initialize tokenizer handles and dataset metadata."""
        dataset = normalize_dataset_name(dataset_name)
        self.parser_tokenizer = parser_tokenizer
        self.placement_tokenizer = placement_tokenizer
        self.dataset_name = str(dataset)
        self.canvas_size = canvas_size or canvas_size_for_dataset(dataset)
        self.id2label = (
            {int(key): str(value) for key, value in id2label.items()}
            if id2label is not None
            else id2label_for_dataset(dataset)
        )
        self.label2id = {label.lower(): idx for idx, label in self.id2label.items()}
        # Keep released spellings as aliases because RICO uses lower-case labels.
        self.label2id.update(label2id_for_dataset(dataset))
        if parser_tokenizer is not None and placement_tokenizer is not None:
            super().__init__(
                parser_tokenizer=parser_tokenizer,
                placement_tokenizer=placement_tokenizer,
            )

    @classmethod
    def from_config(
        cls,
        dataset_name: ParseThenPlaceDatasetName | str = ParseThenPlaceDatasetName.rico,
        *,
        canvas_size: tuple[int, int] | None = None,
        id2label: dict[int, str] | None = None,
    ) -> ParseThenPlaceProcessor:
        """Construct metadata-only processor for tests and local smoke checks."""
        return cls(
            parser_tokenizer=None,
            placement_tokenizer=None,
            dataset_name=dataset_name,
            canvas_size=canvas_size,
            id2label=id2label,
        )

    def preprocess_prompt(
        self,
        prompt: str,
        *,
        replace_explicit_value: bool = True,
    ) -> PromptEncoding:
        """Apply the released text normalization used before stage-1 parsing.

        Args:
            prompt: Natural-language text prompt.
            replace_explicit_value: Whether quoted values should be replaced by
                deterministic ``value_N`` placeholders.

        Returns:
            Normalized prompt and the placeholder recovery map.
        """
        result = prompt.replace("#", "").strip().lower()
        result = (
            result.replace("“", '"')
            .replace("”", '"')
            .replace("‘", "'")
            .replace("’", "'")
        )
        result = _WHITESPACE_RE.sub(" ", result)
        if not replace_explicit_value:
            return {"prompt": result, "value_map": None}
        return self._extract_explicit_values(result)

    def _extract_explicit_values(self, prompt: str) -> PromptEncoding:
        result = re.sub(r"(\w)'s\s+", r"\g<1>`s ", prompt)
        values = _DOUBLE_QUOTED_RE.findall(result)
        single_values = _SINGLE_QUOTED_RE.findall(result)
        if len(single_values) == 1 and any(
            punct in single_values[0] for punct in (",", ".")
        ):
            single_values = []
        values.extend(single_values)
        value_map: dict[str, str] = {}
        for value_idx, value in enumerate(values):
            placeholder = f"value_{value_idx}"
            value_map[placeholder] = value.strip('"').strip("'").strip()
            result = result.replace(value, f'"{placeholder}"', 1)
        result = re.sub(r"(\w)`s\s+", r"\g<1>'s ", result)
        result = _WHITESPACE_RE.sub(" ", result)
        return {"prompt": result, "value_map": value_map}

    def __call__(
        self,
        prompt: str | Sequence[str],
        *,
        replace_explicit_value: bool = True,
        return_tensors: Literal["pt"] = "pt",
    ) -> BatchEncoding:
        """Tokenize prompt text for the semantic parser stage."""
        prompts = [prompt] if isinstance(prompt, str) else list(prompt)
        encodings = [
            self.preprocess_prompt(item, replace_explicit_value=replace_explicit_value)
            for item in prompts
        ]
        texts = [item["prompt"] for item in encodings]
        if self.parser_tokenizer is None:
            return BatchEncoding(
                {
                    "prompt_text": texts,
                    "value_maps": [item["value_map"] for item in encodings],
                }
            )
        parser_tokenizer = self.parser_tokenizer
        tokenized = parser_tokenizer(texts, return_tensors=return_tensors, padding=True)
        tokenized["value_maps"] = [item["value_map"] for item in encodings]
        tokenized["prompt_text"] = texts
        return cast(BatchEncoding, tokenized)

    def postprocess_ir(
        self,
        generated_ids: Int[torch.Tensor, "batch tokens"] | Sequence[str],
        *,
        value_maps: list[dict[str, str] | None] | None = None,
    ) -> list[str]:
        """Decode and lightly normalize stage-1 logical forms."""
        if isinstance(generated_ids, torch.Tensor):
            if self.parser_tokenizer is None:
                raise ValueError("parser_tokenizer is required to decode generated ids")

            parser_tokenizer = self.parser_tokenizer
            logical_forms = parser_tokenizer.batch_decode(
                generated_ids, skip_special_tokens=True
            )
        else:
            logical_forms = list(generated_ids)
        if value_maps is None:
            return [self._normalize_ir(item) for item in logical_forms]
        return [
            self._recover_ir_values(self._normalize_ir(item), value_map)
            for item, value_map in zip(logical_forms, value_maps, strict=True)
        ]

    def _normalize_ir(self, logical_form: str) -> str:
        result = logical_form.replace("[", " [ ").replace("]", " ] ").strip().lower()
        return _WHITESPACE_RE.sub(" ", result)

    def _recover_ir_values(
        self,
        logical_form: str,
        value_map: dict[str, str] | None,
    ) -> str:
        if not value_map:
            return logical_form
        result = logical_form
        for placeholder, value in value_map.items():
            recovered = value.replace("'", "")
            result = result.replace(f"'{placeholder}'", f"'{recovered}'")
            result = result.replace(f" {placeholder},", f" {value},")
            result = result.replace(f" {placeholder}'", f" {value}'")
        return result.replace("&", " and ")

    def ir_to_placement_inputs(self, logical_forms: Sequence[str]) -> list[str]:
        """Convert logical forms to placement-constraint strings.

        Runtime keeps this method deterministic and accepts already-linearized
        constraints, which is also the artifact stored in stage-1 prediction JSON
        files. The current parity scripts do not execute the released grammar
        executor.
        """
        return [self._logical_form_to_constraint(item) for item in logical_forms]

    def _logical_form_to_constraint(self, logical_form: str) -> str:
        text = _WHITESPACE_RE.sub(" ", logical_form.strip())
        if ":" in text and "|" in text:
            return text
        return text

    def encode_placement_inputs(
        self,
        placement_inputs: Sequence[str],
        *,
        return_tensors: Literal["pt"] = "pt",
    ) -> BatchEncoding:
        """Tokenize stage-2 placement constraints."""
        if self.placement_tokenizer is None:
            return BatchEncoding({"placement_text": list(placement_inputs)})
        placement_tokenizer = self.placement_tokenizer
        tokenized = placement_tokenizer(
            list(placement_inputs), return_tensors=return_tensors, padding=True
        )
        tokenized["placement_text"] = list(placement_inputs)
        return cast(BatchEncoding, tokenized)

    def decode_layout_sequences(
        self,
        generated_ids: Int[torch.Tensor, "batch tokens"] | Sequence[str],
        *,
        batch_size: int,
        num_return_sequences: int,
    ) -> list[list[str]]:
        """Decode stage-2 generated ids into grouped layout strings."""
        if isinstance(generated_ids, torch.Tensor):
            if self.placement_tokenizer is None:
                raise ValueError(
                    "placement_tokenizer is required to decode generated ids"
                )

            placement_tokenizer = self.placement_tokenizer
            flat = placement_tokenizer.batch_decode(
                generated_ids, skip_special_tokens=True
            )
        else:
            flat = list(generated_ids)
        expected = batch_size * num_return_sequences
        if len(flat) != expected:
            raise ValueError(
                "Generated layout count does not match batch_size * num_return_sequences: "
                f"{len(flat)} != {expected}"
            )

        return [
            flat[idx * num_return_sequences : (idx + 1) * num_return_sequences]
            for idx in range(batch_size)
        ]

    def layout_text_to_output(
        self,
        layout_text: Sequence[str] | Sequence[Sequence[str]],
        *,
        output_candidate: Literal["first", "all", "best"] = "first",
        output_type: Literal["dataclass", "dict"] = "dataclass",
        return_intermediates: bool = False,
    ) -> LayoutGenerationOutput | ParseThenPlaceOutputDict:
        """Parse generated ``label left top width height`` text into schema."""
        candidate_groups = self._normalize_layout_text_groups(layout_text)
        selected = self._select_candidates(candidate_groups, output_candidate)
        parsed_groups = [self._parse_layout_text(item) for item in selected]
        max_len = max((len(item) for item in parsed_groups), default=0) or 1

        bbox_rows: list[Float[torch.Tensor, "elements 4"]] = []
        label_rows: list[Int[torch.Tensor, ...]] = []
        mask_rows: list[Bool[torch.Tensor, ...]] = []

        for parsed in parsed_groups:
            labels = torch.tensor([item["label"] for item in parsed], dtype=torch.long)
            boxes = torch.tensor([item["bbox"] for item in parsed], dtype=torch.float32)
            mask = torch.ones(len(parsed), dtype=torch.bool)

            if len(parsed) == 0:
                labels = torch.zeros(max_len, dtype=torch.long)
                boxes = torch.zeros(max_len, 4, dtype=torch.float32)
                mask = torch.zeros(max_len, dtype=torch.bool)
            elif len(parsed) < max_len:
                pad = max_len - len(parsed)
                labels = torch.nn.functional.pad(labels, (0, pad))
                boxes = torch.nn.functional.pad(boxes, (0, 0, 0, pad))
                mask = torch.nn.functional.pad(mask, (0, pad))

            label_rows.append(labels)
            bbox_rows.append(boxes)
            mask_rows.append(mask)

        raw_bbox = torch.stack(bbox_rows)
        bbox = normalize_boxes(
            raw_bbox,
            canvas_size=self.canvas_size,
            box_format="ltwh",
        )
        output = LayoutGenerationOutput(
            bbox=bbox.float(),
            labels=torch.stack(label_rows).long(),
            mask=torch.stack(mask_rows).bool(),
            id2label=dict(self.id2label),
            intermediates={
                "layout_text": selected,
                "layout_text_candidates": candidate_groups,
                "dataset_name": self.dataset_name,
                "canvas_size": self.canvas_size,
            }
            if return_intermediates
            else None,
        )
        if output_type == "dict":
            return dict(output)
        if output_type != "dataclass":
            raise ValueError(f"Unsupported output_type: {output_type}")

        return output

    def _normalize_layout_text_groups(
        self,
        layout_text: Sequence[str] | Sequence[Sequence[str]],
    ) -> list[list[str]]:
        if not layout_text:
            return []
        first = layout_text[0]
        if isinstance(first, str):
            return [[item] for item in cast(Sequence[str], layout_text)]
        return [list(item) for item in cast(Sequence[Sequence[str]], layout_text)]

    def _select_candidates(
        self,
        candidate_groups: list[list[str]],
        output_candidate: Literal["first", "all", "best"],
    ) -> list[str]:
        if output_candidate == "first":
            return [group[0] if group else "" for group in candidate_groups]
        if output_candidate == "best":
            return [
                max(group, key=lambda item: len(self._parse_layout_text(item)))
                if group
                else ""
                for group in candidate_groups
            ]
        if output_candidate == "all":
            return ["\n".join(group) for group in candidate_groups]
        raise ValueError(f"Unsupported output_candidate: {output_candidate}")

    def _parse_layout_text(self, layout_text: str) -> list[ParsedElement]:
        elements: list[ParsedElement] = []
        for match in _LAYOUT_PATTERN.finditer(layout_text.lower()):
            label = _WHITESPACE_RE.sub(" ", match.group("label")).strip()
            label_id = self.label2id.get(label)
            if label_id is None:
                continue
            elements.append(
                {
                    "label": label_id,
                    "bbox": [
                        float(match.group("left")),
                        float(match.group("top")),
                        float(match.group("width")),
                        float(match.group("height")),
                    ],
                }
            )
        return elements

__init__

__init__(
    parser_tokenizer: PreTrainedTokenizerBase | None = None,
    placement_tokenizer: PreTrainedTokenizerBase
    | None = None,
    dataset_name: ParseThenPlaceDatasetName
    | str = ParseThenPlaceDatasetName.rico,
    canvas_size: tuple[int, int] | None = None,
    id2label: dict[int, str] | None = None,
) -> None

Initialize tokenizer handles and dataset metadata.

Source code in models/parse-then-place/src/parse_then_place/processing_parse_then_place.py
 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
def __init__(
    self,
    parser_tokenizer: PreTrainedTokenizerBase | None = None,
    placement_tokenizer: PreTrainedTokenizerBase | None = None,
    dataset_name: ParseThenPlaceDatasetName | str = ParseThenPlaceDatasetName.rico,
    canvas_size: tuple[int, int] | None = None,
    id2label: dict[int, str] | None = None,
) -> None:
    """Initialize tokenizer handles and dataset metadata."""
    dataset = normalize_dataset_name(dataset_name)
    self.parser_tokenizer = parser_tokenizer
    self.placement_tokenizer = placement_tokenizer
    self.dataset_name = str(dataset)
    self.canvas_size = canvas_size or canvas_size_for_dataset(dataset)
    self.id2label = (
        {int(key): str(value) for key, value in id2label.items()}
        if id2label is not None
        else id2label_for_dataset(dataset)
    )
    self.label2id = {label.lower(): idx for idx, label in self.id2label.items()}
    # Keep released spellings as aliases because RICO uses lower-case labels.
    self.label2id.update(label2id_for_dataset(dataset))
    if parser_tokenizer is not None and placement_tokenizer is not None:
        super().__init__(
            parser_tokenizer=parser_tokenizer,
            placement_tokenizer=placement_tokenizer,
        )

from_config classmethod

from_config(
    dataset_name: ParseThenPlaceDatasetName
    | str = ParseThenPlaceDatasetName.rico,
    *,
    canvas_size: tuple[int, int] | None = None,
    id2label: dict[int, str] | None = None,
) -> ParseThenPlaceProcessor

Construct metadata-only processor for tests and local smoke checks.

Source code in models/parse-then-place/src/parse_then_place/processing_parse_then_place.py
104
105
106
107
108
109
110
111
112
113
114
115
116
117
118
119
@classmethod
def from_config(
    cls,
    dataset_name: ParseThenPlaceDatasetName | str = ParseThenPlaceDatasetName.rico,
    *,
    canvas_size: tuple[int, int] | None = None,
    id2label: dict[int, str] | None = None,
) -> ParseThenPlaceProcessor:
    """Construct metadata-only processor for tests and local smoke checks."""
    return cls(
        parser_tokenizer=None,
        placement_tokenizer=None,
        dataset_name=dataset_name,
        canvas_size=canvas_size,
        id2label=id2label,
    )

preprocess_prompt

preprocess_prompt(
    prompt: str, *, replace_explicit_value: bool = True
) -> PromptEncoding

Apply the released text normalization used before stage-1 parsing.

Parameters:

Name Type Description Default
prompt str

Natural-language text prompt.

required
replace_explicit_value bool

Whether quoted values should be replaced by deterministic value_N placeholders.

True

Returns:

Type Description
PromptEncoding

Normalized prompt and the placeholder recovery map.

Source code in models/parse-then-place/src/parse_then_place/processing_parse_then_place.py
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
def preprocess_prompt(
    self,
    prompt: str,
    *,
    replace_explicit_value: bool = True,
) -> PromptEncoding:
    """Apply the released text normalization used before stage-1 parsing.

    Args:
        prompt: Natural-language text prompt.
        replace_explicit_value: Whether quoted values should be replaced by
            deterministic ``value_N`` placeholders.

    Returns:
        Normalized prompt and the placeholder recovery map.
    """
    result = prompt.replace("#", "").strip().lower()
    result = (
        result.replace("“", '"')
        .replace("”", '"')
        .replace("‘", "'")
        .replace("’", "'")
    )
    result = _WHITESPACE_RE.sub(" ", result)
    if not replace_explicit_value:
        return {"prompt": result, "value_map": None}
    return self._extract_explicit_values(result)

__call__

__call__(
    prompt: str | Sequence[str],
    *,
    replace_explicit_value: bool = True,
    return_tensors: Literal["pt"] = "pt",
) -> BatchEncoding

Tokenize prompt text for the semantic parser stage.

Source code in models/parse-then-place/src/parse_then_place/processing_parse_then_place.py
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
def __call__(
    self,
    prompt: str | Sequence[str],
    *,
    replace_explicit_value: bool = True,
    return_tensors: Literal["pt"] = "pt",
) -> BatchEncoding:
    """Tokenize prompt text for the semantic parser stage."""
    prompts = [prompt] if isinstance(prompt, str) else list(prompt)
    encodings = [
        self.preprocess_prompt(item, replace_explicit_value=replace_explicit_value)
        for item in prompts
    ]
    texts = [item["prompt"] for item in encodings]
    if self.parser_tokenizer is None:
        return BatchEncoding(
            {
                "prompt_text": texts,
                "value_maps": [item["value_map"] for item in encodings],
            }
        )
    parser_tokenizer = self.parser_tokenizer
    tokenized = parser_tokenizer(texts, return_tensors=return_tensors, padding=True)
    tokenized["value_maps"] = [item["value_map"] for item in encodings]
    tokenized["prompt_text"] = texts
    return cast(BatchEncoding, tokenized)

postprocess_ir

postprocess_ir(
    generated_ids: Int[Tensor, "batch tokens"]
    | Sequence[str],
    *,
    value_maps: list[dict[str, str] | None] | None = None,
) -> list[str]

Decode and lightly normalize stage-1 logical forms.

Source code in models/parse-then-place/src/parse_then_place/processing_parse_then_place.py
194
195
196
197
198
199
200
201
202
203
204
205
206
207
208
209
210
211
212
213
214
215
216
def postprocess_ir(
    self,
    generated_ids: Int[torch.Tensor, "batch tokens"] | Sequence[str],
    *,
    value_maps: list[dict[str, str] | None] | None = None,
) -> list[str]:
    """Decode and lightly normalize stage-1 logical forms."""
    if isinstance(generated_ids, torch.Tensor):
        if self.parser_tokenizer is None:
            raise ValueError("parser_tokenizer is required to decode generated ids")

        parser_tokenizer = self.parser_tokenizer
        logical_forms = parser_tokenizer.batch_decode(
            generated_ids, skip_special_tokens=True
        )
    else:
        logical_forms = list(generated_ids)
    if value_maps is None:
        return [self._normalize_ir(item) for item in logical_forms]
    return [
        self._recover_ir_values(self._normalize_ir(item), value_map)
        for item, value_map in zip(logical_forms, value_maps, strict=True)
    ]

ir_to_placement_inputs

ir_to_placement_inputs(
    logical_forms: Sequence[str],
) -> list[str]

Convert logical forms to placement-constraint strings.

Runtime keeps this method deterministic and accepts already-linearized constraints, which is also the artifact stored in stage-1 prediction JSON files. The current parity scripts do not execute the released grammar executor.

Source code in models/parse-then-place/src/parse_then_place/processing_parse_then_place.py
237
238
239
240
241
242
243
244
245
def ir_to_placement_inputs(self, logical_forms: Sequence[str]) -> list[str]:
    """Convert logical forms to placement-constraint strings.

    Runtime keeps this method deterministic and accepts already-linearized
    constraints, which is also the artifact stored in stage-1 prediction JSON
    files. The current parity scripts do not execute the released grammar
    executor.
    """
    return [self._logical_form_to_constraint(item) for item in logical_forms]

encode_placement_inputs

encode_placement_inputs(
    placement_inputs: Sequence[str],
    *,
    return_tensors: Literal["pt"] = "pt",
) -> BatchEncoding

Tokenize stage-2 placement constraints.

Source code in models/parse-then-place/src/parse_then_place/processing_parse_then_place.py
253
254
255
256
257
258
259
260
261
262
263
264
265
266
267
def encode_placement_inputs(
    self,
    placement_inputs: Sequence[str],
    *,
    return_tensors: Literal["pt"] = "pt",
) -> BatchEncoding:
    """Tokenize stage-2 placement constraints."""
    if self.placement_tokenizer is None:
        return BatchEncoding({"placement_text": list(placement_inputs)})
    placement_tokenizer = self.placement_tokenizer
    tokenized = placement_tokenizer(
        list(placement_inputs), return_tensors=return_tensors, padding=True
    )
    tokenized["placement_text"] = list(placement_inputs)
    return cast(BatchEncoding, tokenized)

decode_layout_sequences

decode_layout_sequences(
    generated_ids: Int[Tensor, "batch tokens"]
    | Sequence[str],
    *,
    batch_size: int,
    num_return_sequences: int,
) -> list[list[str]]

Decode stage-2 generated ids into grouped layout strings.

Source code in models/parse-then-place/src/parse_then_place/processing_parse_then_place.py
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
def decode_layout_sequences(
    self,
    generated_ids: Int[torch.Tensor, "batch tokens"] | Sequence[str],
    *,
    batch_size: int,
    num_return_sequences: int,
) -> list[list[str]]:
    """Decode stage-2 generated ids into grouped layout strings."""
    if isinstance(generated_ids, torch.Tensor):
        if self.placement_tokenizer is None:
            raise ValueError(
                "placement_tokenizer is required to decode generated ids"
            )

        placement_tokenizer = self.placement_tokenizer
        flat = placement_tokenizer.batch_decode(
            generated_ids, skip_special_tokens=True
        )
    else:
        flat = list(generated_ids)
    expected = batch_size * num_return_sequences
    if len(flat) != expected:
        raise ValueError(
            "Generated layout count does not match batch_size * num_return_sequences: "
            f"{len(flat)} != {expected}"
        )

    return [
        flat[idx * num_return_sequences : (idx + 1) * num_return_sequences]
        for idx in range(batch_size)
    ]

layout_text_to_output

layout_text_to_output(
    layout_text: Sequence[str] | Sequence[Sequence[str]],
    *,
    output_candidate: Literal[
        "first", "all", "best"
    ] = "first",
    output_type: Literal["dataclass", "dict"] = "dataclass",
    return_intermediates: bool = False,
) -> LayoutGenerationOutput | ParseThenPlaceOutputDict

Parse generated label left top width height text into schema.

Source code in models/parse-then-place/src/parse_then_place/processing_parse_then_place.py
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
def layout_text_to_output(
    self,
    layout_text: Sequence[str] | Sequence[Sequence[str]],
    *,
    output_candidate: Literal["first", "all", "best"] = "first",
    output_type: Literal["dataclass", "dict"] = "dataclass",
    return_intermediates: bool = False,
) -> LayoutGenerationOutput | ParseThenPlaceOutputDict:
    """Parse generated ``label left top width height`` text into schema."""
    candidate_groups = self._normalize_layout_text_groups(layout_text)
    selected = self._select_candidates(candidate_groups, output_candidate)
    parsed_groups = [self._parse_layout_text(item) for item in selected]
    max_len = max((len(item) for item in parsed_groups), default=0) or 1

    bbox_rows: list[Float[torch.Tensor, "elements 4"]] = []
    label_rows: list[Int[torch.Tensor, ...]] = []
    mask_rows: list[Bool[torch.Tensor, ...]] = []

    for parsed in parsed_groups:
        labels = torch.tensor([item["label"] for item in parsed], dtype=torch.long)
        boxes = torch.tensor([item["bbox"] for item in parsed], dtype=torch.float32)
        mask = torch.ones(len(parsed), dtype=torch.bool)

        if len(parsed) == 0:
            labels = torch.zeros(max_len, dtype=torch.long)
            boxes = torch.zeros(max_len, 4, dtype=torch.float32)
            mask = torch.zeros(max_len, dtype=torch.bool)
        elif len(parsed) < max_len:
            pad = max_len - len(parsed)
            labels = torch.nn.functional.pad(labels, (0, pad))
            boxes = torch.nn.functional.pad(boxes, (0, 0, 0, pad))
            mask = torch.nn.functional.pad(mask, (0, pad))

        label_rows.append(labels)
        bbox_rows.append(boxes)
        mask_rows.append(mask)

    raw_bbox = torch.stack(bbox_rows)
    bbox = normalize_boxes(
        raw_bbox,
        canvas_size=self.canvas_size,
        box_format="ltwh",
    )
    output = LayoutGenerationOutput(
        bbox=bbox.float(),
        labels=torch.stack(label_rows).long(),
        mask=torch.stack(mask_rows).bool(),
        id2label=dict(self.id2label),
        intermediates={
            "layout_text": selected,
            "layout_text_candidates": candidate_groups,
            "dataset_name": self.dataset_name,
            "canvas_size": self.canvas_size,
        }
        if return_intermediates
        else None,
    )
    if output_type == "dict":
        return dict(output)
    if output_type != "dataclass":
        raise ValueError(f"Unsupported output_type: {output_type}")

    return output

canvas_size_for_dataset

canvas_size_for_dataset(
    dataset_name: ParseThenPlaceDatasetName | str,
) -> tuple[int, int]

Return the dataset canvas size as (width, height).

Source code in models/parse-then-place/src/parse_then_place/labels.py
127
128
129
130
131
def canvas_size_for_dataset(
    dataset_name: ParseThenPlaceDatasetName | str,
) -> tuple[int, int]:
    """Return the dataset canvas size as ``(width, height)``."""
    return dataset_metadata(dataset_name)["canvas_size"]

id2label_for_dataset

id2label_for_dataset(
    dataset_name: ParseThenPlaceDatasetName | str,
) -> dict[int, str]

Return the dataset-local integer-id label map.

Source code in models/parse-then-place/src/parse_then_place/labels.py
111
112
113
114
115
def id2label_for_dataset(
    dataset_name: ParseThenPlaceDatasetName | str,
) -> dict[int, str]:
    """Return the dataset-local integer-id label map."""
    return dict(dataset_metadata(dataset_name)["id2label"])

label2id_for_dataset

label2id_for_dataset(
    dataset_name: ParseThenPlaceDatasetName | str,
) -> dict[str, int]

Return lower-case label names mapped to dataset-local ids.

Source code in models/parse-then-place/src/parse_then_place/labels.py
118
119
120
121
122
123
124
def label2id_for_dataset(
    dataset_name: ParseThenPlaceDatasetName | str,
) -> dict[str, int]:
    """Return lower-case label names mapped to dataset-local ids."""
    return {
        label.lower(): idx for idx, label in id2label_for_dataset(dataset_name).items()
    }

normalize_dataset_name

normalize_dataset_name(
    dataset_name: ParseThenPlaceDatasetName | str,
) -> ParseThenPlaceDatasetName

Normalize Parse-Then-Place dataset names.

Parameters:

Name Type Description Default
dataset_name ParseThenPlaceDatasetName | str

Dataset enum value or public/release string.

required

Returns:

Type Description
ParseThenPlaceDatasetName

Canonical Parse-Then-Place dataset name.

Raises:

Type Description
ValueError

If the dataset is unsupported.

Source code in models/parse-then-place/src/parse_then_place/labels.py
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
def normalize_dataset_name(
    dataset_name: ParseThenPlaceDatasetName | str,
) -> ParseThenPlaceDatasetName:
    """Normalize Parse-Then-Place dataset names.

    Args:
        dataset_name: Dataset enum value or public/release string.

    Returns:
        Canonical Parse-Then-Place dataset name.

    Raises:
        ValueError: If the dataset is unsupported.
    """
    if isinstance(dataset_name, ParseThenPlaceDatasetName):
        return dataset_name
    key = dataset_name.lower().replace("-", "_")
    if key == "webui":
        key = "web"
    try:
        return ParseThenPlaceDatasetName(key)
    except ValueError as exc:
        raise ValueError(
            f"Unsupported Parse-Then-Place dataset: {dataset_name}"
        ) from exc

normalize_stage2_mode

normalize_stage2_mode(
    stage2_mode: Stage2Mode | str,
) -> Stage2Mode

Normalize a released stage-2 checkpoint mode.

Source code in models/parse-then-place/src/parse_then_place/labels.py
 94
 95
 96
 97
 98
 99
100
101
def normalize_stage2_mode(stage2_mode: Stage2Mode | str) -> Stage2Mode:
    """Normalize a released stage-2 checkpoint mode."""
    if isinstance(stage2_mode, Stage2Mode):
        return stage2_mode
    try:
        return Stage2Mode(stage2_mode.lower().replace("-", "_"))
    except ValueError as exc:
        raise ValueError(f"Unsupported stage2_mode: {stage2_mode}") from exc

configuration_parse_then_place

Configuration for Parse-Then-Place composite checkpoints.

ParseThenPlaceConfig

Bases: PretrainedConfig

Stores Parse-Then-Place dataset and generation defaults.

Source code in models/parse-then-place/src/parse_then_place/configuration_parse_then_place.py
18
19
20
21
22
23
24
25
26
27
28
29
30
31
32
33
34
35
36
37
38
39
40
41
42
43
44
45
46
47
48
49
50
51
52
53
54
55
56
57
58
59
60
61
62
63
64
65
66
67
68
69
70
71
72
73
74
75
76
77
78
79
80
81
82
83
84
85
86
87
88
89
90
91
92
93
94
95
96
97
98
class ParseThenPlaceConfig(PretrainedConfig):
    """Stores Parse-Then-Place dataset and generation defaults."""

    model_type = "parse-then-place"

    def __init__(
        self,
        dataset_name: str = "rico",
        stage2_mode: Stage2Mode | str = Stage2Mode.finetune,
        parser_model_name: str = "google/t5-v1_1-base",
        parser_generation_max_length: int = 600,
        placement_generation_max_length: int = 500,
        temperature: float = 0.7,
        num_return_sequences: int = 5,
        canvas_size: tuple[int, int] | list[int] | None = None,
        id2label: dict[int | str, str] | None = None,
        parser_subfolder: str = "semantic_parser",
        placement_subfolder: str = "placement",
        pad_token_id: int = 0,
        eos_token_id: int = 1,
        decoder_start_token_id: int = 0,
        is_encoder_decoder: bool = True,
        transformers_version: str | None = None,
        architectures: list[str] | None = None,
        output_hidden_states: bool | None = False,
        return_dict: bool | None = True,
        dtype: str | None = None,
        torch_dtype: str | None = None,
        chunk_size_feed_forward: int = 0,
        problem_type: Literal[
            "regression", "single_label_classification", "multi_label_classification"
        ]
        | None = None,
        name_or_path: str = "",
        _commit_hash: str | None = None,
        attn_implementation: str | None = None,
        **kwargs: str | int | float | bool | None,
    ) -> None:
        """Initialize the composite checkpoint configuration."""
        dataset = normalize_dataset_name(dataset_name)
        mode = normalize_stage2_mode(stage2_mode)

        self.dataset_name = str(dataset)
        self.stage2_mode = str(mode)
        self.parser_model_name = parser_model_name
        self.parser_generation_max_length = parser_generation_max_length
        self.placement_generation_max_length = placement_generation_max_length
        self.temperature = temperature
        self.num_return_sequences = num_return_sequences
        self.canvas_size = tuple(canvas_size or canvas_size_for_dataset(dataset))
        self.parser_subfolder = parser_subfolder
        self.placement_subfolder = placement_subfolder

        label_map = id2label or id2label_for_dataset(dataset)
        normalized_id2label = {int(key): str(value) for key, value in label_map.items()}
        label2id = {value: key for key, value in normalized_id2label.items()}
        _ = kwargs.pop("id2label", None)
        _ = kwargs.pop("label2id", None)

        super().__init__(
            transformers_version=transformers_version,
            architectures=architectures,
            output_hidden_states=output_hidden_states,
            return_dict=return_dict,
            dtype=dtype or torch_dtype,
            chunk_size_feed_forward=chunk_size_feed_forward,
            is_encoder_decoder=is_encoder_decoder,
            id2label=normalized_id2label,
            label2id=label2id,
            problem_type=problem_type,
        )
        # Transformers v5 keeps model-specific token fields on the subclass;
        # only common configuration fields belong in the base call.
        self.pad_token_id = pad_token_id
        self.eos_token_id = eos_token_id
        self.decoder_start_token_id = decoder_start_token_id
        self.name_or_path = name_or_path
        self._commit_hash = _commit_hash
        self._attn_implementation = attn_implementation
        for key, value in kwargs.items():
            setattr(self, key, value)

__init__

__init__(
    dataset_name: str = "rico",
    stage2_mode: Stage2Mode | str = Stage2Mode.finetune,
    parser_model_name: str = "google/t5-v1_1-base",
    parser_generation_max_length: int = 600,
    placement_generation_max_length: int = 500,
    temperature: float = 0.7,
    num_return_sequences: int = 5,
    canvas_size: tuple[int, int] | list[int] | None = None,
    id2label: dict[int | str, str] | None = None,
    parser_subfolder: str = "semantic_parser",
    placement_subfolder: str = "placement",
    pad_token_id: int = 0,
    eos_token_id: int = 1,
    decoder_start_token_id: int = 0,
    is_encoder_decoder: bool = True,
    transformers_version: str | None = None,
    architectures: list[str] | None = None,
    output_hidden_states: bool | None = False,
    return_dict: bool | None = True,
    dtype: str | None = None,
    torch_dtype: str | None = None,
    chunk_size_feed_forward: int = 0,
    problem_type: Literal[
        "regression",
        "single_label_classification",
        "multi_label_classification",
    ]
    | None = None,
    name_or_path: str = "",
    _commit_hash: str | None = None,
    attn_implementation: str | None = None,
    **kwargs: str | int | float | bool | None,
) -> None

Initialize the composite checkpoint configuration.

Source code in models/parse-then-place/src/parse_then_place/configuration_parse_then_place.py
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
def __init__(
    self,
    dataset_name: str = "rico",
    stage2_mode: Stage2Mode | str = Stage2Mode.finetune,
    parser_model_name: str = "google/t5-v1_1-base",
    parser_generation_max_length: int = 600,
    placement_generation_max_length: int = 500,
    temperature: float = 0.7,
    num_return_sequences: int = 5,
    canvas_size: tuple[int, int] | list[int] | None = None,
    id2label: dict[int | str, str] | None = None,
    parser_subfolder: str = "semantic_parser",
    placement_subfolder: str = "placement",
    pad_token_id: int = 0,
    eos_token_id: int = 1,
    decoder_start_token_id: int = 0,
    is_encoder_decoder: bool = True,
    transformers_version: str | None = None,
    architectures: list[str] | None = None,
    output_hidden_states: bool | None = False,
    return_dict: bool | None = True,
    dtype: str | None = None,
    torch_dtype: str | None = None,
    chunk_size_feed_forward: int = 0,
    problem_type: Literal[
        "regression", "single_label_classification", "multi_label_classification"
    ]
    | None = None,
    name_or_path: str = "",
    _commit_hash: str | None = None,
    attn_implementation: str | None = None,
    **kwargs: str | int | float | bool | None,
) -> None:
    """Initialize the composite checkpoint configuration."""
    dataset = normalize_dataset_name(dataset_name)
    mode = normalize_stage2_mode(stage2_mode)

    self.dataset_name = str(dataset)
    self.stage2_mode = str(mode)
    self.parser_model_name = parser_model_name
    self.parser_generation_max_length = parser_generation_max_length
    self.placement_generation_max_length = placement_generation_max_length
    self.temperature = temperature
    self.num_return_sequences = num_return_sequences
    self.canvas_size = tuple(canvas_size or canvas_size_for_dataset(dataset))
    self.parser_subfolder = parser_subfolder
    self.placement_subfolder = placement_subfolder

    label_map = id2label or id2label_for_dataset(dataset)
    normalized_id2label = {int(key): str(value) for key, value in label_map.items()}
    label2id = {value: key for key, value in normalized_id2label.items()}
    _ = kwargs.pop("id2label", None)
    _ = kwargs.pop("label2id", None)

    super().__init__(
        transformers_version=transformers_version,
        architectures=architectures,
        output_hidden_states=output_hidden_states,
        return_dict=return_dict,
        dtype=dtype or torch_dtype,
        chunk_size_feed_forward=chunk_size_feed_forward,
        is_encoder_decoder=is_encoder_decoder,
        id2label=normalized_id2label,
        label2id=label2id,
        problem_type=problem_type,
    )
    # Transformers v5 keeps model-specific token fields on the subclass;
    # only common configuration fields belong in the base call.
    self.pad_token_id = pad_token_id
    self.eos_token_id = eos_token_id
    self.decoder_start_token_id = decoder_start_token_id
    self.name_or_path = name_or_path
    self._commit_hash = _commit_hash
    self._attn_implementation = attn_implementation
    for key, value in kwargs.items():
        setattr(self, key, value)

labels

Dataset metadata for Parse-Then-Place checkpoints.

ParseThenPlaceDatasetName

Bases: StrEnum

Datasets supported by the original Parse-Then-Place release.

Source code in models/parse-then-place/src/parse_then_place/labels.py
11
12
13
14
15
class ParseThenPlaceDatasetName(StrEnum):
    """Datasets supported by the original Parse-Then-Place release."""

    rico = auto()
    web = auto()

Stage2Mode

Bases: StrEnum

Released stage-2 checkpoint modes.

Source code in models/parse-then-place/src/parse_then_place/labels.py
18
19
20
21
22
class Stage2Mode(StrEnum):
    """Released stage-2 checkpoint modes."""

    pretrain = auto()
    finetune = auto()

DatasetMetadata

Bases: TypedDict

Static conversion metadata for one dataset.

Source code in models/parse-then-place/src/parse_then_place/labels.py
25
26
27
28
29
30
class DatasetMetadata(TypedDict):
    """Static conversion metadata for one dataset."""

    id2label: dict[int, str]
    canvas_size: tuple[int, int]
    max_elements: int

normalize_dataset_name

normalize_dataset_name(
    dataset_name: ParseThenPlaceDatasetName | str,
) -> ParseThenPlaceDatasetName

Normalize Parse-Then-Place dataset names.

Parameters:

Name Type Description Default
dataset_name ParseThenPlaceDatasetName | str

Dataset enum value or public/release string.

required

Returns:

Type Description
ParseThenPlaceDatasetName

Canonical Parse-Then-Place dataset name.

Raises:

Type Description
ValueError

If the dataset is unsupported.

Source code in models/parse-then-place/src/parse_then_place/labels.py
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
def normalize_dataset_name(
    dataset_name: ParseThenPlaceDatasetName | str,
) -> ParseThenPlaceDatasetName:
    """Normalize Parse-Then-Place dataset names.

    Args:
        dataset_name: Dataset enum value or public/release string.

    Returns:
        Canonical Parse-Then-Place dataset name.

    Raises:
        ValueError: If the dataset is unsupported.
    """
    if isinstance(dataset_name, ParseThenPlaceDatasetName):
        return dataset_name
    key = dataset_name.lower().replace("-", "_")
    if key == "webui":
        key = "web"
    try:
        return ParseThenPlaceDatasetName(key)
    except ValueError as exc:
        raise ValueError(
            f"Unsupported Parse-Then-Place dataset: {dataset_name}"
        ) from exc

normalize_stage2_mode

normalize_stage2_mode(
    stage2_mode: Stage2Mode | str,
) -> Stage2Mode

Normalize a released stage-2 checkpoint mode.

Source code in models/parse-then-place/src/parse_then_place/labels.py
 94
 95
 96
 97
 98
 99
100
101
def normalize_stage2_mode(stage2_mode: Stage2Mode | str) -> Stage2Mode:
    """Normalize a released stage-2 checkpoint mode."""
    if isinstance(stage2_mode, Stage2Mode):
        return stage2_mode
    try:
        return Stage2Mode(stage2_mode.lower().replace("-", "_"))
    except ValueError as exc:
        raise ValueError(f"Unsupported stage2_mode: {stage2_mode}") from exc

dataset_metadata

dataset_metadata(
    dataset_name: ParseThenPlaceDatasetName | str,
) -> DatasetMetadata

Return static metadata for a Parse-Then-Place dataset.

Source code in models/parse-then-place/src/parse_then_place/labels.py
104
105
106
107
108
def dataset_metadata(
    dataset_name: ParseThenPlaceDatasetName | str,
) -> DatasetMetadata:
    """Return static metadata for a Parse-Then-Place dataset."""
    return DATASET_METADATA[normalize_dataset_name(dataset_name)]

id2label_for_dataset

id2label_for_dataset(
    dataset_name: ParseThenPlaceDatasetName | str,
) -> dict[int, str]

Return the dataset-local integer-id label map.

Source code in models/parse-then-place/src/parse_then_place/labels.py
111
112
113
114
115
def id2label_for_dataset(
    dataset_name: ParseThenPlaceDatasetName | str,
) -> dict[int, str]:
    """Return the dataset-local integer-id label map."""
    return dict(dataset_metadata(dataset_name)["id2label"])

label2id_for_dataset

label2id_for_dataset(
    dataset_name: ParseThenPlaceDatasetName | str,
) -> dict[str, int]

Return lower-case label names mapped to dataset-local ids.

Source code in models/parse-then-place/src/parse_then_place/labels.py
118
119
120
121
122
123
124
def label2id_for_dataset(
    dataset_name: ParseThenPlaceDatasetName | str,
) -> dict[str, int]:
    """Return lower-case label names mapped to dataset-local ids."""
    return {
        label.lower(): idx for idx, label in id2label_for_dataset(dataset_name).items()
    }

canvas_size_for_dataset

canvas_size_for_dataset(
    dataset_name: ParseThenPlaceDatasetName | str,
) -> tuple[int, int]

Return the dataset canvas size as (width, height).

Source code in models/parse-then-place/src/parse_then_place/labels.py
127
128
129
130
131
def canvas_size_for_dataset(
    dataset_name: ParseThenPlaceDatasetName | str,
) -> tuple[int, int]:
    """Return the dataset canvas size as ``(width, height)``."""
    return dataset_metadata(dataset_name)["canvas_size"]

pipeline_parse_then_place

Pipeline wrapper for Parse-Then-Place composite checkpoints.

ParseThenPlacePipeline

Bases: LayoutGenerationPipeline

Compose standard seq2seq parser and placement models.

Source code in models/parse-then-place/src/parse_then_place/pipeline_parse_then_place.py
 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
class ParseThenPlacePipeline(LayoutGenerationPipeline):
    """Compose standard seq2seq parser and placement models."""

    config_class: ClassVar[type[PretrainedConfig]] = ParseThenPlaceConfig
    component_specs: ClassVar[dict[str, PipelineComponentSpec]] = {
        "parser": PipelineComponentSpec(
            attribute_name="parser",
            loader=_load_seq2seq_component,
            config_subfolder_attribute="parser_subfolder",
            required=False,
        ),
        "placement": PipelineComponentSpec(
            attribute_name="placement",
            loader=_load_seq2seq_component,
            config_subfolder_attribute="placement_subfolder",
        ),
        "processor": PipelineComponentSpec(
            attribute_name="processor",
            loader=_load_processor_component,
            marker_file="processor_config.json",
            save_with_is_main_process=False,
        ),
    }

    config: ParseThenPlaceConfig
    parser: PreTrainedModel | None
    placement: PreTrainedModel | None
    processor: ParseThenPlaceProcessor

    def __init__(
        self,
        config: ParseThenPlaceConfig,
        processor: ParseThenPlaceProcessor,
        *,
        parser: PreTrainedModel | None = None,
        placement: PreTrainedModel | None = None,
    ) -> None:
        """Initialize the composite pipeline."""
        super().__init__(config)
        self.config = config
        self.processor = processor
        self.parser = parser
        self.placement = placement

    @classmethod
    def from_pretrained(
        cls,
        pretrained_model_name_or_path: str | Path,
        *,
        parser: PreTrainedModel | None = None,
        placement: PreTrainedModel | None = None,
        processor: ParseThenPlaceProcessor | None = None,
        local_files_only: bool = False,
        config: ParseThenPlaceConfig | PretrainedConfig | None = None,
    ) -> ParseThenPlacePipeline:  # ty: ignore[invalid-method-override]
        """Load a composite pipeline from a root directory."""
        components: dict[str, PreTrainedModel | ParseThenPlaceProcessor] = {}
        if parser is not None:
            components["parser"] = parser
        if placement is not None:
            components["placement"] = placement
        if processor is None:
            loaded = super().from_pretrained(
                pretrained_model_name_or_path,
                local_files_only=local_files_only,
                config=config,
                components=components,
            )
            return cast(ParseThenPlacePipeline, loaded)
        components["processor"] = processor
        loaded = super().from_pretrained(
            pretrained_model_name_or_path,
            local_files_only=local_files_only,
            config=config,
            components=components,
        )
        return cast(ParseThenPlacePipeline, loaded)

    @classmethod
    def _from_pretrained_components(
        cls,
        *,
        config: PretrainedConfig,
        components: Mapping[str, PreTrainedModel | ParseThenPlaceProcessor | None],
    ) -> ParseThenPlacePipeline:
        """Build a pipeline from loaded config and components."""
        return cls(
            config=cast(ParseThenPlaceConfig, config),
            processor=cast(ParseThenPlaceProcessor, components["processor"]),
            parser=cast(PreTrainedModel | None, components.get("parser")),
            placement=cast(PreTrainedModel | None, components["placement"]),
        )

    @torch.no_grad()
    def parse(
        self,
        input_ids: Int[torch.Tensor, "batch tokens"],
        attention_mask: Bool[torch.Tensor, "batch tokens"] | None = None,
        *,
        generation_max_length: int | None = None,
        **generate_kwargs: str | int | float | bool | torch.Generator | None,
    ) -> Int[torch.Tensor, "batch tokens"]:
        """Generate logical-form token ids with the parser stage."""
        if self.parser is None:
            raise ValueError("Parser stage is not loaded")

        generated = cast(_GenerationModel, self.parser).generate(
            input_ids=input_ids,
            attention_mask=attention_mask,
            max_length=generation_max_length
            or self.config.parser_generation_max_length,
            **generate_kwargs,
        )
        return generated

    @torch.no_grad()
    def place(
        self,
        input_ids: Int[torch.Tensor, "batch tokens"],
        attention_mask: Bool[torch.Tensor, "batch tokens"] | None = None,
        *,
        generation_max_length: int | None = None,
        num_return_sequences: int | None = None,
        temperature: float | None = None,
        do_sample: bool = True,
        generator: torch.Generator | None = None,
        **generate_kwargs: str | float | bool | torch.Generator | None,
    ) -> Int[torch.Tensor, "batch tokens"]:
        """Generate layout token ids with the placement stage."""
        if self.placement is None:
            raise ValueError("Placement stage is not loaded")

        if generator is not None:
            generate_kwargs["generator"] = generator
        generated = cast(_GenerationModel, self.placement).generate(
            input_ids=input_ids,
            attention_mask=attention_mask,
            max_length=generation_max_length
            or self.config.placement_generation_max_length,
            num_return_sequences=num_return_sequences
            or self.config.num_return_sequences,
            temperature=temperature or self.config.temperature,
            do_sample=do_sample,
            **generate_kwargs,
        )
        return generated

    def __call__(
        self,
        *,
        prompt: str | Sequence[str] | None = None,
        batch_size: int = 1,
        seed: int | None = None,
        generator: torch.Generator | None = None,
        condition_type: ConditionType | str = ConditionType.text,
        labels: Int[torch.Tensor, "batch elements"]
        | list[ArrayLikeInput]
        | None = None,
        bbox: Float[torch.Tensor, "batch elements 4"]
        | list[ArrayLikeInput]
        | None = None,
        mask: Bool[torch.Tensor, "batch elements"] | list[ArrayLikeInput] | 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,
        num_return_sequences: int | None = None,
        temperature: float | None = None,
        output_candidate: Literal["first", "all", "best"] = "first",
        output_type: Literal["dataclass", "dict"] = "dataclass",
        return_intermediates: bool = False,
        layout_text: str | list[str] | list[list[str]] | None = None,
    ) -> LayoutGenerationOutput | ParseThenPlaceOutputDict:  # ty: ignore[invalid-method-override]
        """Generate a layout from natural-language text."""
        _ = (
            batch_size,
            labels,
            bbox,
            mask,
            num_elements,
            box_format,
            normalized,
            canvas_size,
            num_inference_steps,
        )
        condition = normalize_condition_type(condition_type)
        if condition is not ConditionType.text:
            raise NotImplementedError(
                "Parse-Then-Place only supports condition_type='text'"
            )

        if layout_text is not None:
            layout_items = (
                [layout_text] if isinstance(layout_text, str) else layout_text
            )
            return self.processor.layout_text_to_output(
                layout_items,
                output_candidate=output_candidate,
                output_type=output_type,
                return_intermediates=return_intermediates,
            )
        if prompt is None:
            raise ValueError("prompt is required for Parse-Then-Place generation")

        generation_generator = self.prepare_generator(
            generator=generator,
            seed=seed,
        )
        prompts = [prompt] if isinstance(prompt, str) else list(prompt)
        parser_inputs = self.processor(prompts)
        if "input_ids" not in parser_inputs:
            raise ValueError("processor requires parser_tokenizer for model inference")

        parser_ids = self.parse(
            parser_inputs["input_ids"],
            attention_mask=parser_inputs.get("attention_mask"),
            generation_max_length=self.config.parser_generation_max_length,
        )
        value_maps = cast(list[dict[str, str] | None], parser_inputs.get("value_maps"))
        logical_forms = self.processor.postprocess_ir(
            parser_ids,
            value_maps=value_maps,
        )
        placement_inputs = self.processor.ir_to_placement_inputs(logical_forms)
        placement_encoded = self.processor.encode_placement_inputs(placement_inputs)
        if "input_ids" not in placement_encoded:
            raise ValueError(
                "processor requires placement_tokenizer for model inference"
            )

        return_sequences = num_return_sequences or self.config.num_return_sequences
        placement_ids = self.place(
            placement_encoded["input_ids"],
            attention_mask=placement_encoded.get("attention_mask"),
            num_return_sequences=return_sequences,
            temperature=temperature,
            generator=generation_generator,
        )
        grouped = self.processor.decode_layout_sequences(
            placement_ids,
            batch_size=len(prompts),
            num_return_sequences=return_sequences,
        )
        output = self.processor.layout_text_to_output(
            grouped,
            output_candidate=output_candidate,
            output_type="dataclass",
            return_intermediates=True,
        )
        if isinstance(output, LayoutGenerationOutput):
            intermediates = (
                dict(output.intermediates)
                if isinstance(output.intermediates, dict)
                else {}
            )
            if return_intermediates:
                intermediates.update(
                    {
                        "prompt": prompts,
                        "logical_forms": logical_forms,
                        "placement_inputs": placement_inputs,
                    }
                )
            output.intermediates = intermediates if return_intermediates else None
        if output_type == "dict":
            return dict(output)
        return output

    generate = __call__

__init__

__init__(
    config: ParseThenPlaceConfig,
    processor: ParseThenPlaceProcessor,
    *,
    parser: PreTrainedModel | None = None,
    placement: PreTrainedModel | None = None,
) -> None

Initialize the composite pipeline.

Source code in models/parse-then-place/src/parse_then_place/pipeline_parse_then_place.py
111
112
113
114
115
116
117
118
119
120
121
122
123
124
def __init__(
    self,
    config: ParseThenPlaceConfig,
    processor: ParseThenPlaceProcessor,
    *,
    parser: PreTrainedModel | None = None,
    placement: PreTrainedModel | None = None,
) -> None:
    """Initialize the composite pipeline."""
    super().__init__(config)
    self.config = config
    self.processor = processor
    self.parser = parser
    self.placement = placement

from_pretrained classmethod

from_pretrained(
    pretrained_model_name_or_path: str | Path,
    *,
    parser: PreTrainedModel | None = None,
    placement: PreTrainedModel | None = None,
    processor: ParseThenPlaceProcessor | None = None,
    local_files_only: bool = False,
    config: ParseThenPlaceConfig
    | PretrainedConfig
    | None = None,
) -> ParseThenPlacePipeline

Load a composite pipeline from a root directory.

Source code in models/parse-then-place/src/parse_then_place/pipeline_parse_then_place.py
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
@classmethod
def from_pretrained(
    cls,
    pretrained_model_name_or_path: str | Path,
    *,
    parser: PreTrainedModel | None = None,
    placement: PreTrainedModel | None = None,
    processor: ParseThenPlaceProcessor | None = None,
    local_files_only: bool = False,
    config: ParseThenPlaceConfig | PretrainedConfig | None = None,
) -> ParseThenPlacePipeline:  # ty: ignore[invalid-method-override]
    """Load a composite pipeline from a root directory."""
    components: dict[str, PreTrainedModel | ParseThenPlaceProcessor] = {}
    if parser is not None:
        components["parser"] = parser
    if placement is not None:
        components["placement"] = placement
    if processor is None:
        loaded = super().from_pretrained(
            pretrained_model_name_or_path,
            local_files_only=local_files_only,
            config=config,
            components=components,
        )
        return cast(ParseThenPlacePipeline, loaded)
    components["processor"] = processor
    loaded = super().from_pretrained(
        pretrained_model_name_or_path,
        local_files_only=local_files_only,
        config=config,
        components=components,
    )
    return cast(ParseThenPlacePipeline, loaded)

parse

parse(
    input_ids: Int[Tensor, "batch tokens"],
    attention_mask: Bool[Tensor, "batch tokens"]
    | None = None,
    *,
    generation_max_length: int | None = None,
    **generate_kwargs: str
    | int
    | float
    | bool
    | Generator
    | None,
) -> Int[torch.Tensor, "batch tokens"]

Generate logical-form token ids with the parser stage.

Source code in models/parse-then-place/src/parse_then_place/pipeline_parse_then_place.py
175
176
177
178
179
180
181
182
183
184
185
186
187
188
189
190
191
192
193
194
195
@torch.no_grad()
def parse(
    self,
    input_ids: Int[torch.Tensor, "batch tokens"],
    attention_mask: Bool[torch.Tensor, "batch tokens"] | None = None,
    *,
    generation_max_length: int | None = None,
    **generate_kwargs: str | int | float | bool | torch.Generator | None,
) -> Int[torch.Tensor, "batch tokens"]:
    """Generate logical-form token ids with the parser stage."""
    if self.parser is None:
        raise ValueError("Parser stage is not loaded")

    generated = cast(_GenerationModel, self.parser).generate(
        input_ids=input_ids,
        attention_mask=attention_mask,
        max_length=generation_max_length
        or self.config.parser_generation_max_length,
        **generate_kwargs,
    )
    return generated

place

place(
    input_ids: Int[Tensor, "batch tokens"],
    attention_mask: Bool[Tensor, "batch tokens"]
    | None = None,
    *,
    generation_max_length: int | None = None,
    num_return_sequences: int | None = None,
    temperature: float | None = None,
    do_sample: bool = True,
    generator: Generator | None = None,
    **generate_kwargs: str
    | float
    | bool
    | Generator
    | None,
) -> Int[torch.Tensor, "batch tokens"]

Generate layout token ids with the placement stage.

Source code in models/parse-then-place/src/parse_then_place/pipeline_parse_then_place.py
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
@torch.no_grad()
def place(
    self,
    input_ids: Int[torch.Tensor, "batch tokens"],
    attention_mask: Bool[torch.Tensor, "batch tokens"] | None = None,
    *,
    generation_max_length: int | None = None,
    num_return_sequences: int | None = None,
    temperature: float | None = None,
    do_sample: bool = True,
    generator: torch.Generator | None = None,
    **generate_kwargs: str | float | bool | torch.Generator | None,
) -> Int[torch.Tensor, "batch tokens"]:
    """Generate layout token ids with the placement stage."""
    if self.placement is None:
        raise ValueError("Placement stage is not loaded")

    if generator is not None:
        generate_kwargs["generator"] = generator
    generated = cast(_GenerationModel, self.placement).generate(
        input_ids=input_ids,
        attention_mask=attention_mask,
        max_length=generation_max_length
        or self.config.placement_generation_max_length,
        num_return_sequences=num_return_sequences
        or self.config.num_return_sequences,
        temperature=temperature or self.config.temperature,
        do_sample=do_sample,
        **generate_kwargs,
    )
    return generated

__call__

__call__(
    *,
    prompt: str | Sequence[str] | None = None,
    batch_size: int = 1,
    seed: int | None = None,
    generator: Generator | None = None,
    condition_type: ConditionType
    | str = ConditionType.text,
    labels: Int[Tensor, "batch elements"]
    | list[ArrayLikeInput]
    | None = None,
    bbox: Float[Tensor, "batch elements 4"]
    | list[ArrayLikeInput]
    | None = None,
    mask: Bool[Tensor, "batch elements"]
    | list[ArrayLikeInput]
    | 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,
    num_return_sequences: int | None = None,
    temperature: float | None = None,
    output_candidate: Literal[
        "first", "all", "best"
    ] = "first",
    output_type: Literal["dataclass", "dict"] = "dataclass",
    return_intermediates: bool = False,
    layout_text: str
    | list[str]
    | list[list[str]]
    | None = None,
) -> LayoutGenerationOutput | ParseThenPlaceOutputDict

Generate a layout from natural-language text.

Source code in models/parse-then-place/src/parse_then_place/pipeline_parse_then_place.py
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
def __call__(
    self,
    *,
    prompt: str | Sequence[str] | None = None,
    batch_size: int = 1,
    seed: int | None = None,
    generator: torch.Generator | None = None,
    condition_type: ConditionType | str = ConditionType.text,
    labels: Int[torch.Tensor, "batch elements"]
    | list[ArrayLikeInput]
    | None = None,
    bbox: Float[torch.Tensor, "batch elements 4"]
    | list[ArrayLikeInput]
    | None = None,
    mask: Bool[torch.Tensor, "batch elements"] | list[ArrayLikeInput] | 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,
    num_return_sequences: int | None = None,
    temperature: float | None = None,
    output_candidate: Literal["first", "all", "best"] = "first",
    output_type: Literal["dataclass", "dict"] = "dataclass",
    return_intermediates: bool = False,
    layout_text: str | list[str] | list[list[str]] | None = None,
) -> LayoutGenerationOutput | ParseThenPlaceOutputDict:  # ty: ignore[invalid-method-override]
    """Generate a layout from natural-language text."""
    _ = (
        batch_size,
        labels,
        bbox,
        mask,
        num_elements,
        box_format,
        normalized,
        canvas_size,
        num_inference_steps,
    )
    condition = normalize_condition_type(condition_type)
    if condition is not ConditionType.text:
        raise NotImplementedError(
            "Parse-Then-Place only supports condition_type='text'"
        )

    if layout_text is not None:
        layout_items = (
            [layout_text] if isinstance(layout_text, str) else layout_text
        )
        return self.processor.layout_text_to_output(
            layout_items,
            output_candidate=output_candidate,
            output_type=output_type,
            return_intermediates=return_intermediates,
        )
    if prompt is None:
        raise ValueError("prompt is required for Parse-Then-Place generation")

    generation_generator = self.prepare_generator(
        generator=generator,
        seed=seed,
    )
    prompts = [prompt] if isinstance(prompt, str) else list(prompt)
    parser_inputs = self.processor(prompts)
    if "input_ids" not in parser_inputs:
        raise ValueError("processor requires parser_tokenizer for model inference")

    parser_ids = self.parse(
        parser_inputs["input_ids"],
        attention_mask=parser_inputs.get("attention_mask"),
        generation_max_length=self.config.parser_generation_max_length,
    )
    value_maps = cast(list[dict[str, str] | None], parser_inputs.get("value_maps"))
    logical_forms = self.processor.postprocess_ir(
        parser_ids,
        value_maps=value_maps,
    )
    placement_inputs = self.processor.ir_to_placement_inputs(logical_forms)
    placement_encoded = self.processor.encode_placement_inputs(placement_inputs)
    if "input_ids" not in placement_encoded:
        raise ValueError(
            "processor requires placement_tokenizer for model inference"
        )

    return_sequences = num_return_sequences or self.config.num_return_sequences
    placement_ids = self.place(
        placement_encoded["input_ids"],
        attention_mask=placement_encoded.get("attention_mask"),
        num_return_sequences=return_sequences,
        temperature=temperature,
        generator=generation_generator,
    )
    grouped = self.processor.decode_layout_sequences(
        placement_ids,
        batch_size=len(prompts),
        num_return_sequences=return_sequences,
    )
    output = self.processor.layout_text_to_output(
        grouped,
        output_candidate=output_candidate,
        output_type="dataclass",
        return_intermediates=True,
    )
    if isinstance(output, LayoutGenerationOutput):
        intermediates = (
            dict(output.intermediates)
            if isinstance(output.intermediates, dict)
            else {}
        )
        if return_intermediates:
            intermediates.update(
                {
                    "prompt": prompts,
                    "logical_forms": logical_forms,
                    "placement_inputs": placement_inputs,
                }
            )
        output.intermediates = intermediates if return_intermediates else None
    if output_type == "dict":
        return dict(output)
    return output

processing_parse_then_place

Processor for Parse-Then-Place text, IR, and generated layout strings.

PromptEncoding

Bases: TypedDict

Preprocessed prompt and value placeholders.

Source code in models/parse-then-place/src/parse_then_place/processing_parse_then_place.py
24
25
26
27
28
class PromptEncoding(TypedDict):
    """Preprocessed prompt and value placeholders."""

    prompt: str
    value_map: dict[str, str] | None

ParsedElement

Bases: TypedDict

One parsed generated layout element.

Source code in models/parse-then-place/src/parse_then_place/processing_parse_then_place.py
31
32
33
34
35
class ParsedElement(TypedDict):
    """One parsed generated layout element."""

    label: int
    bbox: list[float]

TensorLike

Bases: Protocol

Runtime-safe protocol for tensor-like processor dictionary values.

Source code in models/parse-then-place/src/parse_then_place/processing_parse_then_place.py
46
47
48
49
50
51
52
53
@runtime_checkable
class TensorLike(Protocol):
    """Runtime-safe protocol for tensor-like processor dictionary values."""

    @property
    def shape(self) -> tuple[int, ...]:
        """Return tensor-like dimensions."""
        ...

shape property

shape: tuple[int, ...]

Return tensor-like dimensions.

ParseThenPlaceProcessor

Bases: ProcessorMixin

Build stage inputs and parse placement-model output text.

Source code in models/parse-then-place/src/parse_then_place/processing_parse_then_place.py
 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
class ParseThenPlaceProcessor(ProcessorMixin):
    """Build stage inputs and parse placement-model output text."""

    attributes = ["parser_tokenizer", "placement_tokenizer"]
    parser_tokenizer_class = "AutoTokenizer"
    placement_tokenizer_class = "T5Tokenizer"

    def __init__(
        self,
        parser_tokenizer: PreTrainedTokenizerBase | None = None,
        placement_tokenizer: PreTrainedTokenizerBase | None = None,
        dataset_name: ParseThenPlaceDatasetName | str = ParseThenPlaceDatasetName.rico,
        canvas_size: tuple[int, int] | None = None,
        id2label: dict[int, str] | None = None,
    ) -> None:
        """Initialize tokenizer handles and dataset metadata."""
        dataset = normalize_dataset_name(dataset_name)
        self.parser_tokenizer = parser_tokenizer
        self.placement_tokenizer = placement_tokenizer
        self.dataset_name = str(dataset)
        self.canvas_size = canvas_size or canvas_size_for_dataset(dataset)
        self.id2label = (
            {int(key): str(value) for key, value in id2label.items()}
            if id2label is not None
            else id2label_for_dataset(dataset)
        )
        self.label2id = {label.lower(): idx for idx, label in self.id2label.items()}
        # Keep released spellings as aliases because RICO uses lower-case labels.
        self.label2id.update(label2id_for_dataset(dataset))
        if parser_tokenizer is not None and placement_tokenizer is not None:
            super().__init__(
                parser_tokenizer=parser_tokenizer,
                placement_tokenizer=placement_tokenizer,
            )

    @classmethod
    def from_config(
        cls,
        dataset_name: ParseThenPlaceDatasetName | str = ParseThenPlaceDatasetName.rico,
        *,
        canvas_size: tuple[int, int] | None = None,
        id2label: dict[int, str] | None = None,
    ) -> ParseThenPlaceProcessor:
        """Construct metadata-only processor for tests and local smoke checks."""
        return cls(
            parser_tokenizer=None,
            placement_tokenizer=None,
            dataset_name=dataset_name,
            canvas_size=canvas_size,
            id2label=id2label,
        )

    def preprocess_prompt(
        self,
        prompt: str,
        *,
        replace_explicit_value: bool = True,
    ) -> PromptEncoding:
        """Apply the released text normalization used before stage-1 parsing.

        Args:
            prompt: Natural-language text prompt.
            replace_explicit_value: Whether quoted values should be replaced by
                deterministic ``value_N`` placeholders.

        Returns:
            Normalized prompt and the placeholder recovery map.
        """
        result = prompt.replace("#", "").strip().lower()
        result = (
            result.replace("“", '"')
            .replace("”", '"')
            .replace("‘", "'")
            .replace("’", "'")
        )
        result = _WHITESPACE_RE.sub(" ", result)
        if not replace_explicit_value:
            return {"prompt": result, "value_map": None}
        return self._extract_explicit_values(result)

    def _extract_explicit_values(self, prompt: str) -> PromptEncoding:
        result = re.sub(r"(\w)'s\s+", r"\g<1>`s ", prompt)
        values = _DOUBLE_QUOTED_RE.findall(result)
        single_values = _SINGLE_QUOTED_RE.findall(result)
        if len(single_values) == 1 and any(
            punct in single_values[0] for punct in (",", ".")
        ):
            single_values = []
        values.extend(single_values)
        value_map: dict[str, str] = {}
        for value_idx, value in enumerate(values):
            placeholder = f"value_{value_idx}"
            value_map[placeholder] = value.strip('"').strip("'").strip()
            result = result.replace(value, f'"{placeholder}"', 1)
        result = re.sub(r"(\w)`s\s+", r"\g<1>'s ", result)
        result = _WHITESPACE_RE.sub(" ", result)
        return {"prompt": result, "value_map": value_map}

    def __call__(
        self,
        prompt: str | Sequence[str],
        *,
        replace_explicit_value: bool = True,
        return_tensors: Literal["pt"] = "pt",
    ) -> BatchEncoding:
        """Tokenize prompt text for the semantic parser stage."""
        prompts = [prompt] if isinstance(prompt, str) else list(prompt)
        encodings = [
            self.preprocess_prompt(item, replace_explicit_value=replace_explicit_value)
            for item in prompts
        ]
        texts = [item["prompt"] for item in encodings]
        if self.parser_tokenizer is None:
            return BatchEncoding(
                {
                    "prompt_text": texts,
                    "value_maps": [item["value_map"] for item in encodings],
                }
            )
        parser_tokenizer = self.parser_tokenizer
        tokenized = parser_tokenizer(texts, return_tensors=return_tensors, padding=True)
        tokenized["value_maps"] = [item["value_map"] for item in encodings]
        tokenized["prompt_text"] = texts
        return cast(BatchEncoding, tokenized)

    def postprocess_ir(
        self,
        generated_ids: Int[torch.Tensor, "batch tokens"] | Sequence[str],
        *,
        value_maps: list[dict[str, str] | None] | None = None,
    ) -> list[str]:
        """Decode and lightly normalize stage-1 logical forms."""
        if isinstance(generated_ids, torch.Tensor):
            if self.parser_tokenizer is None:
                raise ValueError("parser_tokenizer is required to decode generated ids")

            parser_tokenizer = self.parser_tokenizer
            logical_forms = parser_tokenizer.batch_decode(
                generated_ids, skip_special_tokens=True
            )
        else:
            logical_forms = list(generated_ids)
        if value_maps is None:
            return [self._normalize_ir(item) for item in logical_forms]
        return [
            self._recover_ir_values(self._normalize_ir(item), value_map)
            for item, value_map in zip(logical_forms, value_maps, strict=True)
        ]

    def _normalize_ir(self, logical_form: str) -> str:
        result = logical_form.replace("[", " [ ").replace("]", " ] ").strip().lower()
        return _WHITESPACE_RE.sub(" ", result)

    def _recover_ir_values(
        self,
        logical_form: str,
        value_map: dict[str, str] | None,
    ) -> str:
        if not value_map:
            return logical_form
        result = logical_form
        for placeholder, value in value_map.items():
            recovered = value.replace("'", "")
            result = result.replace(f"'{placeholder}'", f"'{recovered}'")
            result = result.replace(f" {placeholder},", f" {value},")
            result = result.replace(f" {placeholder}'", f" {value}'")
        return result.replace("&", " and ")

    def ir_to_placement_inputs(self, logical_forms: Sequence[str]) -> list[str]:
        """Convert logical forms to placement-constraint strings.

        Runtime keeps this method deterministic and accepts already-linearized
        constraints, which is also the artifact stored in stage-1 prediction JSON
        files. The current parity scripts do not execute the released grammar
        executor.
        """
        return [self._logical_form_to_constraint(item) for item in logical_forms]

    def _logical_form_to_constraint(self, logical_form: str) -> str:
        text = _WHITESPACE_RE.sub(" ", logical_form.strip())
        if ":" in text and "|" in text:
            return text
        return text

    def encode_placement_inputs(
        self,
        placement_inputs: Sequence[str],
        *,
        return_tensors: Literal["pt"] = "pt",
    ) -> BatchEncoding:
        """Tokenize stage-2 placement constraints."""
        if self.placement_tokenizer is None:
            return BatchEncoding({"placement_text": list(placement_inputs)})
        placement_tokenizer = self.placement_tokenizer
        tokenized = placement_tokenizer(
            list(placement_inputs), return_tensors=return_tensors, padding=True
        )
        tokenized["placement_text"] = list(placement_inputs)
        return cast(BatchEncoding, tokenized)

    def decode_layout_sequences(
        self,
        generated_ids: Int[torch.Tensor, "batch tokens"] | Sequence[str],
        *,
        batch_size: int,
        num_return_sequences: int,
    ) -> list[list[str]]:
        """Decode stage-2 generated ids into grouped layout strings."""
        if isinstance(generated_ids, torch.Tensor):
            if self.placement_tokenizer is None:
                raise ValueError(
                    "placement_tokenizer is required to decode generated ids"
                )

            placement_tokenizer = self.placement_tokenizer
            flat = placement_tokenizer.batch_decode(
                generated_ids, skip_special_tokens=True
            )
        else:
            flat = list(generated_ids)
        expected = batch_size * num_return_sequences
        if len(flat) != expected:
            raise ValueError(
                "Generated layout count does not match batch_size * num_return_sequences: "
                f"{len(flat)} != {expected}"
            )

        return [
            flat[idx * num_return_sequences : (idx + 1) * num_return_sequences]
            for idx in range(batch_size)
        ]

    def layout_text_to_output(
        self,
        layout_text: Sequence[str] | Sequence[Sequence[str]],
        *,
        output_candidate: Literal["first", "all", "best"] = "first",
        output_type: Literal["dataclass", "dict"] = "dataclass",
        return_intermediates: bool = False,
    ) -> LayoutGenerationOutput | ParseThenPlaceOutputDict:
        """Parse generated ``label left top width height`` text into schema."""
        candidate_groups = self._normalize_layout_text_groups(layout_text)
        selected = self._select_candidates(candidate_groups, output_candidate)
        parsed_groups = [self._parse_layout_text(item) for item in selected]
        max_len = max((len(item) for item in parsed_groups), default=0) or 1

        bbox_rows: list[Float[torch.Tensor, "elements 4"]] = []
        label_rows: list[Int[torch.Tensor, ...]] = []
        mask_rows: list[Bool[torch.Tensor, ...]] = []

        for parsed in parsed_groups:
            labels = torch.tensor([item["label"] for item in parsed], dtype=torch.long)
            boxes = torch.tensor([item["bbox"] for item in parsed], dtype=torch.float32)
            mask = torch.ones(len(parsed), dtype=torch.bool)

            if len(parsed) == 0:
                labels = torch.zeros(max_len, dtype=torch.long)
                boxes = torch.zeros(max_len, 4, dtype=torch.float32)
                mask = torch.zeros(max_len, dtype=torch.bool)
            elif len(parsed) < max_len:
                pad = max_len - len(parsed)
                labels = torch.nn.functional.pad(labels, (0, pad))
                boxes = torch.nn.functional.pad(boxes, (0, 0, 0, pad))
                mask = torch.nn.functional.pad(mask, (0, pad))

            label_rows.append(labels)
            bbox_rows.append(boxes)
            mask_rows.append(mask)

        raw_bbox = torch.stack(bbox_rows)
        bbox = normalize_boxes(
            raw_bbox,
            canvas_size=self.canvas_size,
            box_format="ltwh",
        )
        output = LayoutGenerationOutput(
            bbox=bbox.float(),
            labels=torch.stack(label_rows).long(),
            mask=torch.stack(mask_rows).bool(),
            id2label=dict(self.id2label),
            intermediates={
                "layout_text": selected,
                "layout_text_candidates": candidate_groups,
                "dataset_name": self.dataset_name,
                "canvas_size": self.canvas_size,
            }
            if return_intermediates
            else None,
        )
        if output_type == "dict":
            return dict(output)
        if output_type != "dataclass":
            raise ValueError(f"Unsupported output_type: {output_type}")

        return output

    def _normalize_layout_text_groups(
        self,
        layout_text: Sequence[str] | Sequence[Sequence[str]],
    ) -> list[list[str]]:
        if not layout_text:
            return []
        first = layout_text[0]
        if isinstance(first, str):
            return [[item] for item in cast(Sequence[str], layout_text)]
        return [list(item) for item in cast(Sequence[Sequence[str]], layout_text)]

    def _select_candidates(
        self,
        candidate_groups: list[list[str]],
        output_candidate: Literal["first", "all", "best"],
    ) -> list[str]:
        if output_candidate == "first":
            return [group[0] if group else "" for group in candidate_groups]
        if output_candidate == "best":
            return [
                max(group, key=lambda item: len(self._parse_layout_text(item)))
                if group
                else ""
                for group in candidate_groups
            ]
        if output_candidate == "all":
            return ["\n".join(group) for group in candidate_groups]
        raise ValueError(f"Unsupported output_candidate: {output_candidate}")

    def _parse_layout_text(self, layout_text: str) -> list[ParsedElement]:
        elements: list[ParsedElement] = []
        for match in _LAYOUT_PATTERN.finditer(layout_text.lower()):
            label = _WHITESPACE_RE.sub(" ", match.group("label")).strip()
            label_id = self.label2id.get(label)
            if label_id is None:
                continue
            elements.append(
                {
                    "label": label_id,
                    "bbox": [
                        float(match.group("left")),
                        float(match.group("top")),
                        float(match.group("width")),
                        float(match.group("height")),
                    ],
                }
            )
        return elements

__init__

__init__(
    parser_tokenizer: PreTrainedTokenizerBase | None = None,
    placement_tokenizer: PreTrainedTokenizerBase
    | None = None,
    dataset_name: ParseThenPlaceDatasetName
    | str = ParseThenPlaceDatasetName.rico,
    canvas_size: tuple[int, int] | None = None,
    id2label: dict[int, str] | None = None,
) -> None

Initialize tokenizer handles and dataset metadata.

Source code in models/parse-then-place/src/parse_then_place/processing_parse_then_place.py
 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
def __init__(
    self,
    parser_tokenizer: PreTrainedTokenizerBase | None = None,
    placement_tokenizer: PreTrainedTokenizerBase | None = None,
    dataset_name: ParseThenPlaceDatasetName | str = ParseThenPlaceDatasetName.rico,
    canvas_size: tuple[int, int] | None = None,
    id2label: dict[int, str] | None = None,
) -> None:
    """Initialize tokenizer handles and dataset metadata."""
    dataset = normalize_dataset_name(dataset_name)
    self.parser_tokenizer = parser_tokenizer
    self.placement_tokenizer = placement_tokenizer
    self.dataset_name = str(dataset)
    self.canvas_size = canvas_size or canvas_size_for_dataset(dataset)
    self.id2label = (
        {int(key): str(value) for key, value in id2label.items()}
        if id2label is not None
        else id2label_for_dataset(dataset)
    )
    self.label2id = {label.lower(): idx for idx, label in self.id2label.items()}
    # Keep released spellings as aliases because RICO uses lower-case labels.
    self.label2id.update(label2id_for_dataset(dataset))
    if parser_tokenizer is not None and placement_tokenizer is not None:
        super().__init__(
            parser_tokenizer=parser_tokenizer,
            placement_tokenizer=placement_tokenizer,
        )

from_config classmethod

from_config(
    dataset_name: ParseThenPlaceDatasetName
    | str = ParseThenPlaceDatasetName.rico,
    *,
    canvas_size: tuple[int, int] | None = None,
    id2label: dict[int, str] | None = None,
) -> ParseThenPlaceProcessor

Construct metadata-only processor for tests and local smoke checks.

Source code in models/parse-then-place/src/parse_then_place/processing_parse_then_place.py
104
105
106
107
108
109
110
111
112
113
114
115
116
117
118
119
@classmethod
def from_config(
    cls,
    dataset_name: ParseThenPlaceDatasetName | str = ParseThenPlaceDatasetName.rico,
    *,
    canvas_size: tuple[int, int] | None = None,
    id2label: dict[int, str] | None = None,
) -> ParseThenPlaceProcessor:
    """Construct metadata-only processor for tests and local smoke checks."""
    return cls(
        parser_tokenizer=None,
        placement_tokenizer=None,
        dataset_name=dataset_name,
        canvas_size=canvas_size,
        id2label=id2label,
    )

preprocess_prompt

preprocess_prompt(
    prompt: str, *, replace_explicit_value: bool = True
) -> PromptEncoding

Apply the released text normalization used before stage-1 parsing.

Parameters:

Name Type Description Default
prompt str

Natural-language text prompt.

required
replace_explicit_value bool

Whether quoted values should be replaced by deterministic value_N placeholders.

True

Returns:

Type Description
PromptEncoding

Normalized prompt and the placeholder recovery map.

Source code in models/parse-then-place/src/parse_then_place/processing_parse_then_place.py
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
def preprocess_prompt(
    self,
    prompt: str,
    *,
    replace_explicit_value: bool = True,
) -> PromptEncoding:
    """Apply the released text normalization used before stage-1 parsing.

    Args:
        prompt: Natural-language text prompt.
        replace_explicit_value: Whether quoted values should be replaced by
            deterministic ``value_N`` placeholders.

    Returns:
        Normalized prompt and the placeholder recovery map.
    """
    result = prompt.replace("#", "").strip().lower()
    result = (
        result.replace("“", '"')
        .replace("”", '"')
        .replace("‘", "'")
        .replace("’", "'")
    )
    result = _WHITESPACE_RE.sub(" ", result)
    if not replace_explicit_value:
        return {"prompt": result, "value_map": None}
    return self._extract_explicit_values(result)

__call__

__call__(
    prompt: str | Sequence[str],
    *,
    replace_explicit_value: bool = True,
    return_tensors: Literal["pt"] = "pt",
) -> BatchEncoding

Tokenize prompt text for the semantic parser stage.

Source code in models/parse-then-place/src/parse_then_place/processing_parse_then_place.py
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
def __call__(
    self,
    prompt: str | Sequence[str],
    *,
    replace_explicit_value: bool = True,
    return_tensors: Literal["pt"] = "pt",
) -> BatchEncoding:
    """Tokenize prompt text for the semantic parser stage."""
    prompts = [prompt] if isinstance(prompt, str) else list(prompt)
    encodings = [
        self.preprocess_prompt(item, replace_explicit_value=replace_explicit_value)
        for item in prompts
    ]
    texts = [item["prompt"] for item in encodings]
    if self.parser_tokenizer is None:
        return BatchEncoding(
            {
                "prompt_text": texts,
                "value_maps": [item["value_map"] for item in encodings],
            }
        )
    parser_tokenizer = self.parser_tokenizer
    tokenized = parser_tokenizer(texts, return_tensors=return_tensors, padding=True)
    tokenized["value_maps"] = [item["value_map"] for item in encodings]
    tokenized["prompt_text"] = texts
    return cast(BatchEncoding, tokenized)

postprocess_ir

postprocess_ir(
    generated_ids: Int[Tensor, "batch tokens"]
    | Sequence[str],
    *,
    value_maps: list[dict[str, str] | None] | None = None,
) -> list[str]

Decode and lightly normalize stage-1 logical forms.

Source code in models/parse-then-place/src/parse_then_place/processing_parse_then_place.py
194
195
196
197
198
199
200
201
202
203
204
205
206
207
208
209
210
211
212
213
214
215
216
def postprocess_ir(
    self,
    generated_ids: Int[torch.Tensor, "batch tokens"] | Sequence[str],
    *,
    value_maps: list[dict[str, str] | None] | None = None,
) -> list[str]:
    """Decode and lightly normalize stage-1 logical forms."""
    if isinstance(generated_ids, torch.Tensor):
        if self.parser_tokenizer is None:
            raise ValueError("parser_tokenizer is required to decode generated ids")

        parser_tokenizer = self.parser_tokenizer
        logical_forms = parser_tokenizer.batch_decode(
            generated_ids, skip_special_tokens=True
        )
    else:
        logical_forms = list(generated_ids)
    if value_maps is None:
        return [self._normalize_ir(item) for item in logical_forms]
    return [
        self._recover_ir_values(self._normalize_ir(item), value_map)
        for item, value_map in zip(logical_forms, value_maps, strict=True)
    ]

ir_to_placement_inputs

ir_to_placement_inputs(
    logical_forms: Sequence[str],
) -> list[str]

Convert logical forms to placement-constraint strings.

Runtime keeps this method deterministic and accepts already-linearized constraints, which is also the artifact stored in stage-1 prediction JSON files. The current parity scripts do not execute the released grammar executor.

Source code in models/parse-then-place/src/parse_then_place/processing_parse_then_place.py
237
238
239
240
241
242
243
244
245
def ir_to_placement_inputs(self, logical_forms: Sequence[str]) -> list[str]:
    """Convert logical forms to placement-constraint strings.

    Runtime keeps this method deterministic and accepts already-linearized
    constraints, which is also the artifact stored in stage-1 prediction JSON
    files. The current parity scripts do not execute the released grammar
    executor.
    """
    return [self._logical_form_to_constraint(item) for item in logical_forms]

encode_placement_inputs

encode_placement_inputs(
    placement_inputs: Sequence[str],
    *,
    return_tensors: Literal["pt"] = "pt",
) -> BatchEncoding

Tokenize stage-2 placement constraints.

Source code in models/parse-then-place/src/parse_then_place/processing_parse_then_place.py
253
254
255
256
257
258
259
260
261
262
263
264
265
266
267
def encode_placement_inputs(
    self,
    placement_inputs: Sequence[str],
    *,
    return_tensors: Literal["pt"] = "pt",
) -> BatchEncoding:
    """Tokenize stage-2 placement constraints."""
    if self.placement_tokenizer is None:
        return BatchEncoding({"placement_text": list(placement_inputs)})
    placement_tokenizer = self.placement_tokenizer
    tokenized = placement_tokenizer(
        list(placement_inputs), return_tensors=return_tensors, padding=True
    )
    tokenized["placement_text"] = list(placement_inputs)
    return cast(BatchEncoding, tokenized)

decode_layout_sequences

decode_layout_sequences(
    generated_ids: Int[Tensor, "batch tokens"]
    | Sequence[str],
    *,
    batch_size: int,
    num_return_sequences: int,
) -> list[list[str]]

Decode stage-2 generated ids into grouped layout strings.

Source code in models/parse-then-place/src/parse_then_place/processing_parse_then_place.py
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
def decode_layout_sequences(
    self,
    generated_ids: Int[torch.Tensor, "batch tokens"] | Sequence[str],
    *,
    batch_size: int,
    num_return_sequences: int,
) -> list[list[str]]:
    """Decode stage-2 generated ids into grouped layout strings."""
    if isinstance(generated_ids, torch.Tensor):
        if self.placement_tokenizer is None:
            raise ValueError(
                "placement_tokenizer is required to decode generated ids"
            )

        placement_tokenizer = self.placement_tokenizer
        flat = placement_tokenizer.batch_decode(
            generated_ids, skip_special_tokens=True
        )
    else:
        flat = list(generated_ids)
    expected = batch_size * num_return_sequences
    if len(flat) != expected:
        raise ValueError(
            "Generated layout count does not match batch_size * num_return_sequences: "
            f"{len(flat)} != {expected}"
        )

    return [
        flat[idx * num_return_sequences : (idx + 1) * num_return_sequences]
        for idx in range(batch_size)
    ]

layout_text_to_output

layout_text_to_output(
    layout_text: Sequence[str] | Sequence[Sequence[str]],
    *,
    output_candidate: Literal[
        "first", "all", "best"
    ] = "first",
    output_type: Literal["dataclass", "dict"] = "dataclass",
    return_intermediates: bool = False,
) -> LayoutGenerationOutput | ParseThenPlaceOutputDict

Parse generated label left top width height text into schema.

Source code in models/parse-then-place/src/parse_then_place/processing_parse_then_place.py
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
def layout_text_to_output(
    self,
    layout_text: Sequence[str] | Sequence[Sequence[str]],
    *,
    output_candidate: Literal["first", "all", "best"] = "first",
    output_type: Literal["dataclass", "dict"] = "dataclass",
    return_intermediates: bool = False,
) -> LayoutGenerationOutput | ParseThenPlaceOutputDict:
    """Parse generated ``label left top width height`` text into schema."""
    candidate_groups = self._normalize_layout_text_groups(layout_text)
    selected = self._select_candidates(candidate_groups, output_candidate)
    parsed_groups = [self._parse_layout_text(item) for item in selected]
    max_len = max((len(item) for item in parsed_groups), default=0) or 1

    bbox_rows: list[Float[torch.Tensor, "elements 4"]] = []
    label_rows: list[Int[torch.Tensor, ...]] = []
    mask_rows: list[Bool[torch.Tensor, ...]] = []

    for parsed in parsed_groups:
        labels = torch.tensor([item["label"] for item in parsed], dtype=torch.long)
        boxes = torch.tensor([item["bbox"] for item in parsed], dtype=torch.float32)
        mask = torch.ones(len(parsed), dtype=torch.bool)

        if len(parsed) == 0:
            labels = torch.zeros(max_len, dtype=torch.long)
            boxes = torch.zeros(max_len, 4, dtype=torch.float32)
            mask = torch.zeros(max_len, dtype=torch.bool)
        elif len(parsed) < max_len:
            pad = max_len - len(parsed)
            labels = torch.nn.functional.pad(labels, (0, pad))
            boxes = torch.nn.functional.pad(boxes, (0, 0, 0, pad))
            mask = torch.nn.functional.pad(mask, (0, pad))

        label_rows.append(labels)
        bbox_rows.append(boxes)
        mask_rows.append(mask)

    raw_bbox = torch.stack(bbox_rows)
    bbox = normalize_boxes(
        raw_bbox,
        canvas_size=self.canvas_size,
        box_format="ltwh",
    )
    output = LayoutGenerationOutput(
        bbox=bbox.float(),
        labels=torch.stack(label_rows).long(),
        mask=torch.stack(mask_rows).bool(),
        id2label=dict(self.id2label),
        intermediates={
            "layout_text": selected,
            "layout_text_candidates": candidate_groups,
            "dataset_name": self.dataset_name,
            "canvas_size": self.canvas_size,
        }
        if return_intermediates
        else None,
    )
    if output_type == "dict":
        return dict(output)
    if output_type != "dataclass":
        raise ValueError(f"Unsupported output_type: {output_type}")

    return output