Skip to content

unitorch.cli.models.gemma¤

GemmaProcessor¤

Tip

core/process/gemma is the section for configuration of GemmaProcessor.

Bases: GemmaProcessor

Processor for Gemma decoder-only generation tasks.

Source code in src/unitorch/cli/models/gemma/processing.py
19
20
21
22
23
24
25
26
27
28
29
30
31
32
33
def __init__(
    self,
    tokenizer_file: str,
    tokenizer_config: Optional[str] = None,
    chat_template: Optional[str] = None,
    max_seq_length: Optional[int] = 12800,
    max_gen_seq_length: Optional[int] = 512,
):
    super().__init__(
        tokenizer_file=tokenizer_file,
        tokenizer_config=tokenizer_config,
        chat_template=chat_template,
        max_seq_length=max_seq_length,
        max_gen_seq_length=max_gen_seq_length,
    )

from_config classmethod ¤

from_config(config, **kwargs)
Source code in src/unitorch/cli/models/gemma/processing.py
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
@classmethod
@config_defaults_init("core/process/gemma")
def from_config(cls, config, **kwargs):
    config.set_default_section("core/process/gemma")
    pretrained_name = config.getoption("pretrained_name", "gemma-4-12b")

    tokenizer_file = config.getoption("tokenizer_file", None)
    if tokenizer_file is None:
        tokenizer_file = resolve_pretrained_gemma_path(pretrained_name, "tokenizer")
    else:
        tokenizer_file = cached_path(tokenizer_file)

    tokenizer_config = config.getoption("tokenizer_config", None)
    if tokenizer_config is None:
        tokenizer_config = resolve_pretrained_gemma_path(
            pretrained_name,
            "tokenizer_config",
        )
    else:
        tokenizer_config = cached_path(tokenizer_config)

    chat_template = config.getoption("chat_template", None)
    chat_template = cached_path(chat_template) if chat_template is not None else None

    return {
        "tokenizer_file": tokenizer_file,
        "tokenizer_config": tokenizer_config,
        "chat_template": chat_template,
    }

_chat_template ¤

_chat_template(messages: List[Dict[str, Any]])
Source code in src/unitorch/cli/models/gemma/processing.py
65
66
67
68
69
70
@register_process("core/process/gemma/chat_template")
def _chat_template(
    self,
    messages: List[Dict[str, Any]],
):
    return super().chat_template(messages=messages)

_generation_inputs ¤

_generation_inputs(
    text: str, max_seq_length: Optional[int] = None
)
Source code in src/unitorch/cli/models/gemma/processing.py
72
73
74
75
76
77
78
79
80
81
82
83
84
85
@register_process("core/process/gemma/generation/inputs")
def _generation_inputs(
    self,
    text: str,
    max_seq_length: Optional[int] = None,
):
    outputs = super().generation_inputs(
        text=text,
        max_seq_length=max_seq_length,
    )
    return TensorInputs(
        input_ids=outputs.input_ids,
        attention_mask=outputs.attention_mask,
    )

_generation_labels ¤

_generation_labels(
    text: str, max_gen_seq_length: Optional[int] = None
)
Source code in src/unitorch/cli/models/gemma/processing.py
 87
 88
 89
 90
 91
 92
 93
 94
 95
 96
 97
 98
 99
100
@register_process("core/process/gemma/generation/labels")
def _generation_labels(
    self,
    text: str,
    max_gen_seq_length: Optional[int] = None,
):
    outputs = super().generation_labels(
        text=text,
        max_gen_seq_length=max_gen_seq_length,
    )
    return GenerationTargets(
        refs=outputs.input_ids,
        masks=outputs.attention_mask,
    )

_generation ¤

