Skip to content
Open
Show file tree
Hide file tree
Changes from all commits
Commits
File filter

Filter by extension

Filter by extension

Conversations
Failed to load comments.
Loading
Jump to
Jump to file
Failed to load files.
Loading
Diff view
Diff view
13 changes: 13 additions & 0 deletions litgpt/adapter.py
Original file line number Diff line number Diff line change
Expand Up @@ -19,6 +19,7 @@
from litgpt.model import GPT as BaseModel
from litgpt.model import Block as BaseBlock
from litgpt.model import CausalSelfAttention as BaseCausalSelfAttention
from litgpt.vision import MultiModalProjector, VisionEncoder


@dataclass
Expand All @@ -45,6 +46,18 @@ def __init__(self, config: Config) -> None:
self.mask_cache: torch.Tensor | None = None
self.max_seq_length = self.config.block_size

# Optional vision encoder for multimodal models
if config.is_multimodal:
self.vision_encoder = VisionEncoder(config, pretrained_model_name=config.vision_model_name)
self.mm_projector = MultiModalProjector(
vision_dim=config.vision_feature_dim,
text_dim=config.n_embd,
projector_type=config.mm_projector_type or "linear",
)
else:
self.vision_encoder = None
self.mm_projector = None

@classmethod
def from_name(cls, name: str, **kwargs: Any) -> Self:
return cls(Config.from_name(name, **kwargs))
Expand Down
13 changes: 13 additions & 0 deletions litgpt/adapter_v2.py
Original file line number Diff line number Diff line change
Expand Up @@ -22,6 +22,7 @@
from litgpt.model import Block as BaseBlock
from litgpt.scripts.convert_hf_checkpoint import qkv_reassemble
from litgpt.utils import map_old_state_dict_weights
from litgpt.vision import MultiModalProjector, VisionEncoder


@dataclass
Expand Down Expand Up @@ -80,6 +81,18 @@ def __init__(self, config: Config) -> None:
self.mask_cache: torch.Tensor | None = None
self.max_seq_length = self.config.block_size

# Optional vision encoder for multimodal models
if config.is_multimodal:
self.vision_encoder = VisionEncoder(config, pretrained_model_name=config.vision_model_name)
self.mm_projector = MultiModalProjector(
vision_dim=config.vision_feature_dim,
text_dim=config.n_embd,
projector_type=config.mm_projector_type or "linear",
)
else:
self.vision_encoder = None
self.mm_projector = None

@classmethod
def from_name(cls, name: str, **kwargs: Any) -> Self:
return cls(Config.from_name(name, **kwargs))
Expand Down
22 changes: 22 additions & 0 deletions litgpt/api.py
Original file line number Diff line number Diff line change
Expand Up @@ -33,6 +33,7 @@
load_checkpoint,
save_config,
)
from litgpt.vision import ImagePreprocessor


class LLM(torch.nn.Module):
Expand Down Expand Up @@ -469,6 +470,7 @@ def generate(
top_p: float = 1.0,
return_as_token_ids: bool = False,
stream: bool = False,
image: str | Path | None = None,
) -> str | torch.Tensor:
"""
Takes a conditioning sequence (prompt) as input and continues to generate as many tokens as requested.
Expand Down Expand Up @@ -499,12 +501,31 @@ def generate(
stream: If True, returns a generator that yields tokens as they are generated.
At the moment, this setting is slower and may use more memory than the non-streaming version.
We plan to resolve this in the future.
image: Optional path to an image file for multimodal models.
"""
if self.model is None:
raise AttributeError(
"The model is not initialized yet; use the .distribute() "
"or .trainer_setup() method to initialize the model."
)

# Preprocess image if provided
pixel_values = None
if image is not None:
if not self.config.is_multimodal:
raise ValueError(
"An image was provided but the model is not multimodal. "
"Ensure the model config has vision_feature_dim set."
)
preprocessor = ImagePreprocessor(
image_size=self.config.vision_image_size or 224,
)
if self.fabric is not None:
device = self.fabric.device
else:
device = self.preprocessor.device
pixel_values = preprocessor(image, device=device)

input_ids = self._text_to_token_ids(prompt, sys_prompt)
prompt_length = input_ids.size(0)
max_returned_tokens = prompt_length + max_new_tokens
Expand Down Expand Up @@ -559,6 +580,7 @@ def iterator():
top_p=top_p,
eos_id=self.preprocessor.tokenizer.eos_id,
include_prompt=False,
pixel_values=pixel_values,
)

