Image-to-Text
PEFT
Safetensors
Portuguese
English
vision-language
table-extraction
scientific-figures
markdown-table
qwen2.5-vl
lora
icdar-metric-loss
Instructions to use lucasoc/sci-image-models with libraries, inference providers, notebooks, and local apps. Follow these links to get started.
- Libraries
- PEFT
How to use lucasoc/sci-image-models with PEFT:
from peft import PeftModel from transformers import AutoModelForCausalLM base_model = AutoModelForCausalLM.from_pretrained("Qwen/Qwen2.5-VL-3B-Instruct") model = PeftModel.from_pretrained(base_model, "lucasoc/sci-image-models") - Notebooks
- Google Colab
- Kaggle
Download src/models/loader.py from lucasoc/sci-image-models: direct link, hf CLI and curl.
- Browser
- Download file 11.4 kB
-
https://huggingface.co/lucasoc/sci-image-models/resolve/main/src/models/loader.py
- Command line
-
hf download hf://lucasoc/sci-image-models/src/models/loader.py
-
curl -L -o loader.py https://huggingface.co/lucasoc/sci-image-models/resolve/main/src/models/loader.py
11.4 kB
| """ | |
| Model loader and PEFT / Quantization configuration. | |
| Supports Qwen2.5-VL and Qwen2-VL vision-language architectures with 4-bit QLoRA. | |
| """ | |
| from typing import Any, Dict, Optional, Tuple | |
| import torch | |
| from transformers import ( | |
| AutoProcessor, | |
| AutoModelForImageTextToText, | |
| BitsAndBytesConfig, | |
| Qwen2VLForConditionalGeneration, | |
| Qwen2_5_VLForConditionalGeneration, | |
| ) | |
| from transformers.models.qwen2_vl.modeling_qwen2_vl import Qwen2VLCausalLMOutputWithPast | |
| from peft import ( | |
| LoraConfig, | |
| get_peft_model, | |
| prepare_model_for_kbit_training, | |
| PeftModel, | |
| ) | |
| from ..utils.logging import setup_logger | |
| logger = setup_logger(__name__) | |
| def build_quantization_config(cfg: Dict[str, Any]) -> Optional[BitsAndBytesConfig]: | |
| """Constructs BitsAndBytesConfig for 4-bit / 8-bit QLoRA.""" | |
| q_cfg = cfg.get("quantization", {}) | |
| if not q_cfg.get("load_in_4bit", False): | |
| return None | |
| compute_dtype = torch.bfloat16 if q_cfg.get("bnb_4bit_compute_dtype") == "bfloat16" else torch.float16 | |
| return BitsAndBytesConfig( | |
| load_in_4bit=True, | |
| bnb_4bit_quant_type=q_cfg.get("bnb_4bit_quant_type", "nf4"), | |
| bnb_4bit_compute_dtype=compute_dtype, | |
| bnb_4bit_use_double_quant=q_cfg.get("bnb_4bit_use_double_quant", True), | |
| ) | |
| def build_peft_config(cfg: Dict[str, Any]) -> LoraConfig: | |
| """Constructs LoRA configuration.""" | |
| p_cfg = cfg.get("peft", {}) | |
| return LoraConfig( | |
| r=p_cfg.get("r", 16), | |
| lora_alpha=p_cfg.get("lora_alpha", 32), | |
| lora_dropout=p_cfg.get("lora_dropout", 0.05), | |
| bias=p_cfg.get("bias", "none"), | |
| task_type=p_cfg.get("task_type", "CAUSAL_LM"), | |
| target_modules=p_cfg.get("target_modules", ["q_proj", "v_proj"]), | |
| ) | |
| def load_model_and_processor( | |
| cfg: Dict[str, Any], | |
| is_training: bool = True, | |
| adapter_path: Optional[str] = None | |
| ) -> Tuple[Any, Any]: | |
| """Loads model and processor with optional LoRA / QLoRA wrapping.""" | |
| model_cfg = cfg.get("model", {}) | |
| model_name = model_cfg.get("name_or_path", "Qwen/Qwen2.5-VL-3B-Instruct") | |
| model_type = model_cfg.get("model_type", "qwen2_5_vl") | |
| trust_remote_code = model_cfg.get("trust_remote_code", True) | |
| logger.info(f"Loading processor for: {model_name}") | |
| processor_kwargs = {"trust_remote_code": trust_remote_code} | |
| data_cfg = cfg.get("data", {}) | |
| min_pixels = data_cfg.get("min_pixels") | |
| max_pixels = data_cfg.get("max_pixels") | |
| try: | |
| p_kwargs = dict(processor_kwargs) | |
| if min_pixels is not None: | |
| p_kwargs["min_pixels"] = min_pixels | |
| if max_pixels is not None: | |
| p_kwargs["max_pixels"] = max_pixels | |
| processor = AutoProcessor.from_pretrained(model_name, **p_kwargs) | |
| except TypeError: | |
| processor = AutoProcessor.from_pretrained(model_name, **processor_kwargs) | |
| load_4bit = cfg.get("quantization", {}).get("load_in_4bit", False) or (is_training and cfg.get("training", {}).get("method") == "qlora") | |
| bnb_config = build_quantization_config(cfg) if load_4bit else None | |
| torch_dtype = torch.bfloat16 if model_cfg.get("torch_dtype") == "bfloat16" else torch.float16 | |
| attn_impl = model_cfg.get("attn_implementation") | |
| if attn_impl == "flash_attention_2" and not (torch.cuda.is_available() and torch.cuda.get_device_capability()[0] >= 8): | |
| logger.info("FlashAttention-2 requires compute capability >= 8.0. Falling back to 'sdpa'.") | |
| attn_impl = "sdpa" | |
| logger.info(f"Loading base model: {model_name} (dtype={torch_dtype}, 4bit={bnb_config is not None}, attn={attn_impl})") | |
| kwargs = { | |
| "quantization_config": bnb_config, | |
| "torch_dtype": torch_dtype, | |
| "device_map": "auto", | |
| "trust_remote_code": trust_remote_code, | |
| } | |
| if attn_impl: | |
| kwargs["attn_implementation"] = attn_impl | |
| if "qwen2_5" in model_type or "qwen2.5" in model_type or "Qwen2.5" in model_name: | |
| model = Qwen2_5_VLForConditionalGeneration.from_pretrained(model_name, **kwargs) | |
| elif "qwen2" in model_type or "qwen" in model_type or "Qwen2" in model_name: | |
| model = Qwen2VLForConditionalGeneration.from_pretrained(model_name, **kwargs) | |
| else: | |
| model = AutoModelForImageTextToText.from_pretrained(model_name, **kwargs) | |
| if adapter_path: | |
| logger.info(f"Loading trained LoRA adapter from: {adapter_path}") | |
| model = PeftModel.from_pretrained(model, adapter_path) | |
| elif is_training and cfg.get("training", {}).get("method") in ["qlora", "lora"]: | |
| if bnb_config is not None: | |
| model = prepare_model_for_kbit_training( | |
| model, | |
| use_gradient_checkpointing=cfg.get("training", {}).get("gradient_checkpointing", True) | |
| ) | |
| peft_config = build_peft_config(cfg) | |
| model = get_peft_model(model, peft_config) | |
| model.print_trainable_parameters() | |
| # Apply selective memory-efficient forward pass to prevent OOM on large vocabulary logits | |
| if is_training: | |
| target_base = getattr(model, "base_model", None) | |
| if target_base is not None and hasattr(target_base, "model"): | |
| base_inner = target_base.model | |
| if isinstance(base_inner, (Qwen2VLForConditionalGeneration, Qwen2_5_VLForConditionalGeneration)): | |
| base_inner.forward = _memory_efficient_qwen2_vl_forward.__get__(base_inner, type(base_inner)) | |
| logger.info("Applied selective memory-efficient forward pass to base model.") | |
| elif isinstance(model, (Qwen2VLForConditionalGeneration, Qwen2_5_VLForConditionalGeneration)): | |
| model.forward = _memory_efficient_qwen2_vl_forward.__get__(model, type(model)) | |
| logger.info("Applied selective memory-efficient forward pass to model.") | |
| if cfg.get("training", {}).get("gradient_checkpointing", True) and hasattr(model, "gradient_checkpointing_enable"): | |
| model.gradient_checkpointing_enable() | |
| return model, processor | |
| def _memory_efficient_qwen2_vl_forward( | |
| self, | |
| input_ids=None, | |
| attention_mask=None, | |
| position_ids=None, | |
| past_key_values=None, | |
| inputs_embeds=None, | |
| labels=None, | |
| use_cache=None, | |
| pixel_values=None, | |
| pixel_values_videos=None, | |
| image_grid_thw=None, | |
| video_grid_thw=None, | |
| mm_token_type_ids=None, | |
| logits_to_keep=0, | |
| **kwargs, | |
| ): | |
| """Memory-efficient forward pass that only computes lm_head and cross-entropy on non-masked tokens.""" | |
| import torch.nn as nn | |
| outputs = self.model( | |
| input_ids=input_ids, | |
| pixel_values=pixel_values, | |
| pixel_values_videos=pixel_values_videos, | |
| image_grid_thw=image_grid_thw, | |
| video_grid_thw=video_grid_thw, | |
| mm_token_type_ids=mm_token_type_ids, | |
| position_ids=position_ids, | |
| attention_mask=attention_mask, | |
| past_key_values=past_key_values, | |
| inputs_embeds=inputs_embeds, | |
| use_cache=use_cache, | |
| **kwargs, | |
| ) | |
| hidden_states = outputs.last_hidden_state | |
| loss = None | |
| logits = None | |
| if labels is not None: | |
| shift_labels = nn.functional.pad(labels, (0, 1), value=-100) | |
| shift_labels = shift_labels[..., 1:].contiguous() | |
| shift_labels_flat = shift_labels.view(-1) | |
| valid_mask = shift_labels_flat != -100 | |
| if valid_mask.any(): | |
| shift_hidden = hidden_states.view(-1, hidden_states.shape[-1]) | |
| valid_hidden = shift_hidden[valid_mask] | |
| valid_labels = shift_labels_flat[valid_mask].to(valid_hidden.device) | |
| token_weights = kwargs.get("token_weights", getattr(self, "token_weights", None)) | |
| if token_weights is not None: | |
| token_weights = token_weights.to(valid_hidden.device) | |
| cfg_obj = getattr(self, "config", None) | |
| vocab_size = getattr(cfg_obj, "vocab_size", None) or getattr(getattr(cfg_obj, "text_config", None), "vocab_size", None) or 152064 | |
| if token_weights.shape[0] < vocab_size: | |
| token_weights = nn.functional.pad( | |
| token_weights, | |
| (0, vocab_size - token_weights.shape[0]), | |
| value=1.0, | |
| ) | |
| elif token_weights.shape[0] > vocab_size: | |
| token_weights = token_weights[:vocab_size] | |
| sample_weights = token_weights[valid_labels] if token_weights is not None else None | |
| focal_gamma = kwargs.get("focal_gamma", getattr(self, "focal_gamma", 0.0)) | |
| num_items_in_batch = kwargs.get("num_items_in_batch", None) | |
| reduction = "sum" if num_items_in_batch is not None else "mean" | |
| chunk_size = 64 | |
| if valid_hidden.shape[0] > chunk_size: | |
| losses = [] | |
| w_splits = sample_weights.split(chunk_size) if sample_weights is not None else [None] * ((valid_hidden.shape[0] + chunk_size - 1) // chunk_size) | |
| for chunk_h, chunk_y, chunk_w in zip(valid_hidden.split(chunk_size), valid_labels.split(chunk_size), w_splits): | |
| chunk_logits = self.lm_head(chunk_h).float() | |
| ce = nn.functional.cross_entropy(chunk_logits, chunk_y, reduction="none") | |
| if chunk_w is not None: | |
| ce = ce * chunk_w | |
| if focal_gamma > 0.0: | |
| pt = torch.exp(-ce.detach()) | |
| focal_w = (1.0 - pt) ** focal_gamma | |
| ce = focal_w * ce | |
| losses.append(ce.sum()) | |
| total_loss = torch.stack(losses).sum() | |
| if reduction == "mean": | |
| loss = total_loss / valid_hidden.shape[0] | |
| else: | |
| if torch.is_tensor(num_items_in_batch): | |
| num_items_in_batch = num_items_in_batch.to(total_loss.device) | |
| loss = total_loss / (num_items_in_batch if num_items_in_batch is not None else valid_hidden.shape[0]) | |
| else: | |
| valid_logits = self.lm_head(valid_hidden).float() | |
| ce = nn.functional.cross_entropy(valid_logits, valid_labels, reduction="none") | |
| if sample_weights is not None: | |
| ce = ce * sample_weights | |
| if focal_gamma > 0.0: | |
| pt = torch.exp(-ce.detach()) | |
| focal_w = (1.0 - pt) ** focal_gamma | |
| ce = focal_w * ce | |
| if reduction == "mean": | |
| loss = ce.mean() | |
| else: | |
| total_loss = ce.sum() | |
| if torch.is_tensor(num_items_in_batch): | |
| num_items_in_batch = num_items_in_batch.to(total_loss.device) | |
| loss = total_loss / (num_items_in_batch if num_items_in_batch is not None else valid_hidden.shape[0]) | |
| else: | |
| loss = torch.tensor(0.0, device=hidden_states.device, requires_grad=True) | |
| else: | |
| slice_indices = slice(-logits_to_keep, None) if isinstance(logits_to_keep, int) else logits_to_keep | |
| logits = self.lm_head(hidden_states[:, slice_indices, :]) | |
| return Qwen2VLCausalLMOutputWithPast( | |
| loss=loss, | |
| logits=logits, | |
| past_key_values=outputs.past_key_values, | |
| hidden_states=outputs.hidden_states, | |
| attentions=outputs.attentions, | |
| rope_deltas=outputs.rope_deltas, | |
| ) | |