_generation(
    text: str,
    text_pair: str,
    max_seq_length: Optional[int] = None,
    max_gen_seq_length: Optional[int] = None,
)
Source code in src/unitorch/cli/models/gemma/processing.py
102
103
104
105
106
107
108
109
110
111
112
113
114
115
116
117
118
119
120
121
122
@register_process("core/process/gemma/generation")
def _generation(
    self,
    text: str,
    text_pair: str,
    max_seq_length: Optional[int] = None,
    max_gen_seq_length: Optional[int] = None,
):
    outputs = super().generation(
        text=text,
        text_pair=text_pair,
        max_seq_length=max_seq_length,
        max_gen_seq_length=max_gen_seq_length,
    )
    return TensorInputs(
        input_ids=outputs.input_ids,
        attention_mask=outputs.attention_mask,
    ), GenerationTargets(
        refs=outputs.input_ids_label,
        masks=outputs.attention_mask_label,
    )

_messages_generation ¤

_messages_generation(
    messages: List[Dict[str, Any]],
    max_seq_length: Optional[int] = None,
)
Source code in src/unitorch/cli/models/gemma/processing.py
124
125
126
127
128
129
130
131
132
133
134
135
136
137
138
139
140
@register_process("core/process/gemma/messages/generation")
def _messages_generation(
    self,
    messages: List[Dict[str, Any]],
    max_seq_length: Optional[int] = None,
):
    outputs = super().messages_generation(
        messages=messages,
        max_seq_length=max_seq_length,
    )
    return TensorInputs(
        input_ids=outputs.input_ids,
        attention_mask=outputs.attention_mask,
    ), GenerationTargets(
        refs=outputs.input_ids_label,
        masks=outputs.attention_mask_label,
    )

_detokenize ¤

_detokenize(outputs: GenerationOutputs)
Source code in src/unitorch/cli/models/gemma/processing.py
142
143
144
145
146
147
148
149
150
151
152
153
154
155
156
157
158
159
160
161
@register_process("core/postprocess/gemma/detokenize")
def _detokenize(
    self,
    outputs: GenerationOutputs,
):
    results = outputs.to_pandas()
    assert results.shape[0] == 0 or results.shape[0] == outputs.sequences.shape[0]

    decoded = super().detokenize(sequences=outputs.sequences)
    cleanup_string = lambda text: re.sub(r"\n", " ", text).strip()
    if isinstance(decoded[0], list):
        decoded = [list(map(cleanup_string, sequence)) for sequence in decoded]
    elif isinstance(decoded[0], str):
        decoded = list(map(cleanup_string, decoded))
    else:
        raise ValueError(
            f"Unsupported type for Gemma detokenize: {type(decoded[0])}"
        )
    results["decoded"] = decoded
    return WriterOutputs(results)

GemmaVLProcessor¤

Tip

core/process/gemma_vl is the section for configuration of GemmaVLProcessor.

Bases: GemmaVLProcessor

Processor for Gemma multimodal generation tasks.

Source code in src/unitorch/cli/models/gemma/processing_vl.py
19
20
21
22
23
24
25
26
27
28
29
30
31
32
33
34
35
def __init__(
    self,
    tokenizer_file: str,
    processor_config_path: str,
    tokenizer_config: Optional[str] = None,
    chat_template: Optional[str] = None,
    max_seq_length: Optional[int] = 12800,
    max_gen_seq_length: Optional[int] = 512,
):
    super().__init__(
        tokenizer_file=tokenizer_file,
        processor_config_path=processor_config_path,
        tokenizer_config=tokenizer_config,
        chat_template=chat_template,
        max_seq_length=max_seq_length,
        max_gen_seq_length=max_gen_seq_length,
    )

from_config classmethod ¤

from_config(config, **kwargs)
Source code in src/unitorch/cli/models/gemma/processing_vl.py
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
@classmethod
@config_defaults_init("core/process/gemma_vl")
def from_config(cls, config, **kwargs):
    config.set_default_section("core/process/gemma_vl")
    pretrained_name = config.getoption("pretrained_name", "gemma-4-12b")

    tokenizer_file = config.getoption("tokenizer_file", None)
    if tokenizer_file is None:
        tokenizer_file = resolve_pretrained_gemma_path(pretrained_name, "tokenizer")
    else:
        tokenizer_file = cached_path(tokenizer_file)

    processor_config_path = config.getoption("processor_config_path", None)
    if processor_config_path is None:
        processor_config_path = resolve_pretrained_gemma_path(
            pretrained_name,
            "processor_config",
        )
    else:
        processor_config_path = cached_path(processor_config_path)

    tokenizer_config = config.getoption("tokenizer_config", None)
    if tokenizer_config is None:
        tokenizer_config = resolve_pretrained_gemma_path(
            pretrained_name,
            "tokenizer_config",
        )
    else:
        tokenizer_config = cached_path(tokenizer_config)

    chat_template = config.getoption("chat_template", None)
    chat_template = cached_path(chat_template) if chat_template is not None else None

    return {
        "tokenizer_file": tokenizer_file,
        "processor_config_path": processor_config_path,
        "tokenizer_config": tokenizer_config,
        "chat_template": chat_template,
    }