if stream:
Expand Down
54 changes: 51 additions & 3 deletions litgpt/chat/base.py
Original file line number Diff line number Diff line change
Expand Up @@ -35,6 +35,7 @@ def generate(
top_k: int | None = None,
top_p: float = 1.0,
stop_tokens: tuple[list[int], ...] = (),
pixel_values: torch.Tensor | None = None,
) -> Iterator[torch.Tensor]:
"""Takes a conditioning sequence (prompt) as input and continues to generate as many tokens as possible.

Expand Down Expand Up @@ -72,11 +73,22 @@ def generate(
top_k=top_k,
top_p=top_p,
stop_tokens=stop_tokens,
pixel_values=pixel_values,
)


def process_prompt(
prompt, model, tokenizer, prompt_style, fabric, temperature, max_new_tokens, top_k, top_p, stop_tokens
prompt,
model,
tokenizer,
prompt_style,
fabric,
temperature,
max_new_tokens,
top_k,
top_p,
stop_tokens,
pixel_values: torch.Tensor | None = None,
):
prompt = prompt_style.apply(prompt=prompt)
encoded_prompt = tokenizer.encode(prompt, device=fabric.device)
Expand All @@ -98,6 +110,7 @@ def process_prompt(
top_k=top_k,
top_p=top_p,
stop_tokens=stop_tokens,
pixel_values=pixel_values,
)
token_generator: Iterator[str] = tokenizer.decode_stream(y, device=fabric.device)

Expand All @@ -121,7 +134,30 @@ def process_prompt(
fabric.print()


def interact(multiline, model, tokenizer, prompt_style, fabric, temperature, max_new_tokens, top_k, top_p, stop_tokens):
def interact(
multiline,
model,
tokenizer,
prompt_style,
fabric,
temperature,
max_new_tokens,
top_k,
top_p,
stop_tokens,
initial_image: Path | None = None,
):
pixel_values = None
if initial_image is not None:
if not model.config.is_multimodal:
fabric.print("Warning: An image was provided but the model is not multimodal.", file=sys.stderr)
else:
from litgpt.vision import ImagePreprocessor

preprocessor = ImagePreprocessor(image_size=model.config.vision_image_size or 224)
pixel_values = preprocessor(initial_image, device=fabric.device)
fabric.print(f">> Loaded image: {initial_image}")

while True:
try:
if not multiline:
Expand All @@ -144,7 +180,17 @@ def interact(multiline, model, tokenizer, prompt_style, fabric, temperature, max
break

process_prompt(
prompt, model, tokenizer, prompt_style, fabric, temperature, max_new_tokens, top_k, top_p, stop_tokens
prompt,
model,
tokenizer,
prompt_style,
fabric,
temperature,
max_new_tokens,
top_k,
top_p,
stop_tokens,
pixel_values,
)


Expand All @@ -161,6 +207,7 @@ def main(
compile: bool = False,
multiline: bool = False,
access_token: str | None = None,
image: Path | None = None,
) -> None:
"""Chat with a model.

Expand Down Expand Up @@ -267,6 +314,7 @@ def main(
top_k=top_k,
top_p=top_p,
stop_tokens=stop_tokens,
initial_image=image,
)

if fabric.device.type == "cuda":
Expand Down
11 changes: 11 additions & 0 deletions litgpt/config.py
Original file line number Diff line number Diff line change
Expand Up @@ -114,6 +114,17 @@ class Config:
# `rope_base` is used, for 1 `rope_local_base_freq` is used. If
# `len(rope_indices) > n_layer`, we only use the initial part.
rope_indices: list[int] | None = None
# Vision encoder config (optional, for multimodal models)
vision_feature_dim: int | None = None # Output dim of vision encoder (e.g. 1152 for SigLIP)
vision_start_token_id: int | None = None # Token ID for <image> placeholder
vision_patch_size: int | None = None # Patch size (e.g. 14)
vision_image_size: int | None = None # Input image size (e.g. 224)
mm_projector_type: str | None = None # "linear" or "mlp2x"
vision_model_name: str | None = None # HF model name for the vision encoder

@property
def is_multimodal(self) -> bool:
return self.vision_feature_dim is not None

def __post_init__(self):
if not self.name:
Expand Down
9 changes: 8 additions & 1 deletion litgpt/generate/base.py
Original file line number Diff line number Diff line change
Expand Up @@ -79,9 +79,10 @@ def next_token(
input_pos: torch.Tensor,
x: torch.Tensor,
input_pos_maxp1: int | None = None,
pixel_values: torch.Tensor | None = None,
**sample_kwargs: dict[str, Any],
) -> torch.Tensor:
logits = model(x, input_pos, input_pos_maxp1=input_pos_maxp1)
logits = model(x, input_pos, input_pos_maxp1=input_pos_maxp1, pixel_values=pixel_values)
_next = sample(logits, **sample_kwargs).to(dtype=torch.int64)
return _next

