cad2program / modeling_cad2program.py
turiya-ai's picture
Update modeling_cad2program.py
cffc3bf verified
Raw
History Blame Contribute Delete
28.6 kB
"""
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!")