_chat_template ¤

_chat_template(messages: List[Dict[str, Any]])
Source code in src/unitorch/cli/models/gemma/processing_vl.py
77
78
79
80
81
82
@register_process("core/process/gemma_vl/chat_template")
def _chat_template(
    self,
    messages: List[Dict[str, Any]],
):
    return super().chat_template(messages=messages)

_generation_inputs ¤

_generation_inputs(
    text: str,
    images: Union[Image, str, List[Image], List[str]],
    max_seq_length: Optional[int] = None,
)
Source code in src/unitorch/cli/models/gemma/processing_vl.py
 84
 85
 86
 87
 88
 89
 90
 91
 92
 93
 94
 95
 96
 97
 98
 99
100
101
102
@register_process("core/process/gemma_vl/generation/inputs")
def _generation_inputs(
    self,
    text: str,
    images: Union[Image.Image, str, List[Image.Image], List[str]],
    max_seq_length: Optional[int] = None,
):
    outputs = super().generation_inputs(
        text=text,
        images=images,
        max_seq_length=max_seq_length,
    )
    return TensorInputs(
        input_ids=outputs.input_ids,
        attention_mask=outputs.attention_mask,
        mm_token_type_ids=outputs.mm_token_type_ids,
        pixel_values=outputs.pixel_values,
        image_position_ids=outputs.image_position_ids,
    )

_generation_labels ¤

_generation_labels(
    text: str, max_gen_seq_length: Optional[int] = None
)
Source code in src/unitorch/cli/models/gemma/processing_vl.py
104
105
106
107
108
109
110
111
112
113
114
115
116
117
@register_process("core/process/gemma_vl/generation/labels")
def _generation_labels(
    self,
    text: str,
    max_gen_seq_length: Optional[int] = None,
):
    outputs = super().generation_labels(
        text=text,
        max_gen_seq_length=max_gen_seq_length,
    )
    return GenerationTargets(
        refs=outputs.input_ids,
        masks=outputs.attention_mask,
    )

_generation ¤

_generation(
    text: str,
    images: Union[Image, str, List[Image], List[str]],
    text_pair: str,
    max_seq_length: Optional[int] = None,
    max_gen_seq_length: Optional[int] = None,
)
Source code in src/unitorch/cli/models/gemma/processing_vl.py
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
@register_process("core/process/gemma_vl/generation")
def _generation(
    self,
    text: str,
    images: Union[Image.Image, str, List[Image.Image], List[str]],
    text_pair: str,
    max_seq_length: Optional[int] = None,
    max_gen_seq_length: Optional[int] = None,
):
    outputs = super().generation(
        text=text,
        images=images,
        text_pair=text_pair,
        max_seq_length=max_seq_length,
        max_gen_seq_length=max_gen_seq_length,
    )
    return TensorInputs(
        input_ids=outputs.input_ids,
        attention_mask=outputs.attention_mask,
        mm_token_type_ids=outputs.mm_token_type_ids,
        pixel_values=outputs.pixel_values,
        image_position_ids=outputs.image_position_ids,
    ), GenerationTargets(
        refs=outputs.input_ids_label,
        masks=outputs.attention_mask_label,
    )

_messages_generation ¤

