> For clean Markdown of any page, append .md to the page URL.
> For a complete documentation index, see https://docs.nvidia.com/nemo/automodel/llms.txt.
> For AI client integration (Claude Code, Cursor, etc.), connect to the MCP server at https://docs.nvidia.com/nemo/automodel/_mcp/server.

# nemo_automodel.components.speculative.dspark.draft_gemma4

## Module Contents

### Classes

| Name                                                                                                              | Description |
| ----------------------------------------------------------------------------------------------------------------- | ----------- |
| [`Gemma4DSparkAttention`](#nemo_automodel-components-speculative-dspark-draft_gemma4-Gemma4DSparkAttention)       | -           |
| [`Gemma4DSparkDecoderLayer`](#nemo_automodel-components-speculative-dspark-draft_gemma4-Gemma4DSparkDecoderLayer) | -           |
| [`Gemma4DSparkModel`](#nemo_automodel-components-speculative-dspark-draft_gemma4-Gemma4DSparkModel)               | -           |

### Data

[`__all__`](#nemo_automodel-components-speculative-dspark-draft_gemma4-__all__)

### API

```python
class nemo_automodel.components.speculative.dspark.draft_gemma4.Gemma4DSparkAttention(
    config,
    layer_idx: int
)
```

**Bases:** `Module`

**`attention_dropout`** `= float(config.attention_dropout)`

---

**`head_dim`** `= int(config.global_head_dim)`

---

**`k_norm`**

---

**`k_proj`**

---

**`layer_idx`** `= int(layer_idx)`

---

**`num_attention_heads`** `= int(config.num_attention_heads)`

---

**`num_key_value_groups`**

---

**`num_key_value_heads`** `= int(config.num_global_key_value_heads)`

---

**`o_proj`**

---

**`q_norm`**

---

**`q_proj`**

---

**`scaling`** `= 1.0`

---

**`use_alternative_attention`** `= bool(config.attention_k_eq_v)`

---

**`v_norm`**

---

```python
nemo_automodel.components.speculative.dspark.draft_gemma4.Gemma4DSparkAttention._repeat_kv(
    hidden_states: torch.Tensor
) -> torch.Tensor
```

```python
nemo_automodel.components.speculative.dspark.draft_gemma4.Gemma4DSparkAttention.forward(
    hidden_states: torch.Tensor,
    target_hidden_states: torch.Tensor,
    position_embeddings: tuple[torch.Tensor, torch.Tensor],
    attention_mask: torch.Tensor | None,
    past_key_values: transformers.cache_utils.Cache | None = None,
    cache_position: torch.LongTensor | None = None,
    kwargs = {}
) -> tuple[torch.Tensor, torch.Tensor | None]
```

```python
class nemo_automodel.components.speculative.dspark.draft_gemma4.Gemma4DSparkDecoderLayer(
    config,
    layer_idx: int
)
```

**Bases:** `GradientCheckpointingLayer`

**`hidden_size`** `= config.hidden_size`

---

**`input_layernorm`**

---

**`mlp`** `= Gemma4TextMLP(config, layer_idx)`

---

**`post_attention_layernorm`**

---

**`post_feedforward_layernorm`**

---

**`pre_feedforward_layernorm`**

---

**`self_attn`**

---

```python
nemo_automodel.components.speculative.dspark.draft_gemma4.Gemma4DSparkDecoderLayer.forward(
    target_hidden_states: torch.Tensor | None = None,
    hidden_states: torch.Tensor | None = None,
    attention_mask: torch.Tensor | None = None,
    position_ids: torch.LongTensor | None = None,
    past_key_value: transformers.cache_utils.Cache | None = None,
    output_attentions: bool | None = False,
    use_cache: bool | None = False,
    cache_position: torch.LongTensor | None = None,
    position_embeddings: tuple[torch.Tensor, torch.Tensor] | None = None,
    kwargs = {}
) -> torch.Tensor
```

```python
class nemo_automodel.components.speculative.dspark.draft_gemma4.Gemma4DSparkModel(
    config
)
```

**Bases:** `Gemma4PreTrainedModel`

**`_no_split_modules`** `= ['Gemma4DSparkDecoderLayer']`

---

**`base_model_prefix`** `= 'model'`

---

**`block_size`** `= int(config.block_size)`

---

**`embed_tokens`**

---

**`enable_confidence_head`** `= bool(config.enable_confidence_head)`

---

**`fc`**

---

**`hidden_norm`**

---

**`layers`**

---

**`lm_head`**

---

**`markov_head`** `= build_markov_head(config)`

---

**`mask_token_id`** `= config.mask_token_id`

---

**`norm`**

---

**`num_anchors`** `= int(config.num_anchors)`

---

**`rotary_emb`**

---

**`target_layer_ids`** `= config.target_layer_ids`

---

```python
nemo_automodel.components.speculative.dspark.draft_gemma4.Gemma4DSparkModel._apply(
    fn,
    recurse: bool = True
)
```

Keep the RoPE `inv_freq` buffer in fp32 across dtype casts.

`model.to(bfloat16)` (the training build path) would otherwise round
`inv_freq` to bf16 and dephase RoPE with absolute position, eroding
draft acceptance (see `pin_rope_inv_freq_fp32`).

```python
nemo_automodel.components.speculative.dspark.draft_gemma4.Gemma4DSparkModel._forward_backbone(
    position_ids: torch.LongTensor,
    attention_mask: torch.Tensor | None = None,
    noise_embedding: torch.Tensor | None = None,
    target_hidden_states: torch.Tensor | None = None,
    past_key_values: transformers.cache_utils.Cache | None = None,
    use_cache: bool = False,
    kwargs = {}
) -> torch.Tensor
```

```python
nemo_automodel.components.speculative.dspark.draft_gemma4.Gemma4DSparkModel.compute_logits(
    hidden_states: torch.Tensor
) -> torch.Tensor
```

```python
nemo_automodel.components.speculative.dspark.draft_gemma4.Gemma4DSparkModel.forward(
    input_ids: torch.Tensor,
    target_hidden_states: torch.Tensor,
    loss_mask: torch.Tensor,
    target_last_hidden_states: torch.Tensor | None = None
) -> nemo_automodel.components.speculative.dspark.common.DSparkForwardOutput
```

```python
nemo_automodel.components.speculative.dspark.draft_gemma4.Gemma4DSparkModel.initialize_embeddings_and_head(
    embed_tokens: torch.nn.Module,
    lm_head: torch.nn.Module,
    freeze: bool = True
)
```

```python
nemo_automodel.components.speculative.dspark.draft_gemma4.Gemma4DSparkModel.predict_confidence_step(
    hidden_states: torch.Tensor,
    prev_token_ids: torch.Tensor | None = None
) -> torch.Tensor | None
```

```python
nemo_automodel.components.speculative.dspark.draft_gemma4.Gemma4DSparkModel.sample_draft_token_step(
    base_logits: torch.Tensor,
    prev_token_ids: torch.Tensor,
    temperature: float = 0.0,
    hidden_states: torch.Tensor | None = None
) -> tuple[torch.Tensor, torch.Tensor]
```

```python
nemo_automodel.components.speculative.dspark.draft_gemma4.Gemma4DSparkModel.sample_draft_tokens(
    base_logits: torch.Tensor,
    first_prev_token_ids: torch.Tensor,
    temperature: float = 0.0,
    hidden_states: torch.Tensor | None = None
) -> tuple[torch.Tensor, torch.Tensor]
```

```python
nemo_automodel.components.speculative.dspark.draft_gemma4.Gemma4DSparkModel.set_embedding_head_trainable(
    trainable: bool
)
```

```python
nemo_automodel.components.speculative.dspark.draft_gemma4.__all__ = ['Gemma4DSparkModel']
```