import torch as t import argparse from sae.SAE_Tools import load_sae_model from typing import List, Tuple, Dict, Any, Union, Callable, cast, Literal import os from PIL import Image def load_image_and_text_from_folder(folder_path: str) -> Tuple[Image.Image, str]: # Load the text file text_file_path = os.path.join(folder_path, "caption.txt") with open(text_file_path, "r") as file: text_prompt = file.read() # Find the image file (assuming only one image file exists in the folder) image_file = next( (f for f in os.listdir(folder_path) if f.lower().endswith((".jpg", ".png", ".jpeg"))), None ) if image_file: image_path = os.path.join(folder_path, image_file) image = Image.open(image_path).convert("RGB") else: raise FileNotFoundError("No image file found in the folder.") return image, text_prompt def str_to_bool(value): if isinstance(value, bool): return value if value.lower() in ('yes', 'true', 't', 'y', '1'): return True elif value.lower() in ('no', 'false', 'f', 'n', '0'): return False else: raise argparse.ArgumentTypeError('Boolean value expected.') def load_saes( model: Any, model_type: Literal["llava", "blip"], text_sae_types: List[str], vis_sae_types: List[str], device: t.device | str = "cpu", dtype: t.dtype = t.float32, ): text_sae_paths = [] for sae_type in text_sae_types: if sae_type == "attn": text_sae_paths.append( # "checkpoints/topk_16.0_32_text_decoder.bert.encoder.layer.{layer}.attention.self.hook_resid_pre_0.001_256_0.0_42.ckpt" "cc3m_checkpoints/topk_32.0_32_text_decoder.bert.encoder.layer.{layer}.attention.self.hook_resid_pre_0.001_256_0.03125_42.ckpt" ) elif sae_type == "crossattn": text_sae_paths.append( # "checkpoints/topk_16.0_32_text_decoder.bert.encoder.layer.{layer}.crossattention.self.hook_resid_pre_0.001_256_0.0_42.ckpt" "cc3m_checkpoints/topk_32.0_32_text_decoder.bert.encoder.layer.{layer}.crossattention.self.hook_resid_pre_0.001_256_0.03125_42.ckpt" ) else: raise vision_sae_paths = [] for sae_type in vis_sae_types: if sae_type == "pre": vision_sae_paths.append( "checkpoints/topk_16.0_32_vision_model.encoder.layers.{layer}.hook_resid_pre_0.001_256_0.0_42.ckpt" ) elif sae_type == "post": vision_sae_paths.append( "cc3m_checkpoints/topk_32.0_32_vision_model.encoder.layers.{layer}.hook_resid_post_0.001_256_0.03125_42.ckpt" ) else: raise text_saes = {} # layer: list[str, sae] for path in text_sae_paths: for layer in range(model.cfg.n_layers): if layer not in text_saes: text_saes[layer] = [] sae = load_sae_model( file_path=path.format(layer=layer), model_type=model_type, hook_type="text", device=device, ).to(dtype) text_saes[layer].append((sae.cfg.hook_name, sae)) vision_saes = {} # layer: list[str, sae] for path in vision_sae_paths: for layer in range(model.cfg.n_layers): if layer not in vision_saes: vision_saes[layer] = [] sae = load_sae_model( file_path=path.format(layer=layer), model_type=model_type, hook_type="vision", device=device, ).to(dtype) vision_saes[layer].append((sae.cfg.hook_name, sae)) return text_saes, vision_saes def str_to_dtype(dtype_str: str) -> t.dtype: if dtype_str == "float32": return t.float32 elif dtype_str == "float16": return t.float16 elif dtype_str == "bfloat16": return t.bfloat16 else: raise ValueError(f"Unsupported dtype string: {dtype_str}")