_messages_generation(
    messages: List[Dict[str, Any]],
    images: Union[Image, str, List[Image], List[str]],
    max_seq_length: Optional[int] = None,
)
Source code in src/unitorch/cli/models/gemma/processing_vl.py
146
147
148
149
150
151
152
153
154
155
156
157
158
159
160
161
162
163
164
165
166
167
@register_process("core/process/gemma_vl/messages/generation")
def _messages_generation(
    self,
    messages: List[Dict[str, Any]],
    images: Union[Image.Image, str, List[Image.Image], List[str]],
    max_seq_length: Optional[int] = None,
):
    outputs = super().messages_generation(
        messages=messages,
        images=images,
        max_seq_length=max_seq_length,
    )
    return TensorInputs(
        input_ids=outputs.input_ids,
        attention_mask=outputs.attention_mask,
        mm_token_type_ids=outputs.mm_token_type_ids,
        pixel_values=outputs.pixel_values,
        image_position_ids=outputs.image_position_ids,
    ), GenerationTargets(
        refs=outputs.input_ids_label,
        masks=outputs.attention_mask_label,
    )

_detokenize ¤

_detokenize(outputs: GenerationOutputs)
Source code in src/unitorch/cli/models/gemma/processing_vl.py
169
170
171
172
173
174
175
176
177
178
179
180
181
182
183
184
185
186
187
188
@register_process("core/postprocess/gemma_vl/detokenize")
def _detokenize(
    self,
    outputs: GenerationOutputs,
):
    results = outputs.to_pandas()
    assert results.shape[0] == 0 or results.shape[0] == outputs.sequences.shape[0]

    decoded = super().detokenize(sequences=outputs.sequences)
    cleanup_string = lambda text: re.sub(r"\n", " ", text).strip()
    if isinstance(decoded[0], list):
        decoded = [list(map(cleanup_string, sequence)) for sequence in decoded]
    elif isinstance(decoded[0], str):
        decoded = list(map(cleanup_string, decoded))
    else:
        raise ValueError(
            f"Unsupported type for Gemma detokenize: {type(decoded[0])}"
        )
    results["decoded"] = decoded
    return WriterOutputs(results)

GemmaForGeneration¤

Tip

core/model/generation/gemma is the section for configuration of GemmaForGeneration.

Bases: GemmaForGeneration

Gemma model for text generation.

Source code in src/unitorch/cli/models/gemma/modeling.py
28
29
30
31
32
33
34
35
36
def __init__(
    self,
    config_path: str,
    gradient_checkpointing: Optional[bool] = False,
):
    super().__init__(
        config_path=config_path,
        gradient_checkpointing=gradient_checkpointing,
    )

from_config classmethod ¤

from_config(config, **kwargs)
Source code in src/unitorch/cli/models/gemma/modeling.py
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
@classmethod
@config_defaults_init("core/model/generation/gemma")
def from_config(cls, config, **kwargs):
    config.set_default_section("core/model/generation/gemma")
    pretrained_name = config.getoption("pretrained_name", "gemma-4-12b")
    pretrained_lora_name = config.getoption("pretrained_lora_name", None)

    config_path = config.getoption("config_path", None)
    if config_path is None:
        config_path = resolve_pretrained_gemma_path(pretrained_name, "config")
    else:
        config_path = cached_path(config_path)

    gradient_checkpointing = config.getoption("gradient_checkpointing", False)
    inst = cls(
        config_path=config_path,
        gradient_checkpointing=gradient_checkpointing,
    )

    pretrained_weight_path = config.getoption("pretrained_weight_path", None)
    weight_path = (
        pretrained_weight_path
        if pretrained_weight_path is not None
        else resolve_pretrained_gemma_path(pretrained_name, "weight")
    )
    if pretrained_weight_path is not None:
        weight_path = cached_path(weight_path)
    if weight_path is not None:
        inst.from_pretrained(weight_path)

    pretrained_lora_weight_path = config.getoption(
        "pretrained_lora_weight_path", None
    )
    lora_weight_path = (
        pretrained_lora_weight_path
        if pretrained_lora_weight_path is not None
        else pretrained_gemma_extensions_infos.get(pretrained_lora_name)
    )
    if pretrained_lora_weight_path is not None:
        lora_weight_path = cached_path(lora_weight_path)
    pretrained_lora_weight = config.getoption("pretrained_lora_weight", 1.0)
    pretrained_lora_alpha = config.getoption("pretrained_lora_alpha", 32.0)
    if lora_weight_path is not None:
        inst.load_lora_weights(
            lora_weight_path,
            lora_weights=pretrained_lora_weight,
            lora_alphas=pretrained_lora_alpha,
            save_base_state=False,
        )

    return inst

