""" CAD2Program Model Implementation Based on "From 2D CAD Drawings to 3D Parametric Models: A Vision-Language Approach" This implements the core vision-language model for converting 2D CAD drawings to 3D parametric models. """ import torch import torch.nn as nn import torch.nn.functional as F from torch.nn import CrossEntropyLoss from transformers import ( AutoModel, AutoTokenizer, AutoImageProcessor, PreTrainedModel, GenerationMixin, ViTModel, ViTConfig, GPT2LMHeadModel, GPT2Config ) from transformers.modeling_outputs import BaseModelOutput, CausalLMOutputWithPast from typing import Optional, Tuple, Union, List, Dict, Any import numpy as np from PIL import Image from model_config import CAD2ProgramConfig from utils.cad_primitives import PRIMITIVE_REGISTRY, CabinetAssembly class VisionLanguageProjector(nn.Module): """Projects vision features to language model dimension""" def __init__(self, vision_dim: int, language_dim: int, hidden_dim: int = None): super().__init__() if hidden_dim is None: hidden_dim = max(vision_dim, language_dim) self.projector = nn.Sequential( nn.Linear(vision_dim, hidden_dim), nn.GELU(), nn.Dropout(0.1), nn.Linear(hidden_dim, language_dim), nn.LayerNorm(language_dim) ) def forward(self, vision_features: torch.Tensor) -> torch.Tensor: """ Project vision features to language space Args: vision_features: [batch_size, seq_len, vision_dim] Returns: projected_features: [batch_size, seq_len, language_dim] """ return self.projector(vision_features) class PrimitiveEmbedding(nn.Module): """Special embeddings for CAD primitives as mentioned in the paper""" def __init__(self, num_primitives: int, embedding_dim: int): super().__init__() self.primitive_embeddings = nn.Embedding(num_primitives, embedding_dim) self.embedding_dim = embedding_dim # Initialize with small random values nn.init.normal_(self.primitive_embeddings.weight, std=0.02) def forward(self, primitive_ids: torch.Tensor) -> torch.Tensor: """ Get embeddings for primitive IDs Args: primitive_ids: [batch_size, num_primitives] Returns: embeddings: [batch_size, num_primitives, embedding_dim] """ return self.primitive_embeddings(primitive_ids) class CAD2ProgramModel(PreTrainedModel, GenerationMixin): """ Main CAD2Program model combining vision encoder and language decoder """ config_class = CAD2ProgramConfig def __init__(self, config: CAD2ProgramConfig): super().__init__(config) self.config = config # Vision encoder (ViT) vision_config = ViTConfig( image_size=config.image_size, patch_size=config.patch_size, hidden_size=config.vision_hidden_size, num_attention_heads=config.num_attention_heads, num_hidden_layers=config.num_hidden_layers // 2, # Smaller vision model intermediate_size=config.intermediate_size, dropout_prob=config.dropout_prob ) self.vision_model = ViTModel(vision_config) # Language decoder (GPT-2 style) language_config = GPT2Config( vocab_size=config.vocab_size, n_embd=config.language_hidden_size, n_head=config.num_attention_heads, n_layer=config.num_hidden_layers, n_positions=config.max_position_embeddings, resid_pdrop=config.dropout_prob, attn_pdrop=config.dropout_prob ) self.language_model = GPT2LMHeadModel(language_config) # Vision-Language projection self.vision_projector = VisionLanguageProjector( vision_dim=config.vision_hidden_size, language_dim=config.language_hidden_size, hidden_dim=config.projector_hidden_size ) # Special primitive embeddings self.primitive_embeddings = PrimitiveEmbedding( num_primitives=config.num_primitive_types, embedding_dim=config.language_hidden_size ) # Image processor and tokenizer will be set during training/inference self.image_processor = None self.tokenizer = None # Post-initialization self.post_init() def get_vision_features(self, pixel_values: torch.Tensor) -> torch.Tensor: """ Extract features from images using vision encoder Args: pixel_values: [batch_size, channels, height, width] Returns: vision_features: [batch_size, num_patches + 1, vision_hidden_size] """ vision_outputs = self.vision_model(pixel_values=pixel_values) return vision_outputs.last_hidden_state def prepare_inputs_embeds( self, pixel_values: Optional[torch.Tensor] = None, input_ids: Optional[torch.Tensor] = None, vision_features: Optional[torch.Tensor] = None ) -> Tuple[torch.Tensor, torch.Tensor]: """ Prepare combined input embeddings from vision and text Args: pixel_values: [batch_size, channels, height, width] input_ids: [batch_size, sequence_length] vision_features: Pre-computed vision features Returns: inputs_embeds: [batch_size, total_sequence_length, hidden_size] attention_mask: [batch_size, total_sequence_length] """ batch_size = pixel_values.shape[0] if pixel_values is not None else input_ids.shape[0] # Get vision features if vision_features is None and pixel_values is not None: vision_features = self.get_vision_features(pixel_values) # Project vision features to language space if vision_features is not None: vision_embeds = self.vision_projector(vision_features) vision_seq_len = vision_embeds.shape[1] else: vision_embeds = None vision_seq_len = 0 # Get text embeddings if input_ids is not None: text_embeds = self.language_model.transformer.wte(input_ids) text_seq_len = text_embeds.shape[1] else: text_embeds = None text_seq_len = 0 # Combine embeddings if vision_embeds is not None and text_embeds is not None: # Concatenate vision and text embeddings inputs_embeds = torch.cat([vision_embeds, text_embeds], dim=1) # Create attention mask attention_mask = torch.ones( batch_size, vision_seq_len + text_seq_len, dtype=torch.long, device=inputs_embeds.device ) elif vision_embeds is not None: inputs_embeds = vision_embeds attention_mask = torch.ones( batch_size, vision_seq_len, dtype=torch.long, device=inputs_embeds.device ) elif text_embeds is not None: inputs_embeds = text_embeds attention_mask = torch.ones( batch_size, text_seq_len, dtype=torch.long, device=inputs_embeds.device ) else: raise ValueError("Either pixel_values or input_ids must be provided") return inputs_embeds, attention_mask def forward( self, pixel_values: Optional[torch.Tensor] = None, input_ids: Optional[torch.Tensor] = None, attention_mask: Optional[torch.Tensor] = None, labels: Optional[torch.Tensor] = None, past_key_values: Optional[Tuple[Tuple[torch.Tensor]]] = None, use_cache: Optional[bool] = None, output_attentions: Optional[bool] = None, output_hidden_states: Optional[bool] = None, return_dict: Optional[bool] = None, **kwargs ) -> Union[Tuple, CausalLMOutputWithPast]: """ Forward pass of the model """ return_dict = return_dict if return_dict is not None else self.config.use_return_dict # Prepare input embeddings if past_key_values is None: inputs_embeds, computed_attention_mask = self.prepare_inputs_embeds( pixel_values=pixel_values, input_ids=input_ids ) # Use computed attention mask if none provided if attention_mask is None: attention_mask = computed_attention_mask else: # For generation, only use text inputs inputs_embeds = None # Forward through language model outputs = self.language_model( input_ids=input_ids if inputs_embeds is None else None, inputs_embeds=inputs_embeds, attention_mask=attention_mask, labels=labels, past_key_values=past_key_values, use_cache=use_cache, output_attentions=output_attentions, output_hidden_states=output_hidden_states, return_dict=return_dict ) return outputs def generate_cad_program( self, images: Union[Image.Image, List[Image.Image], torch.Tensor], prompt: str = "Reconstruct cabinet from image:", max_new_tokens: int = 512, temperature: float = 0.7, top_p: float = 0.9, do_sample: bool = True, **kwargs ) -> str: """ Generate CAD program from input image(s) Args: images: Input CAD drawing image(s) prompt: Text prompt to guide generation max_new_tokens: Maximum tokens to generate temperature: Sampling temperature top_p: Top-p sampling do_sample: Whether to use sampling Returns: generated_program: Python CAD program as string """ self.eval() # Process images if isinstance(images, Image.Image): images = [images] if isinstance(images, list): if self.image_processor is None: raise ValueError("Image processor not set. Call set_image_processor() first.") pixel_values = self.image_processor(images, return_tensors="pt")["pixel_values"] else: pixel_values = images # Ensure on correct device pixel_values = pixel_values.to(self.device) # Tokenize prompt if self.tokenizer is None: raise ValueError("Tokenizer not set. Call set_tokenizer() first.") prompt_ids = self.tokenizer.encode(prompt, return_tensors="pt").to(self.device) # Get vision features with torch.no_grad(): vision_features = self.get_vision_features(pixel_values) vision_embeds = self.vision_projector(vision_features) # Prepare initial input embeddings with vision + prompt prompt_embeds = self.language_model.transformer.wte(prompt_ids) initial_embeds = torch.cat([vision_embeds, prompt_embeds], dim=1) # Generate output_ids = self.language_model.generate( inputs_embeds=initial_embeds, max_new_tokens=max_new_tokens, temperature=temperature, top_p=top_p, do_sample=do_sample, pad_token_id=self.tokenizer.eos_token_id, **kwargs ) # Decode only the generated part (skip vision and prompt tokens) vision_seq_len = vision_embeds.shape[1] prompt_len = prompt_ids.shape[1] generated_ids = output_ids[0, vision_seq_len + prompt_len:] generated_text = self.tokenizer.decode(generated_ids, skip_special_tokens=True) return generated_text.strip() def parse_program_to_assembly(self, program_text: str) -> CabinetAssembly: """ Parse generated program text to CAD assembly Args: program_text: Generated Python program Returns: assembly: CabinetAssembly object """ # This is a simplified parser - in production you'd want more robust parsing try: # Create assembly assembly = CabinetAssembly("Generated Cabinet") # Execute program in controlled environment exec_globals = { "PRIMITIVE_REGISTRY": PRIMITIVE_REGISTRY, "CabinetAssembly": CabinetAssembly, "assembly": assembly } # Add primitive classes to globals for prim_id in PRIMITIVE_REGISTRY.list_primitives(): prim_class = PRIMITIVE_REGISTRY.get_primitive_class(prim_id) exec_globals[prim_class.__name__] = prim_class # Execute program exec(program_text, exec_globals) return assembly except Exception as e: print(f"Error parsing program: {e}") return CabinetAssembly("Error Cabinet") def set_image_processor(self, image_processor): """Set the image processor""" self.image_processor = image_processor def set_tokenizer(self, tokenizer): """Set the tokenizer""" self.tokenizer = tokenizer def prepare_inputs_for_generation(self, input_ids, past_key_values=None, **kwargs): """Prepare inputs for generation""" # This is needed for the generation mixin if past_key_values is not None: input_ids = input_ids[:, -1:] return { "input_ids": input_ids, "past_key_values": past_key_values, "pixel_values": kwargs.get("pixel_values"), } def _reorder_cache(self, past_key_values, beam_idx): """Reorder cache for beam search""" return self.language_model._reorder_cache(past_key_values, beam_idx) class CAD2ProgramForTraining(CAD2ProgramModel): """ Training-specific version with additional loss components """ def __init__(self, config: CAD2ProgramConfig): super().__init__(config) # Additional heads for auxiliary losses self.primitive_classifier = nn.Linear( config.language_hidden_size, config.num_primitive_types ) # Position regression head self.position_regressor = nn.Linear( config.language_hidden_size, 3 # x, y, z coordinates ) # Size regression head self.size_regressor = nn.Linear( config.language_hidden_size, 3 # width, depth, height ) def compute_auxiliary_losses( self, hidden_states: torch.Tensor, primitive_labels: Optional[torch.Tensor] = None, position_labels: Optional[torch.Tensor] = None, size_labels: Optional[torch.Tensor] = None ) -> Dict[str, torch.Tensor]: """ Compute auxiliary losses for better training Args: hidden_states: [batch_size, seq_len, hidden_size] primitive_labels: [batch_size, num_primitives] - primitive type labels position_labels: [batch_size, num_primitives, 3] - position labels size_labels: [batch_size, num_primitives, 3] - size labels Returns: losses: Dictionary of auxiliary losses """ losses = {} # Use the mean of hidden states for classification/regression pooled_states = hidden_states.mean(dim=1) # [batch_size, hidden_size] if primitive_labels is not None: # Primitive classification loss primitive_logits = self.primitive_classifier(pooled_states) primitive_loss = F.cross_entropy( primitive_logits.view(-1, primitive_logits.size(-1)), primitive_labels.view(-1), ignore_index=-1 ) losses["primitive_loss"] = primitive_loss if position_labels is not None: # Position regression loss position_preds = self.position_regressor(pooled_states) position_loss = F.mse_loss( position_preds, position_labels.mean(dim=1) # Average across primitives ) losses["position_loss"] = position_loss if size_labels is not None: # Size regression loss size_preds = self.size_regressor(pooled_states) size_loss = F.mse_loss( size_preds, size_labels.mean(dim=1) # Average across primitives ) losses["size_loss"] = size_loss return losses def forward( self, pixel_values: Optional[torch.Tensor] = None, input_ids: Optional[torch.Tensor] = None, attention_mask: Optional[torch.Tensor] = None, labels: Optional[torch.Tensor] = None, primitive_labels: Optional[torch.Tensor] = None, position_labels: Optional[torch.Tensor] = None, size_labels: Optional[torch.Tensor] = None, **kwargs ) -> Union[Tuple, CausalLMOutputWithPast]: """ Forward pass with auxiliary losses for training """ # Main forward pass outputs = super().forward( pixel_values=pixel_values, input_ids=input_ids, attention_mask=attention_mask, labels=labels, **kwargs ) # Compute auxiliary losses if training if self.training and outputs.hidden_states is not None: aux_losses = self.compute_auxiliary_losses( hidden_states=outputs.hidden_states[-1], # Last layer hidden states primitive_labels=primitive_labels, position_labels=position_labels, size_labels=size_labels ) # Combine losses total_loss = outputs.loss if outputs.loss is not None else 0 for loss_name, loss_value in aux_losses.items(): total_loss = total_loss + 0.1 * loss_value # Weight auxiliary losses # Update output outputs.loss = total_loss outputs.auxiliary_losses = aux_losses return outputs # Model factory functions def create_model_from_config(config: CAD2ProgramConfig) -> CAD2ProgramModel: """Create model from configuration""" return CAD2ProgramModel(config) def create_training_model_from_config(config: CAD2ProgramConfig) -> CAD2ProgramForTraining: """Create training model from configuration""" return CAD2ProgramForTraining(config) def load_pretrained_model(model_path: str) -> CAD2ProgramModel: """Load a pretrained model""" return CAD2ProgramModel.from_pretrained(model_path) # Evaluation utilities class CAD2ProgramEvaluator: """Evaluation utilities for CAD2Program model""" def __init__(self, model: CAD2ProgramModel, tokenizer, image_processor): self.model = model self.tokenizer = tokenizer self.image_processor = image_processor # Set processors self.model.set_tokenizer(tokenizer) self.model.set_image_processor(image_processor) def evaluate_reconstruction_accuracy( self, test_images: List[Image.Image], ground_truth_programs: List[str], metrics: List[str] = ["bleu", "program_similarity", "geometric_accuracy"] ) -> Dict[str, float]: """ Evaluate model on reconstruction accuracy Args: test_images: List of test CAD drawings ground_truth_programs: List of ground truth Python programs metrics: List of metrics to compute Returns: results: Dictionary of metric scores """ results = {metric: 0.0 for metric in metrics} generated_programs = [] # Generate programs for all test images for image in test_images: try: program = self.model.generate_cad_program(image) generated_programs.append(program) except Exception as e: print(f"Generation failed: {e}") generated_programs.append("") # Compute metrics if "bleu" in metrics: results["bleu"] = self._compute_bleu_score( generated_programs, ground_truth_programs ) if "program_similarity" in metrics: results["program_similarity"] = self._compute_program_similarity( generated_programs, ground_truth_programs ) if "geometric_accuracy" in metrics: results["geometric_accuracy"] = self._compute_geometric_accuracy( generated_programs, ground_truth_programs ) return results def _compute_bleu_score(self, generated: List[str], ground_truth: List[str]) -> float: """Compute BLEU score between generated and ground truth programs""" try: from nltk.translate.bleu_score import corpus_bleu # Tokenize programs references = [[gt.split()] for gt in ground_truth] candidates = [gen.split() for gen in generated] return corpus_bleu(references, candidates) except ImportError: print("NLTK not available for BLEU computation") return 0.0 def _compute_program_similarity(self, generated: List[str], ground_truth: List[str]) -> float: """Compute semantic similarity between programs""" total_similarity = 0.0 count = 0 for gen, gt in zip(generated, ground_truth): try: # Parse both programs to assemblies gen_assembly = self.model.parse_program_to_assembly(gen) gt_assembly = self.model.parse_program_to_assembly(gt) # Compare primitive counts and types similarity = self._compare_assemblies(gen_assembly, gt_assembly) total_similarity += similarity count += 1 except: continue return total_similarity / count if count > 0 else 0.0 def _compute_geometric_accuracy(self, generated: List[str], ground_truth: List[str]) -> float: """Compute geometric accuracy of generated models""" total_accuracy = 0.0 count = 0 for gen, gt in zip(generated, ground_truth): try: # Parse to assemblies gen_assembly = self.model.parse_program_to_assembly(gen) gt_assembly = self.model.parse_program_to_assembly(gt) # Compare geometric properties accuracy = self._compare_geometry(gen_assembly, gt_assembly) total_accuracy += accuracy count += 1 except: continue return total_accuracy / count if count > 0 else 0.0 def _compare_assemblies(self, assembly1: CabinetAssembly, assembly2: CabinetAssembly) -> float: """Compare two cabinet assemblies for similarity""" if len(assembly1.primitives) == 0 and len(assembly2.primitives) == 0: return 1.0 if len(assembly1.primitives) == 0 or len(assembly2.primitives) == 0: return 0.0 # Compare primitive types types1 = [type(p).__name__ for p in assembly1.primitives] types2 = [type(p).__name__ for p in assembly2.primitives] # Jaccard similarity on primitive types set1, set2 = set(types1), set(types2) intersection = len(set1 & set2) union = len(set1 | set2) return intersection / union if union > 0 else 0.0 def _compare_geometry(self, assembly1: CabinetAssembly, assembly2: CabinetAssembly) -> float: """Compare geometric properties of assemblies""" if len(assembly1.primitives) != len(assembly2.primitives): return 0.0 # Compare overall dimensions dims1 = assembly1.get_dimensions() dims2 = assembly2.get_dimensions() # Compute relative error total_error = 0.0 for key in ["width", "depth", "height"]: if dims2[key] > 0: error = abs(dims1[key] - dims2[key]) / dims2[key] total_error += error # Convert error to accuracy (1.0 = perfect, 0.0 = completely wrong) accuracy = max(0.0, 1.0 - (total_error / 3.0)) return accuracy # Utility functions for model deployment def save_model_for_huggingface( model: CAD2ProgramModel, tokenizer, image_processor, save_directory: str, push_to_hub: bool = False, hub_model_id: str = None ): """ Save model in HuggingFace format Args: model: Trained CAD2Program model tokenizer: Associated tokenizer image_processor: Associated image processor save_directory: Local directory to save push_to_hub: Whether to push to HF Hub hub_model_id: Model ID for HF Hub """ import os # Create directory os.makedirs(save_directory, exist_ok=True) # Save model model.save_pretrained(save_directory) # Save tokenizer tokenizer.save_pretrained(save_directory) # Save image processor image_processor.save_pretrained(save_directory) # Save additional config additional_config = { "model_type": "cad2program", "task": "image-to-text", "tags": ["cad", "3d-reconstruction", "vision-language"], "license": "apache-2.0" } import json with open(os.path.join(save_directory, "additional_config.json"), "w") as f: json.dump(additional_config, f, indent=2) # Push to hub if requested if push_to_hub and hub_model_id: model.push_to_hub(hub_model_id) tokenizer.push_to_hub(hub_model_id) image_processor.push_to_hub(hub_model_id) # Example usage and testing if __name__ == "__main__": from model_config import create_default_config from transformers import GPT2Tokenizer, ViTImageProcessor # Create model config = create_default_config() model = create_model_from_config(config) # Set up tokenizer and image processor tokenizer = GPT2Tokenizer.from_pretrained("microsoft/DialoGPT-small") tokenizer.pad_token = tokenizer.eos_token image_processor = ViTImageProcessor.from_pretrained("google/vit-base-patch16-224") model.set_tokenizer(tokenizer) model.set_image_processor(image_processor) print(f"Model created with {sum(p.numel() for p in model.parameters())} parameters") print(f"Vision model: {config.vision_model_name}") print(f"Language model: {config.language_model_name}") print(f"Max primitives: {config.max_primitives}") print(f"Supported formats: {config.supported_formats}") # Test forward pass with dummy data batch_size = 2 pixel_values = torch.randn(batch_size, 3, 224, 224) input_ids = torch.randint(0, 1000, (batch_size, 50)) # Forward pass with torch.no_grad(): outputs = model(pixel_values=pixel_values, input_ids=input_ids) print(f"Output logits shape: {outputs.logits.shape}") # Test generation (would need a real image in practice) from PIL import Image dummy_image = Image.new('RGB', (224, 224), color='white') try: generated_program = model.generate_cad_program( dummy_image, max_new_tokens=100, temperature=0.8 ) print(f"Generated program: {generated_program}") except Exception as e: print(f"Generation test failed (expected with dummy data): {e}") print("Model implementation complete!")