Expand Down Expand Up @@ -137,6 +138,7 @@ def generate_fn(
stop_tokens: tuple[list[int], ...] = (),
include_prompt: bool,
include_eos: bool,
pixel_values: torch.Tensor | None = None,
) -> Iterator[torch.Tensor]:
"""
Generates tokens for a single prompt.
Expand Down Expand Up @@ -182,11 +184,13 @@ def generate_fn(
input_pos_maxp1 = prompt_size if all(m.__class__.__name__ != "ThunderModule" for m in model.modules()) else None
for current_idx in range(max_returned_tokens - prompt_size):
# Generate the token
# pixel_values is only needed for the prefill (first forward pass)
token = next_token(
model,
input_pos,
token.view(1, -1),
input_pos_maxp1=input_pos_maxp1,
pixel_values=pixel_values if prefill_token else None,
temperature=temperature,
top_k=top_k,
top_p=top_p,
Expand Down Expand Up @@ -380,6 +384,7 @@ def generate(
top_p: float = 1.0,
eos_id: int | None = None,
include_prompt: bool = True,
pixel_values: torch.Tensor | None = None,
) -> torch.Tensor:
"""
Takes a conditioning sequence (prompt) as input and continues to generate as many tokens as requested.
Expand Down Expand Up @@ -407,6 +412,7 @@ def generate(
or https://huyenchip.com/2024/01/16/sampling.html#top_p
eos_id: If specified, stop generating any more token once the <eos> token is triggered.
include_prompt: If true (default) prepends the prompt (after applying the prompt style) to the output.
pixel_values: Optional image tensor for multimodal models.
"""

token_list = list(
Expand All @@ -420,6 +426,7 @@ def generate(
top_k=top_k,
top_p=top_p,
stop_tokens=(([eos_id],) if eos_id is not None else ()),
pixel_values=pixel_values,
)
)

Expand Down
13 changes: 13 additions & 0 deletions litgpt/lora.py
Original file line number Diff line number Diff line change
Expand Up @@ -59,6 +59,7 @@
from litgpt.model import CausalSelfAttention as BaseCausalSelfAttention
from litgpt.scripts.convert_hf_checkpoint import qkv_reassemble
from litgpt.utils import map_old_state_dict_weights
from litgpt.vision import MultiModalProjector, VisionEncoder


class LoRALayer(nn.Module):
Expand Down Expand Up @@ -503,6 +504,18 @@ def __init__(self, config: Config) -> None:
self.mask_cache: torch.Tensor | None = None
self.max_seq_length = self.config.block_size

# Optional vision encoder for multimodal models
if config.is_multimodal:
self.vision_encoder = VisionEncoder(config, pretrained_model_name=config.vision_model_name)
self.mm_projector = MultiModalProjector(
vision_dim=config.vision_feature_dim,
text_dim=config.n_embd,
projector_type=config.mm_projector_type or "linear",
)
else:
self.vision_encoder = None
self.mm_projector = None

@classmethod
def from_name(cls, name: str, **kwargs: Any) -> Self:
return cls(Config.from_name(name, **kwargs))
Expand Down
20 changes: 20 additions & 0 deletions litgpt/model.py
Original file line number Diff line number Diff line change
Expand Up @@ -17,6 +17,7 @@

from litgpt.config import Config
from litgpt.scripts.convert_hf_checkpoint import qkv_reassemble
from litgpt.vision import MultiModalProjector, VisionEncoder, merge_input_embeds


class GPT(nn.Module):
Expand All @@ -36,6 +37,18 @@ def __init__(self, config: Config) -> None:
self.mask_cache: torch.Tensor | None = None
self.max_seq_length = self.config.block_size

# Optional vision encoder for multimodal models
if config.is_multimodal:
self.vision_encoder = VisionEncoder(config, pretrained_model_name=config.vision_model_name)
self.mm_projector = MultiModalProjector(
vision_dim=config.vision_feature_dim,
text_dim=config.n_embd,
projector_type=config.mm_projector_type or "linear",
)
else:
self.vision_encoder = None
self.mm_projector = None

@property
def max_seq_length(self) -> int:
return self._max_seq_length
Expand Down Expand Up @@ -88,6 +101,7 @@ def forward(
input_pos: torch.Tensor | None = None,
input_pos_maxp1: int | None = None,
lm_head_chunk_size: int = 0,
pixel_values: torch.Tensor | None = None,
) -> torch.Tensor | list[torch.Tensor]:
"""
If `input_pos` is provided, the KV cache uses K and V vectors for
Expand Down Expand Up @@ -156,6 +170,12 @@ def forward(
if self.config.scale_embeddings:
x = x * torch.tensor(self.config.n_embd**0.5, dtype=x.dtype)

# Merge image embeddings if pixel_values are provided
if pixel_values is not None and self.vision_encoder is not None:
image_features = self.vision_encoder(pixel_values)
image_embeds = self.mm_projector(image_features)
x = merge_input_embeds(x, image_embeds, self.config.vision_start_token_id, idx)

for block_idx, block in enumerate(self.transformer.h):
if self.config.rope_indices is not None:
x = block(
Expand Down
Loading