forward ¤

forward(
    input_ids: Tensor,
    attention_mask: Optional[Tensor] = None,
)
Source code in src/unitorch/cli/models/gemma/modeling.py
 90
 91
 92
 93
 94
 95
 96
 97
 98
 99
100
101
102
103
@autocast(
    device_type=("cuda" if torch.cuda.is_available() else "cpu"),
    dtype=(torch.bfloat16 if is_bfloat16_available() else torch.float32),
)
def forward(
    self,
    input_ids: torch.Tensor,
    attention_mask: Optional[torch.Tensor] = None,
):
    outputs = super().forward(
        input_ids=input_ids,
        attention_mask=attention_mask,
    )
    return GenerationOutputs(sequences=outputs)

generate ¤

generate(
    input_ids: Tensor,
    attention_mask: Optional[Tensor] = None,
    num_beams: Optional[int] = 5,
    decoder_start_token_id: Optional[int] = 2,
    decoder_end_token_id: Optional[
        Union[int, List[int]]
    ] = 1,
    decoder_pad_token_id: Optional[int] = 0,
    num_return_sequences: Optional[int] = 1,
    min_gen_seq_length: Optional[int] = 0,
    max_gen_seq_length: Optional[int] = 512,
    repetition_penalty: Optional[float] = 1.0,
    no_repeat_ngram_size: Optional[int] = 0,
    early_stopping: Optional[bool] = True,
    length_penalty: Optional[float] = 1.0,
    num_beam_groups: Optional[int] = 1,
    diversity_penalty: Optional[float] = 0.0,
    do_sample: Optional[bool] = False,
    temperature: Optional[float] = 1.0,
    top_k: Optional[int] = 50,
    top_p: Optional[float] = 1.0,
)
Source code in src/unitorch/cli/models/gemma/modeling.py
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
@config_defaults_method("core/model/generation/gemma")
@torch.no_grad()
@autocast(
    device_type=("cuda" if torch.cuda.is_available() else "cpu"),
    dtype=(torch.bfloat16 if is_bfloat16_available() else torch.float32),
)
def generate(
    self,
    input_ids: torch.Tensor,
    attention_mask: Optional[torch.Tensor] = None,
    num_beams: Optional[int] = 5,
    decoder_start_token_id: Optional[int] = 2,
    decoder_end_token_id: Optional[Union[int, List[int]]] = 1,
    decoder_pad_token_id: Optional[int] = 0,
    num_return_sequences: Optional[int] = 1,
    min_gen_seq_length: Optional[int] = 0,
    max_gen_seq_length: Optional[int] = 512,
    repetition_penalty: Optional[float] = 1.0,
    no_repeat_ngram_size: Optional[int] = 0,
    early_stopping: Optional[bool] = True,
    length_penalty: Optional[float] = 1.0,
    num_beam_groups: Optional[int] = 1,
    diversity_penalty: Optional[float] = 0.0,
    do_sample: Optional[bool] = False,
    temperature: Optional[float] = 1.0,
    top_k: Optional[int] = 50,
    top_p: Optional[float] = 1.0,
):
    outputs = super().generate(
        input_ids=input_ids,
        attention_mask=attention_mask,
        num_beams=num_beams,
        decoder_start_token_id=decoder_start_token_id,
        decoder_end_token_id=decoder_end_token_id,
        decoder_pad_token_id=decoder_pad_token_id,
        num_return_sequences=num_return_sequences,
        min_gen_seq_length=min_gen_seq_length,
        max_gen_seq_length=max_gen_seq_length,
        repetition_penalty=repetition_penalty,
        no_repeat_ngram_size=no_repeat_ngram_size,
        early_stopping=early_stopping,
        length_penalty=length_penalty,
        num_beam_groups=num_beam_groups,
        diversity_penalty=diversity_penalty,
        do_sample=do_sample,
        temperature=temperature,
        top_k=top_k,
        top_p=top_p,
    )
    return GenerationOutputs(
        sequences=outputs.sequences,
        sequences_scores=outputs.sequences_scores,
    )

GemmaVLForGeneration¤

Tip

core/model/generation/gemma_vl is the section for configuration of GemmaVLForGeneration.

Bases: GemmaVLForGeneration

Gemma multimodal model for image-grounded generation.

Source code in src/unitorch/cli/models/gemma/modeling_vl.py
28
29
30
31
32
33
34
35
36
def __init__(
    self,
    config_path: str,
    gradient_checkpointing: Optional[bool] = False,
):
    super().__init__(
        config_path=config_path,
        gradient_checkpointing=gradient_checkpointing,
    )

from_config classmethod ¤

from_config(config, **kwargs)
Source code in src/unitorch/cli/models/gemma/modeling_vl.py
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
@classmethod
@config_defaults_init("core/model/generation/gemma_vl")
def from_config(cls, config, **kwargs):
    config.set_default_section("core/model/generation/gemma_vl")
    pretrained_name = config.getoption("pretrained_name", "gemma-4-12b")
    pretrained_lora_name = config.getoption("pretrained_lora_name", None)

    config_path = config.getoption("config_path", None)
    if config_path is None:
        config_path = resolve_pretrained_gemma_path(pretrained_name, "config")
    else:
        config_path = cached_path(config_path)

    gradient_checkpointing = config.getoption("gradient_checkpointing", False)
    inst = cls(
        config_path=config_path,
        gradient_checkpointing=gradient_checkpointing,
    )

    pretrained_weight_path = config.getoption("pretrained_weight_path", None)
    weight_path = (
        pretrained_weight_path
        if pretrained_weight_path is not None
        else resolve_pretrained_gemma_path(pretrained_name, "weight")
    )
    if pretrained_weight_path is not None:
        weight_path = cached_path(weight_path)
    if weight_path is not None:
        inst.from_pretrained(weight_path)

    pretrained_lora_weight_path = config.getoption(
        "pretrained_lora_weight_path", None
    )
    lora_weight_path = (
        pretrained_lora_weight_path
        if pretrained_lora_weight_path is not None
        else pretrained_gemma_extensions_infos.get(pretrained_lora_name)
    )
    if pretrained_lora_weight_path is not None:
        lora_weight_path = cached_path(lora_weight_path)
    pretrained_lora_weight = config.getoption("pretrained_lora_weight", 1.0)
    pretrained_lora_alpha = config.getoption("pretrained_lora_alpha", 32.0)
    if lora_weight_path is not None:
        inst.load_lora_weights(
            lora_weight_path,
            lora_weights=pretrained_lora_weight,
            lora_alphas=pretrained_lora_alpha,
            save_base_state=False,
        )

    return inst

forward ¤

forward(
    input_ids: Tensor,
    pixel_values: Optional[Tensor] = None,
    image_position_ids: Optional[Tensor] = None,
    attention_mask: Optional[Tensor] = None,
    mm_token_type_ids: Optional[Tensor] = None,
)
Source code in src/unitorch/cli/models/gemma/modeling_vl.py
 90
 91
 92
 93
 94
 95
 96
 97
 98
 99
100
101
102
103
104
105
106
107
108
109
@autocast(
    device_type=("cuda" if torch.cuda.is_available() else "cpu"),
    dtype=(torch.bfloat16 if is_bfloat16_available() else torch.float32),
)
def forward(
    self,
    input_ids: torch.Tensor,
    pixel_values: Optional[torch.Tensor] = None,
    image_position_ids: Optional[torch.Tensor] = None,
    attention_mask: Optional[torch.Tensor] = None,
    mm_token_type_ids: Optional[torch.Tensor] = None,
):
    outputs = super().forward(
        input_ids=input_ids,
        pixel_values=pixel_values,
        image_position_ids=image_position_ids,
        attention_mask=attention_mask,
        mm_token_type_ids=mm_token_type_ids,
    )
    return GenerationOutputs(sequences=outputs)

generate ¤

generate(
    input_ids: Tensor,
    pixel_values: Optional[Tensor] = None,
    image_position_ids: Optional[Tensor] = None,
    attention_mask: Optional[Tensor] = None,
    mm_token_type_ids: Optional[Tensor] = None,
    num_beams: Optional[int] = 5,
    decoder_start_token_id: Optional[int] = 2,
    decoder_end_token_id: Optional[
        Union[int, List[int]]
    ] = 1,
    decoder_pad_token_id: Optional[int] = 0,
    num_return_sequences: Optional[int] = 1,
    min_gen_seq_length: Optional[int] = 0,
    max_gen_seq_length: Optional[int] = 512,
    repetition_penalty: Optional[float] = 1.0,
    no_repeat_ngram_size: Optional[int] = 0,
    early_stopping: Optional[bool] = True,
    length_penalty: Optional[float] = 1.0,
    num_beam_groups: Optional[int] = 1,
    diversity_penalty: Optional[float] = 0.0,
    do_sample: Optional[bool] = False,
    temperature: Optional[float] = 1.0,
    top_k: Optional[int] = 50,
    top_p: Optional[float] = 1.0,
)
Source code in src/unitorch/cli/models/gemma/modeling_vl.py
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
@config_defaults_method("core/model/generation/gemma_vl")
@torch.no_grad()
@autocast(
    device_type=("cuda" if torch.cuda.is_available() else "cpu"),
    dtype=(torch.bfloat16 if is_bfloat16_available() else torch.float32),
)
def generate(
    self,
    input_ids: torch.Tensor,
    pixel_values: Optional[torch.Tensor] = None,
    image_position_ids: Optional[torch.Tensor] = None,
    attention_mask: Optional[torch.Tensor] = None,
    mm_token_type_ids: Optional[torch.Tensor] = None,
    num_beams: Optional[int] = 5,
    decoder_start_token_id: Optional[int] = 2,
    decoder_end_token_id: Optional[Union[int, List[int]]] = 1,
    decoder_pad_token_id: Optional[int] = 0,
    num_return_sequences: Optional[int] = 1,
    min_gen_seq_length: Optional[int] = 0,
    max_gen_seq_length: Optional[int] = 512,
    repetition_penalty: Optional[float] = 1.0,
    no_repeat_ngram_size: Optional[int] = 0,
    early_stopping: Optional[bool] = True,
    length_penalty: Optional[float] = 1.0,
    num_beam_groups: Optional[int] = 1,
    diversity_penalty: Optional[float] = 0.0,
    do_sample: Optional[bool] = False,
    temperature: Optional[float] = 1.0,
    top_k: Optional[int] = 50,
    top_p: Optional[float] = 1.0,
):
    outputs = super().generate(
        input_ids=input_ids,
        pixel_values=pixel_values,
        image_position_ids=image_position_ids,
        attention_mask=attention_mask,
        mm_token_type_ids=mm_token_type_ids,
        num_beams=num_beams,
        decoder_start_token_id=decoder_start_token_id,
        decoder_end_token_id=decoder_end_token_id,
        decoder_pad_token_id=decoder_pad_token_id,
        num_return_sequences=num_return_sequences,
        min_gen_seq_length=min_gen_seq_length,
        max_gen_seq_length=max_gen_seq_length,
        repetition_penalty=repetition_penalty,
        no_repeat_ngram_size=no_repeat_ngram_size,
        early_stopping=early_stopping,
        length_penalty=length_penalty,
        num_beam_groups=num_beam_groups,
        diversity_penalty=diversity_penalty,
        do_sample=do_sample,
        temperature=temperature,
        top_k=top_k,
        top_p=top_p,
    )
    return GenerationOutputs(
        sequences=outputs.sequences,
        sequences_scores=outputs.sequences_scores,
    )