diff --git a/bioscanclip/config/global_config.yaml b/bioscanclip/config/global_config.yaml index 480bafe..186035c 100644 --- a/bioscanclip/config/global_config.yaml +++ b/bioscanclip/config/global_config.yaml @@ -33,6 +33,8 @@ insect_data: species_to_other: ${insect_data.dir}/specie_to_other_labels.json save_ckpt: true bioscan_bert_checkpoint: ${project_root_path}/ckpt/BarcodeBERT/5_mer/model_41.pth +bioscan_bert_checkpoint_trained_with_canada_1_5_m: ${project_root_path}/ckpt/BarcodeBERT/new_checkpoints/trained_with_canada_1_5M/CANADA-1.5M-BEST_k4_4_4_w1_m0_r0_wd.pt +bioscan_bert_checkpoint_trained_with_bioscan_5_m: ${project_root_path}/ckpt/BarcodeBERT/new_checkpoints/trained_with_5m/BIOSCAN-5M-BEST_k4_6_6_w1_m0_r0.pt model_output_dir: ${project_root_path}/ckpt/bioscan_clip inference_and_eval_setting: @@ -65,3 +67,8 @@ general_fine_tune_setting: hf_repo_id: bioscan-ml/clibd default_seed: 42 + +barcodebert_setting: + old_model_setting: + k: 5 + max_len: 660 diff --git a/bioscanclip/config/model_config/for_bioscan_1m/debug_config_for_softmax_issue/image_dna_text_seed_42.yaml b/bioscanclip/config/model_config/for_bioscan_1m/debug_config_for_softmax_issue/image_dna_text_seed_42.yaml new file mode 100644 index 0000000..f0824fd --- /dev/null +++ b/bioscanclip/config/model_config/for_bioscan_1m/debug_config_for_softmax_issue/image_dna_text_seed_42.yaml @@ -0,0 +1,47 @@ +batch_size: 200 +epochs: 30 +wandb_project_name: BIOSCAN-CLIP_softmax_issue +using_train_seen_for_pre_train: true +dataset: bioscan_1m + +image: + input_type: image + model: vit +dna: + input_type: sequence + model: barcode_bert +language: + input_type: sequence + model: bert_small + +model_output_name: image_dna_text_4gpu_softmax_issue +evaluation_period: 1 +ckpt_path: ${project_root_path}/ckpt/bioscan_clip/image_dna_text_4gpu_softmax_issue/best.pth +output_dim: 768 +port: 29531 + +disable_lora: true +lr_scheduler: one_cycle +lr_config: + lr: 1e-6 + max_lr: 5e-5 + + +all_gather: true +loss_setup: + gather_with_grad: true + use_horovod: false + local_loss: false +fix_temperature: false +amp: true + +random_seed: false + +eval_skip_epoch: 23 + +default_seed: 42 + +fine_tuning_set: + batch_size: 150 + epochs: 15 + fine_tune_model_output_dir: ${model_output_dir}/${model_config.model_output_name}/supervise_fine_tune_ckpt \ No newline at end of file diff --git a/bioscanclip/config/model_config/for_bioscan_1m/final_experiments/image_dna_seed_42.yaml b/bioscanclip/config/model_config/for_bioscan_1m/final_experiments/image_dna_seed_42.yaml index 6064945..62647e4 100644 --- a/bioscanclip/config/model_config/for_bioscan_1m/final_experiments/image_dna_seed_42.yaml +++ b/bioscanclip/config/model_config/for_bioscan_1m/final_experiments/image_dna_seed_42.yaml @@ -1,5 +1,5 @@ batch_size: 500 -epochs: 30 +epochs: 50 labels_for_driven_positive_and_negative_pairs: wandb_project_name: BIOSCAN-CLIP using_train_seen_for_pre_train: true @@ -7,10 +7,10 @@ dataset: bioscan_1m image: input_type: image - model: lora_vit + model: vit dna: input_type: sequence - model: lora_barcode_bert + model: barcode_bert model_output_name: image_dna_4gpu evaluation_period: 1 diff --git a/bioscanclip/config/model_config/for_bioscan_1m/final_experiments/image_dna_text_no_loading.yaml b/bioscanclip/config/model_config/for_bioscan_1m/final_experiments/image_dna_text_no_loading.yaml index 82e363e..5e46229 100644 --- a/bioscanclip/config/model_config/for_bioscan_1m/final_experiments/image_dna_text_no_loading.yaml +++ b/bioscanclip/config/model_config/for_bioscan_1m/final_experiments/image_dna_text_no_loading.yaml @@ -1,5 +1,5 @@ batch_size: 500 -epochs: 30 +epochs: 50 labels_for_driven_positive_and_negative_pairs: wandb_project_name: BIOSCAN-CLIP using_train_seen_for_pre_train: true @@ -7,13 +7,13 @@ dataset: bioscan_1m image: input_type: image - model: lora_vit + model: vit dna: input_type: sequence - model: lora_barcode_bert + model: barcode_bert language: input_type: sequence - model: lora_bert + model: bert_small load_ckpt: false model_output_name: no_align_baseline diff --git a/bioscanclip/config/model_config/for_bioscan_1m/final_experiments/image_dna_text_seed_42.yaml b/bioscanclip/config/model_config/for_bioscan_1m/final_experiments/image_dna_text_seed_42.yaml index ada525d..534bd1b 100644 --- a/bioscanclip/config/model_config/for_bioscan_1m/final_experiments/image_dna_text_seed_42.yaml +++ b/bioscanclip/config/model_config/for_bioscan_1m/final_experiments/image_dna_text_seed_42.yaml @@ -1,5 +1,5 @@ batch_size: 500 -epochs: 30 +epochs: 50 wandb_project_name: BIOSCAN-CLIP using_train_seen_for_pre_train: true dataset: bioscan_1m diff --git a/bioscanclip/config/model_config/for_bioscan_1m/final_experiments/image_dna_text_seed_42_load_pre_trained_image_encoder_trained_with_simclr_style.yaml b/bioscanclip/config/model_config/for_bioscan_1m/final_experiments/image_dna_text_seed_42_load_pre_trained_image_encoder_trained_with_simclr_style.yaml index 7e76a6f..efe204c 100644 --- a/bioscanclip/config/model_config/for_bioscan_1m/final_experiments/image_dna_text_seed_42_load_pre_trained_image_encoder_trained_with_simclr_style.yaml +++ b/bioscanclip/config/model_config/for_bioscan_1m/final_experiments/image_dna_text_seed_42_load_pre_trained_image_encoder_trained_with_simclr_style.yaml @@ -1,5 +1,5 @@ batch_size: 500 -epochs: 30 +epochs: 50 labels_for_driven_positive_and_negative_pairs: wandb_project_name: BIOSCAN-CLIP using_train_seen_for_pre_train: true diff --git a/bioscanclip/config/model_config/for_bioscan_1m/final_experiments/image_dna_text_seed_42_old.yaml b/bioscanclip/config/model_config/for_bioscan_1m/final_experiments/image_dna_text_seed_42_old.yaml index 82e363e..80f569e 100644 --- a/bioscanclip/config/model_config/for_bioscan_1m/final_experiments/image_dna_text_seed_42_old.yaml +++ b/bioscanclip/config/model_config/for_bioscan_1m/final_experiments/image_dna_text_seed_42_old.yaml @@ -1,5 +1,5 @@ batch_size: 500 -epochs: 30 +epochs: 50 labels_for_driven_positive_and_negative_pairs: wandb_project_name: BIOSCAN-CLIP using_train_seen_for_pre_train: true diff --git a/bioscanclip/config/model_config/for_bioscan_1m/final_experiments/image_text_seed_42.yaml b/bioscanclip/config/model_config/for_bioscan_1m/final_experiments/image_text_seed_42.yaml index 0bf49c0..04879c0 100644 --- a/bioscanclip/config/model_config/for_bioscan_1m/final_experiments/image_text_seed_42.yaml +++ b/bioscanclip/config/model_config/for_bioscan_1m/final_experiments/image_text_seed_42.yaml @@ -1,16 +1,15 @@ batch_size: 500 -epochs: 30 -labels_for_driven_positive_and_negative_pairs: +epochs: 50 wandb_project_name: BIOSCAN-CLIP using_train_seen_for_pre_train: true dataset: bioscan_1m image: input_type: image - model: lora_vit + model: vit language: input_type: sequence - model: lora_bert + model: bert_small model_output_name: image_text_4gpu evaluation_period: 1 diff --git a/bioscanclip/config/model_config/for_bioscan_5m/barcodeBERT_trained_with_5m/image_dna_text_seed_42_new_barcodeBERT_1M.yaml b/bioscanclip/config/model_config/for_bioscan_5m/barcodeBERT_trained_with_5m/image_dna_text_seed_42_new_barcodeBERT_1M.yaml index c82becf..5ee61c8 100644 --- a/bioscanclip/config/model_config/for_bioscan_5m/barcodeBERT_trained_with_5m/image_dna_text_seed_42_new_barcodeBERT_1M.yaml +++ b/bioscanclip/config/model_config/for_bioscan_5m/barcodeBERT_trained_with_5m/image_dna_text_seed_42_new_barcodeBERT_1M.yaml @@ -1,4 +1,4 @@ -batch_size: 300 +batch_size: 200 epochs: 15 labels_for_driven_positive_and_negative_pairs: wandb_project_name: BIOSCAN-CLIP-5M @@ -16,7 +16,7 @@ language: pre_train_model: prajjwal1/bert-small model_output_name: image_dna_text_4gpu -evaluation_period: 5 +evaluation_period: 3 output_dim: 768 port: 29531 diff --git a/bioscanclip/config/model_config/for_bioscan_5m/barcodeBERT_trained_with_5m/image_dna_text_seed_42_new_barcodeBERT_5M.yaml b/bioscanclip/config/model_config/for_bioscan_5m/barcodeBERT_trained_with_5m/image_dna_text_seed_42_new_barcodeBERT_5M.yaml index d1ceebb..bafaec3 100644 --- a/bioscanclip/config/model_config/for_bioscan_5m/barcodeBERT_trained_with_5m/image_dna_text_seed_42_new_barcodeBERT_5M.yaml +++ b/bioscanclip/config/model_config/for_bioscan_5m/barcodeBERT_trained_with_5m/image_dna_text_seed_42_new_barcodeBERT_5M.yaml @@ -1,5 +1,5 @@ batch_size: 200 -epochs: 10 +epochs: 15 labels_for_driven_positive_and_negative_pairs: wandb_project_name: BIOSCAN-CLIP-5M using_train_seen_for_pre_train: true diff --git a/bioscanclip/config/model_config/for_bioscan_5m/barcodeBERT_trained_with_5m/image_dna_text_seed_42_old_barcodeBERT.yaml b/bioscanclip/config/model_config/for_bioscan_5m/barcodeBERT_trained_with_5m/image_dna_text_seed_42_old_barcodeBERT.yaml index a0fefdb..68bf2eb 100644 --- a/bioscanclip/config/model_config/for_bioscan_5m/barcodeBERT_trained_with_5m/image_dna_text_seed_42_old_barcodeBERT.yaml +++ b/bioscanclip/config/model_config/for_bioscan_5m/barcodeBERT_trained_with_5m/image_dna_text_seed_42_old_barcodeBERT.yaml @@ -1,5 +1,5 @@ batch_size: 200 -epochs: 10 +epochs: 15 labels_for_driven_positive_and_negative_pairs: wandb_project_name: BIOSCAN-CLIP-5M using_train_seen_for_pre_train: true diff --git a/bioscanclip/config/model_config/for_bioscan_5m/final_experiments/image_dna_seed_42.yaml b/bioscanclip/config/model_config/for_bioscan_5m/final_experiments/image_dna_seed_42.yaml index d1b09e9..cc6ac73 100644 --- a/bioscanclip/config/model_config/for_bioscan_5m/final_experiments/image_dna_seed_42.yaml +++ b/bioscanclip/config/model_config/for_bioscan_5m/final_experiments/image_dna_seed_42.yaml @@ -7,10 +7,10 @@ dataset: bioscan_5m image: input_type: image - pre_train_model: vit_base_patch16_224 + model: vit dna: input_type: sequence - pre_train_model: barcode_bert + model: barcode_bert model_output_name: image_dna_4gpu evaluation_period: 1 diff --git a/bioscanclip/config/model_config/for_bioscan_5m/final_experiments/image_dna_text_seed_42.yaml b/bioscanclip/config/model_config/for_bioscan_5m/final_experiments/image_dna_text_seed_42.yaml index 215e932..cc05174 100644 --- a/bioscanclip/config/model_config/for_bioscan_5m/final_experiments/image_dna_text_seed_42.yaml +++ b/bioscanclip/config/model_config/for_bioscan_5m/final_experiments/image_dna_text_seed_42.yaml @@ -7,13 +7,13 @@ dataset: bioscan_5m image: input_type: image - pre_train_model: vit_base_patch16_224 + model: vit dna: input_type: sequence - pre_train_model: barcode_bert + model: barcode_bert language: input_type: sequence - pre_train_model: prajjwal1/bert-small + model: bert_small model_output_name: image_dna_text_4gpu evaluation_period: 1 diff --git a/bioscanclip/config/model_config/for_bioscan_5m/final_experiments/image_text_seed_42.yaml b/bioscanclip/config/model_config/for_bioscan_5m/final_experiments/image_text_seed_42.yaml index 8329180..ae6c8ed 100644 --- a/bioscanclip/config/model_config/for_bioscan_5m/final_experiments/image_text_seed_42.yaml +++ b/bioscanclip/config/model_config/for_bioscan_5m/final_experiments/image_text_seed_42.yaml @@ -7,10 +7,10 @@ dataset: bioscan_5m image: input_type: image - pre_train_model: vit_base_patch16_224 + model: vit language: input_type: sequence - pre_train_model: prajjwal1/bert-small + model: bert_small model_output_name: image_text_4gpu evaluation_period: 1 diff --git a/bioscanclip/config/model_config/for_bioscan_5m/final_experiments/temp_config.yaml b/bioscanclip/config/model_config/for_bioscan_5m/final_experiments/temp_config.yaml index 3fa1507..37f2fe0 100644 --- a/bioscanclip/config/model_config/for_bioscan_5m/final_experiments/temp_config.yaml +++ b/bioscanclip/config/model_config/for_bioscan_5m/final_experiments/temp_config.yaml @@ -7,13 +7,13 @@ dataset: bioscan_5m image: input_type: image - pre_train_model: vit_base_patch16_224 + model: vit_base_patch16_224 dna: input_type: sequence - pre_train_model: barcode_bert + model: barcode_bert language: input_type: sequence - pre_train_model: prajjwal1/bert-small + model: prajjwal1/bert-small model_output_name: image_dna_text_4gpu-testing_tokenizer evaluation_period: 1 diff --git a/bioscanclip/epoch/fine_tuning_epoch.py b/bioscanclip/epoch/fine_tuning_epoch.py index defbbae..82f540a 100644 --- a/bioscanclip/epoch/fine_tuning_epoch.py +++ b/bioscanclip/epoch/fine_tuning_epoch.py @@ -3,6 +3,8 @@ import torch import numpy as np +# TODO: these functions either need to be updated or removed. + def label_batch_to_species_idx(label_batch, unique_species_for_seen): species_list = label_batch['species'] target = torch.tensor([unique_species_for_seen.index(species) for species in species_list]) diff --git a/bioscanclip/epoch/inference_epoch.py b/bioscanclip/epoch/inference_epoch.py index 4164106..8fea570 100644 --- a/bioscanclip/epoch/inference_epoch.py +++ b/bioscanclip/epoch/inference_epoch.py @@ -3,7 +3,9 @@ import torch.nn.functional as F import torch from transformers import AutoTokenizer - +import matplotlib as plt +import seaborn as sns +from sklearn.metrics import confusion_matrix def convert_label_dict_to_list_of_dict(label_batch): order = label_batch['order'] @@ -51,7 +53,6 @@ def get_feature_and_label(dataloader, model, device, for_open_clip=False, multi_ label_list = [] file_name_list =[] - tokenizer = AutoTokenizer.from_pretrained("bioscan-ml/BarcodeBERT", trust_remote_code=True) # Load tokenizer pbar = tqdm(enumerate(dataloader), total=len(dataloader)) model.eval() with torch.no_grad(): @@ -60,28 +61,16 @@ def get_feature_and_label(dataloader, model, device, for_open_clip=False, multi_ processid_batch, image_input_batch, dna_input_batch, input_ids, token_type_ids, attention_mask, label_batch = batch if for_open_clip: - language_input = input_ids - else: - language_input = {'input_ids': input_ids.to(device), 'token_type_ids': token_type_ids.to(device), - 'attention_mask': attention_mask.to(device)} - - if isinstance(dna_input_batch, torch.Tensor): - dna_input_batch = dna_input_batch.to(device) + language_input_batch = input_ids else: - # Tokenizing DNA sequences - tokenized_dna_sequences = [] - for dna_seq in dna_input_batch: - tokenized_output = tokenizer(dna_seq, padding='max_length', truncation=True, max_length=133, return_tensors="pt") - input_seq = tokenized_output["input_ids"] - tokenized_dna_sequences.append(input_seq) - # Convert DNA tokenized sequences into tensors - dna_input_batch = torch.stack(tokenized_dna_sequences).squeeze(1).to(device) + language_input_batch = {'input_ids': input_ids.to(device), 'token_type_ids': token_type_ids.to(device), + 'attention_mask': attention_mask.to(device)} # Forward pass through model image_output, dna_output, language_output, logit_scale, logit_bias = model( image_input_batch.to(device), - dna_input_batch, # Passing tokenized DNA sequences - language_input + dna_input_batch, + language_input_batch ) # Normalizing and storing outputs diff --git a/bioscanclip/epoch/train_epoch.py b/bioscanclip/epoch/train_epoch.py index 46ff071..77e3a46 100644 --- a/bioscanclip/epoch/train_epoch.py +++ b/bioscanclip/epoch/train_epoch.py @@ -15,7 +15,6 @@ def train_epoch(activate_wandb, total_epochs, epoch, dataloader, model, optimize pbar = enumerate(dataloader) epoch_loss = 0.0 total_step = len(dataloader) - tokenizer = AutoTokenizer.from_pretrained("bioscan-ml/BarcodeBERT", trust_remote_code=True) model.train() stop_flag = False for step, batch in pbar: @@ -27,17 +26,8 @@ def train_epoch(activate_wandb, total_epochs, epoch, dataloader, model, optimize 'attention_mask': attention_mask.to(device)} optimizer.zero_grad() image_input_batch = image_input_batch.to(device) + dna_input_batch = dna_input_batch - if isinstance(dna_input_batch, torch.Tensor): - dna_input_batch = dna_input_batch.to(device) - # if dna_input_batch is not a tensor, tokenize it - else: - tokenized_dna_sequences = [] - for dna_seq in dna_input_batch: - tokenized_output = tokenizer(dna_seq, padding='max_length', truncation=True, max_length=133, return_tensors="pt") - input_seq = tokenized_output["input_ids"] - tokenized_dna_sequences.append(input_seq) - dna_input_batch = torch.stack(tokenized_dna_sequences).squeeze(1).to(device) if enable_autocast: with torch.autocast(device_type='cuda', dtype=torch.bfloat16): @@ -78,4 +68,5 @@ def train_epoch(activate_wandb, total_epochs, epoch, dataloader, model, optimize if activate_wandb: wandb.log({"loss": loss.item(), "step": step + epoch * len(dataloader), "learning_rate": current_lr}) + print(f'Epoch [{epoch}/{total_epochs}], Loss: {epoch_loss / len(dataloader)}') \ No newline at end of file diff --git a/bioscanclip/model/dna_encoder.py b/bioscanclip/model/dna_encoder.py index 0c3fd7d..59b36e9 100644 --- a/bioscanclip/model/dna_encoder.py +++ b/bioscanclip/model/dna_encoder.py @@ -1,13 +1,14 @@ import math from itertools import product +import functools import torch import torch.nn as nn from torch import Tensor from torchtext.vocab import build_vocab_from_iterator from transformers import BertConfig, BertForMaskedLM -from bioscanclip.util.util import PadSequence, KmerTokenizer, load_bert_model, remove_extra_pre_fix - +from bioscanclip.util.util import remove_extra_pre_fix +from transformers import AutoTokenizer device = "cuda" if torch.cuda.is_available() else "cpu" @@ -19,7 +20,6 @@ def load_pre_trained_bioscan_bert(bioscan_bert_checkpoint, k=5): bioscan_bert_checkpoint: Path to checkpoint file k: k-mer size (default: 5) """ - print(f"\nLoading model from {bioscan_bert_checkpoint}") # Build k-mer vocabulary kmer_iter = (["".join(kmer)] for kmer in product("ACGT", repeat=k)) @@ -49,6 +49,104 @@ def load_pre_trained_bioscan_bert(bioscan_bert_checkpoint, k=5): model.load_state_dict(model_ckpt, strict=False) return model.to(device) +class NewKmerTokenizer(object): + def __init__(self, model_name="bioscan-ml/BarcodeBERT", max_length=660): + self.tokenizer = AutoTokenizer.from_pretrained(model_name, trust_remote_code=True) + self.max_length = max_length + + def __call__(self, text): + return self.tokenizer( + text, + padding="max_length", + truncation=True, + max_length=self.max_length + ) +class KmerTokenizerWithAttMask(object): + def __init__(self, k: int = 5, max_len: int = 660, stride: int = None): + """ + A tokenizer that: + 1) pads/truncates DNA to max_len with 'N' + 2) splits into k-mers (stride defaults to k for non-overlap) + 3) builds a vocab over all ACGT k-mers + specials + 4) produces input_ids, attention_mask, token_type_ids + """ + self.k = k + self.max_len = max_len + self.stride = stride or k + + if stride is None: + self.stride = k + else: + self.stride = stride + + # build vocab once + # kmer_iter = ("".join(kmer) for kmer in product("ACGT", repeat=k)) + # specials = ["", "", ""] + # vocab = build_vocab_from_iterator(kmer_iter, specials=specials) + # vocab.set_default_index(vocab[""]) + + kmer_iter = (["".join(kmer)] for kmer in product("ACGT", repeat=k)) + vocab = build_vocab_from_iterator(kmer_iter, specials=["", "", ""]) + vocab.set_default_index(vocab[""]) + + self.vocab = vocab + + def __call__(self, dna_sequence: str): + # 1) pad or truncate to max_len + if len(dna_sequence) >= self.max_len: + seq = dna_sequence[: self.max_len] + else: + seq = dna_sequence + "N" * (self.max_len - len(dna_sequence)) + + # 2) split into k-mers + tokens = [ + seq[i : i + self.k] + for i in range(0, len(seq) - self.k + 1, self.stride) + ] + + + # 3) attention mask: 1 for any k-mer containing A/C/G/T, 0 if all 'N' + attention_mask = [ + 0 if any(base == "N" for base in token) else 1 + for token in tokens + ] + # 4) convert tokens to IDs + input_ids = [self.vocab[token] for token in tokens] + input_ids = [0, *input_ids] # add CLS token at the beginning + # 4.5) Fix attention mask + attention_mask = [1, *attention_mask] + + # 5) single-segment token types (all zeros) + token_type_ids = [0] * len(input_ids) + + return { + "input_ids": input_ids, + "attention_mask": attention_mask, + "token_type_ids": token_type_ids + } + +# Old kmer tokenizer +class KmerTokenizer(object): + def __init__(self, k, stride=1): + self.k = k + self.stride = stride + + def __call__(self, dna_sequence): + tokens = [] + for i in range(0, len(dna_sequence) - self.k + 1, self.stride): + k_mer = dna_sequence[i : i + self.k] + tokens.append(k_mer) + return tokens + +class PadSequence(object): + def __init__(self, max_len): + self.max_len = max_len + + def __call__(self, dna_sequence): + if len(dna_sequence) > self.max_len: + return dna_sequence[: self.max_len] + else: + return dna_sequence + "N" * (self.max_len - len(dna_sequence)) def get_sequence_pipeline(k=5): kmer_iter = (["".join(kmer)] for kmer in product("ACGT", repeat=k)) @@ -128,13 +226,27 @@ def reset_parameters(self) -> None: for w_B in self.w_Bs: nn.init.zeros_(w_B.weight) - def forward(self, sequence) -> Tensor: - """ - TODO: change to "return self.base_dna_encoder(x).hidden_states[-1].mean(dim=1)" - TODO: Then also retrain the models. - """ - - return self.base_dna_encoder(sequence).logits.softmax(dim=-1).mean(dim=1) + def forward(self, input) -> Tensor: + input_ids = input["input_ids"].to(device) + attention_mask = input["attention_mask"].to(device) + token_type_ids = input["token_type_ids"].to(device) + outputs = self.base_dna_encoder( + input_ids=input_ids, + attention_mask=attention_mask, + token_type_ids=token_type_ids, + output_hidden_states=True + ) + + hidden_states = outputs.hidden_states[-1] + if attention_mask is not None: + mask = attention_mask.unsqueeze(-1) + masked_h = hidden_states * mask + sum_h = masked_h.sum(dim=1) + lengths = mask.sum(dim=1) + mean_h = sum_h / lengths + return mean_h + else: + return hidden_states.mean(dim=1) class Freeze_DNA_Encoder(nn.Module): def __init__(self): @@ -142,3 +254,240 @@ def __init__(self): def forward(self, x: Tensor) -> Tensor: return x + + """ + """ + + def __init__(self, base_tokenizer=None, batch_size=128): + self.tokenizer = base_tokenizer + self.batch_size = batch_size + + self.k = base_tokenizer.k + self.stride = base_tokenizer.stride + self.max_len = base_tokenizer.max_len + self.vocab_dict = base_tokenizer.vocab_dict + self.unk_token_id = self.vocab_dict.get("[UNK]", 1) + + self.process_count = 0 + self.batch_count = 0 + + self.cache = {} + self.cache_hits = 0 + self.cache_misses = 0 + + self.batch_cache = {} + self.batch_cache_hits = 0 + self.batch_cache_misses = 0 + + @functools.lru_cache(maxsize=16384) + def _tokenize_subsequence(self, seq): + if len(seq) < self.k: + return [] + + tokens = [seq[i:i + self.k] for i in range(0, len(seq) - self.k + 1, self.stride)] + + ids = [self.vocab_dict.get(token, self.unk_token_id) for token in tokens] + + return ids + + def _create_attention_mask_and_token_types(self, ids): + attention_mask = [1] * len(ids) + token_type_ids = [0] * len(ids) + return attention_mask, token_type_ids + + def batch_tokenize(self, texts, padding=False): + batch_tokens = [] + batch_ids = [] + batch_attention_masks = [] + batch_token_type_ids = [] + + for text in texts: + if len(text) > self.max_len: + text = text[:self.max_len] + if padding: + if len(text) < self.max_len: + text = text + 'N' * (self.max_len - len(text)) + + cache_key = (text, padding) + + if cache_key in self.cache: + self.cache_hits += 1 + cached_result = self.cache[cache_key] + batch_ids.append(cached_result["ids"]) + batch_attention_masks.append(cached_result["attention_mask"]) + batch_token_type_ids.append(cached_result["token_type_ids"]) + continue + + self.cache_misses += 1 + + if len(text) <= 32: + tokens = [text[i:i + self.k] for i in range(0, len(text) - self.k + 1, self.stride)] + ids = [self.vocab_dict.get(token, self.unk_token_id) for token in tokens] + else: + ids = [] + pos = 0 + + while pos < len(text) - self.k + 1: + chunk_size = min(32, len(text) - pos) + chunk = text[pos:pos+chunk_size] + + chunk_ids = self._tokenize_subsequence(chunk) + ids.extend(chunk_ids) + + pos += chunk_size + + attention_mask, token_type_ids = self._create_attention_mask_and_token_types(ids) + + batch_ids.append(ids) + batch_attention_masks.append(attention_mask) + batch_token_type_ids.append(token_type_ids) + + self.cache[cache_key] = { + "ids": ids, + "attention_mask": attention_mask, + "token_type_ids": token_type_ids + } + + if len(self.cache) > 100000: + keys_to_remove = list(self.cache.keys())[:20000] + for k in keys_to_remove: + del self.cache[k] + + return batch_ids, batch_attention_masks, batch_token_type_ids + + def __call__(self, texts, padding=False, truncation=True, max_length=660, return_tensors="pt"): + import torch + + single_input = False + if isinstance(texts, str): + texts = [texts] + single_input = True + + self.process_count += len(texts) + self.batch_count += 1 + + cache_key = (tuple(texts), padding, truncation, max_length, return_tensors) + if len(texts) <= 10 and cache_key in self.batch_cache: + self.batch_cache_hits += 1 + return self.batch_cache[cache_key] + + self.batch_cache_misses += 1 + + batch_ids, batch_attention_masks, batch_token_type_ids = self.batch_tokenize(texts, padding=padding) + + if return_tensors == "pt": + if padding == "max_length": + max_len = max_length + + padded_ids = [] + padded_attention_masks = [] + padded_token_type_ids = [] + + for ids, mask, type_ids in zip(batch_ids, batch_attention_masks, batch_token_type_ids): + if truncation and len(ids) > max_len: + ids = ids[:max_len] + mask = mask[:max_len] + type_ids = type_ids[:max_len] + + padding_length = max_len - len(ids) + if padding_length > 0: + ids = ids + [0] * padding_length + mask = mask + [0] * padding_length + type_ids = type_ids + [0] * padding_length + + padded_ids.append(ids) + padded_attention_masks.append(mask) + padded_token_type_ids.append(type_ids) + + batch_ids = padded_ids + batch_attention_masks = padded_attention_masks + batch_token_type_ids = padded_token_type_ids + + batch_ids = torch.tensor(batch_ids) + batch_attention_masks = torch.tensor(batch_attention_masks) + batch_token_type_ids = torch.tensor(batch_token_type_ids) + + result = { + "input_ids": batch_ids, + "attention_mask": batch_attention_masks, + "token_type_ids": batch_token_type_ids + } + + if len(texts) <= 10: + self.batch_cache[cache_key] = result + + if len(self.batch_cache) > 1000: + keys_to_remove = list(self.batch_cache.keys())[:200] + for k in keys_to_remove: + del self.batch_cache[k] + + if single_input and return_tensors == "pt": + return { + "input_ids": result["input_ids"][0].unsqueeze(0), + "attention_mask": result["attention_mask"][0].unsqueeze(0), + "token_type_ids": result["token_type_ids"][0].unsqueeze(0) + } + + return result + + def process_large_dataset(self, dna_input, padding="max_length", truncation=True, max_length=660, return_tensors="pt"): + import torch + from tqdm import tqdm + + all_tokenized_sequences = [] + + for i in tqdm(range(0, len(dna_input), self.batch_size), desc="Tokenizing DNA sequences in batches (with cascade caching)"): + batch = dna_input[i:i+self.batch_size] + + batch_results = self( + batch, + padding=padding, + truncation=truncation, + max_length=max_length, + return_tensors=return_tensors + ) + + if return_tensors == "pt": + for j in range(len(batch)): + sequence_tensor = batch_results["input_ids"][j].unsqueeze(0) + all_tokenized_sequences.append(sequence_tensor) + else: + for j in range(len(batch)): + all_tokenized_sequences.append([batch_results["input_ids"][j]]) + + return all_tokenized_sequences + + def get_cache_stats(self): + total = self.cache_hits + self.cache_misses + hit_rate = self.cache_hits / total if total > 0 else 0 + + batch_total = self.batch_cache_hits + self.batch_cache_misses + batch_hit_rate = self.batch_cache_hits / batch_total if batch_total > 0 else 0 + + subsequence_info = self._tokenize_subsequence.cache_info() + + return { + "process_stats": { + "sequences_processed": self.process_count, + "batches_processed": self.batch_count, + "batch_size": self.batch_size + }, + "sequence_cache": { + "hits": self.cache_hits, + "misses": self.cache_misses, + "hit_rate": hit_rate, + "cache_size": len(self.cache) + }, + "batch_cache": { + "hits": self.batch_cache_hits, + "misses": self.batch_cache_misses, + "hit_rate": batch_hit_rate, + "cache_size": len(self.batch_cache) + }, + "subsequence_cache": { + "hits": subsequence_info.hits, + "misses": subsequence_info.misses, + "size": subsequence_info.currsize, + "maxsize": subsequence_info.maxsize + } + } \ No newline at end of file diff --git a/bioscanclip/model/pre_trained_bert.py b/bioscanclip/model/pre_trained_bert.py deleted file mode 100644 index 2579e06..0000000 --- a/bioscanclip/model/pre_trained_bert.py +++ /dev/null @@ -1,81 +0,0 @@ -from transformers import AutoTokenizer, BertModel -import torch -import torch.nn as nn -import math -from torch import Tensor -def load_pre_trained_bert(): - tokenizer = AutoTokenizer.from_pretrained("prajjwal1/bert-small") - model = BertModel.from_pretrained("prajjwal1/bert-small") - for param in model.parameters(): - param.requires_grad = False - - return tokenizer, model - -# MODIFIED FROM https://github.com/JamesQFreeman/LoRA-barcode_bert/blob/main/lora.py - -class _LoRALayer(nn.Module): - def __init__(self, w: nn.Module, w_a: nn.Module, w_b: nn.Module): - super().__init__() - self.w = w - self.w_a = w_a - self.w_b = w_b - - def forward(self, x): - x = self.w(x) + self.w_b(self.w_a(x)) - return x - - -class LoRA_bert(nn.Module): - def __init__(self, model, r: int, num_classes: int = 0, lora_layer=None): - super(LoRA_bert, self).__init__() - - assert r > 0 - if lora_layer: - self.lora_layer = lora_layer - else: - self.lora_layer = list(range(len(model.encoder.layer))) - - # create for storage, then we can init them or load weights - self.w_As = [] # These are linear layers - self.w_Bs = [] - - # lets freeze first - for param in model.parameters(): - param.requires_grad = False - - for layer_idx, layer in enumerate(model.encoder.layer): - if layer_idx not in self.lora_layer: - continue - w_q_linear = layer.attention.self.query - w_v_linear = layer.attention.self.value - dim = layer.attention.self.query.in_features - - w_a_linear_q = nn.Linear(dim, r, bias=False) - w_b_linear_q = nn.Linear(r, dim, bias=False) - w_a_linear_v = nn.Linear(dim, r, bias=False) - w_b_linear_v = nn.Linear(r, dim, bias=False) - - self.w_As.append(w_a_linear_q) - self.w_Bs.append(w_b_linear_q) - self.w_As.append(w_a_linear_v) - self.w_Bs.append(w_b_linear_v) - - layer.attention.self.query = _LoRALayer(w_q_linear, w_a_linear_q, w_b_linear_q) - layer.attention.self.value = _LoRALayer(w_v_linear, w_a_linear_v, w_b_linear_v) - - self.reset_parameters() - self.lora_bert = model - - if num_classes > 0: - self.proj = nn.Linear(self.lora_bert.pooler.dense.out_features, num_classes) - - - def reset_parameters(self) -> None: - for w_A in self.w_As: - nn.init.kaiming_uniform_(w_A.weight, a=math.sqrt(5)) - for w_B in self.w_Bs: - nn.init.zeros_(w_B.weight) - - def forward(self, x) -> Tensor: - - return self.proj(self.lora_bert(**x).last_hidden_state.mean(dim=1)) diff --git a/bioscanclip/model/simple_clip.py b/bioscanclip/model/simple_clip.py index 0e2bac5..771e6a6 100644 --- a/bioscanclip/model/simple_clip.py +++ b/bioscanclip/model/simple_clip.py @@ -175,6 +175,8 @@ def load_clip_model(args, device=None): language_model_name = 'prajjwal1/bert-small' if hasattr(args.model_config.language, 'pre_train_model'): language_model_name = args.model_config.language.pre_train_model + if language_model_name == "bert-small" or language_model_name == "bert_small": + language_model_name = 'prajjwal1/bert-small' _, pre_trained_bert = load_pre_trained_bert(language_model_name) if disable_lora: language_encoder = CLIBDLanguageEncoder(model=pre_trained_bert, r=4, num_classes=args.model_config.output_dim, @@ -193,7 +195,7 @@ def load_clip_model(args, device=None): if hasattr(args.model_config, 'pre_train_for_barcode_bert') and args.model_config.pre_train_for_barcode_bert == "BIOSCAN-5M": barcode_bert_ckpt = args.bioscan_bert_checkpoint_trained_with_bioscan_5_m - elif hasattr(args.model_config, 'pre_train_for_barcode_bert') and args.model_config.pre_train_for_barcode_bert == "CANADA-1M": + elif hasattr(args.model_config, 'pre_train_for_barcode_bert') and args.model_config.pre_train_for_barcode_bert == "BIOSCAN-1M": barcode_bert_ckpt = args.bioscan_bert_checkpoint_trained_with_canada_1_5_m pre_trained_barcode_bert = load_pre_trained_bioscan_bert( diff --git a/bioscanclip/util/dataset.py b/bioscanclip/util/dataset.py index c1a96e3..d280a8a 100644 --- a/bioscanclip/util/dataset.py +++ b/bioscanclip/util/dataset.py @@ -7,17 +7,18 @@ import pandas as pd import scipy.io as sio import torch +from tqdm import tqdm from PIL import Image from torch.utils.data import Dataset import torchvision.transforms as transforms -from bioscanclip.model.dna_encoder import get_sequence_pipeline +from bioscanclip.model.dna_encoder import KmerTokenizerWithAttMask, NewKmerTokenizer from torch.utils.data.distributed import DistributedSampler import json import time from transformers import AutoTokenizer from bioscanclip.model.language_encoder import load_pre_trained_bert import open_clip -from bioscanclip.util.util import load_kmer_tokenizer, TensorResizeLongEdge +from bioscanclip.util.util import TensorResizeLongEdge DEVICE = torch.device("cuda" if torch.cuda.is_available() else "cpu") @@ -34,13 +35,6 @@ def get_label_ids(input_labels): return label_ids, label_to_id -def tokenize_dna_sequence(pipeline, dna_input): - list_of_output = [] - for i in dna_input: - list_of_output.append(pipeline(i)) - return list_of_output - - def prepare(dataset, rank, world_size, batch_size=32, pin_memory=False, num_workers=0, shuffle=False): sampler = DistributedSampler(dataset, num_replicas=world_size, rank=rank, shuffle=shuffle, drop_last=True) @@ -110,6 +104,9 @@ def __init__( labels=None, for_training=False, for_open_clip=False, + dna_tokenizer=None, + + tokenizer=None, ): if hasattr(args.model_config, "dataset") and args.model_config.dataset == "bioscan_5m": if hasattr(args.model_config, "train_with_small_subset") and args.model_config.train_with_small_subset: @@ -134,18 +131,20 @@ def __init__( self.pre_train_with_small_set = False if hasattr(args.model_config, "train_with_small_subset"): self.pre_train_with_small_set = args.model_config.train_with_small_subset - language_model_name = "prajjwal1/bert-small" - self.tokenizer, _ = load_pre_trained_bert(language_model_name) + self.dna_tokenizer = dna_tokenizer + if self.for_open_clip: # self.tokenizer = open_clip.get_tokenizer('ViT-B-32') - self.tokenizer = None + self.language_tokenizer = None else: if hasattr(args.model_config, "language"): language_model_name = "prajjwal1/bert-small" - if hasattr(args.model_config.language, "pre_train_model"): - language_model_name = args.model_config.language.pre_train_model - self.tokenizer, _ = load_pre_trained_bert(language_model_name) + if hasattr(args.model_config.language, "model"): + language_model_name = args.model_config.language.model + if language_model_name == "bert-small" or language_model_name == "bert_small": + language_model_name = "prajjwal1/bert-small" + self.language_tokenizer, _ = load_pre_trained_bert(language_model_name) list_of_label_dict = get_array_of_label_dicts(self.hdf5_inputs_path, split) self.list_of_label_string = [] @@ -260,10 +259,26 @@ def __getitem__(self, idx): if self.dna_inout_type == "sequence": if self.dna_tokens is None: curr_dna_input = self.hdf5_split_group["barcode"][idx].decode("utf-8") + if self.dna_tokenizer is not None: + curr_dna_input = self.dna_tokenizer(curr_dna_input) + curr_dna_input['input_ids'] = torch.tensor(curr_dna_input['input_ids']).clone() + curr_dna_input['token_type_ids'] = torch.tensor(curr_dna_input['token_type_ids']).clone() + curr_dna_input['attention_mask'] = torch.tensor(curr_dna_input['attention_mask']).clone() + else: + raise TypeError( + f"DNA input type is sequence, but dna_tokenizer is None. Please check the config file." + ) else: - curr_dna_input = self.dna_tokens[idx] + # Using preprocessed DNA tokens + raise NotImplementedError( + f"Using pre-tokenized DNA tokens is not supported now." + ) else: - curr_dna_input = self.hdf5_split_group["dna_features"][idx].astype(np.float32) + # Using pre-extracted DNA features + raise NotImplementedError( + f"DNA input can only be sequence now. Please check the config file." + ) + # curr_dna_input = self.hdf5_split_group["dna_features"][idx].astype(np.float32) if self.dataset == "bioscan_5m": curr_processid = self.hdf5_split_group["processid"][idx].decode("utf-8") @@ -276,16 +291,30 @@ def __getitem__(self, idx): language_token_type_ids = torch.zeros(1, ) language_attention_mask = torch.zeros(1, ) else: + if hasattr(self, "language_tokenizer") and self.language_tokenizer is not None: + language_tokens = self.language_tokenizer([self.list_of_label_string[idx]], padding="max_length", max_length=20, + truncation=True) - language_tokens = self.tokenizer([self.list_of_label_string[idx]], padding="max_length", max_length=20, - truncation=True) - language_input_ids = language_tokens['input_ids'] - language_token_type_ids = language_tokens['token_type_ids'] - language_attention_mask = language_tokens['attention_mask'] + language_input_ids = language_tokens['input_ids'] + language_token_type_ids = language_tokens['token_type_ids'] + language_attention_mask = language_tokens['attention_mask'] - language_input_ids = torch.tensor(language_input_ids[0]) - language_token_type_ids = torch.tensor(language_token_type_ids[0]) - language_attention_mask = torch.tensor(language_attention_mask[0]) + language_input_ids = torch.tensor(language_input_ids[0]) + language_token_type_ids = torch.tensor(language_token_type_ids[0]) + language_attention_mask = torch.tensor(language_attention_mask[0]) + else: + """ + TODO: Edit the code for correctly manage the language tokenization for pre-trained openclip encoder. + """ + # set ids and others to empty tensors + language_tokens = { + 'input_ids': torch.zeros(1, ), + 'token_type_ids': torch.zeros(1, ), + 'attention_mask': torch.zeros(1, ) + } + language_input_ids = torch.zeros(1, ) + language_token_type_ids = torch.zeros(1, ) + language_attention_mask = torch.zeros(1, ) # language_input_ids = self.hdf5_split_group["language_tokens_input_ids"][idx] # language_token_type_ids = self.hdf5_split_group["language_tokens_token_type_ids"][idx] @@ -391,7 +420,6 @@ def construct_dataloader( args, split, length, - sequence_pipeline, return_language=False, labels=None, for_pre_train=False, @@ -414,19 +442,15 @@ def construct_dataloader( dna_type = args.model_config.dna.input_type if dna_type == "sequence": - if hasattr(args.model_config, "pre_train_for_barcode_bert") and (args.model_config.pre_train_for_barcode_bert == "BIOSCAN-5M" or args.model_config.pre_train_for_barcode_bert == "CANADA-1M"): - pass + if hasattr(args.model_config, "pre_train_for_barcode_bert") and (args.model_config.pre_train_for_barcode_bert == "BIOSCAN-5M" or args.model_config.pre_train_for_barcode_bert == "BIOSCAN-1M"): + dna_tokenizer = NewKmerTokenizer(max_length=args.barcodebert_setting.old_model_setting.max_len) else: - - if args.model_config.dataset == "bioscan_5m": - if hasattr(args.model_config, "train_with_small_subset") and args.model_config.train_with_small_subset: - hdf5_file = h5py.File(args.bioscan_5m_data.path_to_smaller_hdf5_data, "r", libver="latest") - else: - hdf5_file = h5py.File(args.bioscan_5m_data.path_to_hdf5_data, "r", libver="latest") - else: - hdf5_file = h5py.File(args.bioscan_data.path_to_hdf5_data, "r", libver="latest") - unprocessed_dna_barcode = np.array([item.decode("utf-8") for item in hdf5_file[split]["barcode"][:]]) - barcode_bert_dna_tokens = tokenize_dna_sequence(sequence_pipeline, unprocessed_dna_barcode) + dna_tokenizer = KmerTokenizerWithAttMask(k=args.barcodebert_setting.old_model_setting.k, + max_len=args.barcodebert_setting.old_model_setting.max_len) + else: + raise NotImplementedError( + f"DNA input type {dna_type} is not supported. Please check the config file." + ) dataset = Dataset_for_CL( args, @@ -439,6 +463,7 @@ def construct_dataloader( labels=labels, for_training=for_pre_train, for_open_clip=for_open_clip, + dna_tokenizer=dna_tokenizer, ) num_workers = 8 @@ -476,13 +501,10 @@ def load_bioscan_dataloader_with_train_seen_and_separate_keys(args, world_size=N return_language = True - sequence_pipeline = get_sequence_pipeline() - train_seen_dataloader = construct_dataloader( args, "train_seen", length_dict["train_seen"], - sequence_pipeline, return_language=return_language, labels=None, for_pre_train=False, @@ -495,7 +517,7 @@ def load_bioscan_dataloader_with_train_seen_and_separate_keys(args, world_size=N args, "val_seen", length_dict["val_seen"], - sequence_pipeline, + return_language=return_language, labels=None, for_pre_train=False, @@ -507,7 +529,7 @@ def load_bioscan_dataloader_with_train_seen_and_separate_keys(args, world_size=N args, "val_unseen", length_dict["val_unseen"], - sequence_pipeline, + return_language=return_language, labels=None, for_pre_train=False, @@ -519,7 +541,7 @@ def load_bioscan_dataloader_with_train_seen_and_separate_keys(args, world_size=N args, "seen_keys", length_dict["seen_keys"], - sequence_pipeline, + return_language=return_language, labels=None, for_pre_train=False, @@ -531,7 +553,7 @@ def load_bioscan_dataloader_with_train_seen_and_separate_keys(args, world_size=N args, "val_unseen_keys", length_dict["val_unseen_keys"], - sequence_pipeline, + return_language=return_language, labels=None, for_pre_train=False, @@ -542,7 +564,7 @@ def load_bioscan_dataloader_with_train_seen_and_separate_keys(args, world_size=N args, "test_unseen_keys", length_dict["test_unseen_keys"], - sequence_pipeline, + return_language=return_language, labels=None, for_pre_train=False, @@ -564,13 +586,13 @@ def load_dataloader_for_everything_in_5m(args, world_size=None, rank=None): return_language = True - sequence_pipeline = get_sequence_pipeline() + dna_tokenizer = KmerTokenizerWithAttMask(k=args.barcodebert_setting.old_model_setting.k, max_len=args.barcodebert_setting.old_model_setting.max_len) pre_train_dataloader = construct_dataloader( args, "no_split_and_seen_train", length_dict["no_split_and_seen_train"], - sequence_pipeline, + return_language=return_language, labels=None, for_pre_train=False, @@ -583,7 +605,7 @@ def load_dataloader_for_everything_in_5m(args, world_size=None, rank=None): args, "all_keys", length_dict["all_keys"], - sequence_pipeline, + return_language=return_language, labels=None, for_pre_train=False, @@ -595,7 +617,7 @@ def load_dataloader_for_everything_in_5m(args, world_size=None, rank=None): args, "val_seen", length_dict["val_seen"], - sequence_pipeline, + return_language=return_language, labels=None, for_pre_train=False, @@ -607,7 +629,7 @@ def load_dataloader_for_everything_in_5m(args, world_size=None, rank=None): args, "val_unseen", length_dict["val_unseen"], - sequence_pipeline, + return_language=return_language, labels=None, for_pre_train=False, @@ -619,7 +641,7 @@ def load_dataloader_for_everything_in_5m(args, world_size=None, rank=None): args, "test_seen", length_dict["test_seen"], - sequence_pipeline, + return_language=return_language, labels=None, for_pre_train=False, @@ -631,7 +653,7 @@ def load_dataloader_for_everything_in_5m(args, world_size=None, rank=None): args, "test_unseen", length_dict["test_unseen"], - sequence_pipeline, + return_language=return_language, labels=None, for_pre_train=False, @@ -642,7 +664,7 @@ def load_dataloader_for_everything_in_5m(args, world_size=None, rank=None): args, "other_heldout", length_dict["other_heldout"], - sequence_pipeline, + return_language=return_language, labels=None, for_pre_train=False, @@ -658,13 +680,10 @@ def load_dataloader(args, world_size=None, rank=None, for_pretrain=True): return_language = True - sequence_pipeline = get_sequence_pipeline() - seen_val_dataloader = construct_dataloader( args, "val_seen", length_dict["val_seen"], - sequence_pipeline, return_language=return_language, labels=None, for_pre_train=False, @@ -676,7 +695,6 @@ def load_dataloader(args, world_size=None, rank=None, for_pretrain=True): args, "val_unseen", length_dict["val_unseen"], - sequence_pipeline, return_language=return_language, labels=None, for_pre_train=False, @@ -688,7 +706,6 @@ def load_dataloader(args, world_size=None, rank=None, for_pretrain=True): args, "all_keys", length_dict["all_keys"], - sequence_pipeline, return_language=return_language, labels=None, for_pre_train=False, @@ -704,7 +721,6 @@ def load_dataloader(args, world_size=None, rank=None, for_pretrain=True): args, "no_split_and_seen_train", length_dict["no_split_and_seen_train"], - sequence_pipeline, return_language=return_language, labels=None, for_pre_train=True, @@ -717,7 +733,6 @@ def load_dataloader(args, world_size=None, rank=None, for_pretrain=True): args, "no_split", length_dict["no_split"], - sequence_pipeline, return_language=return_language, labels=None, for_pre_train=True, @@ -731,7 +746,6 @@ def load_dataloader(args, world_size=None, rank=None, for_pretrain=True): args, "train_seen", length_dict["train_seen"], - sequence_pipeline, return_language=return_language, labels=None, for_pre_train=False, @@ -747,14 +761,12 @@ def load_bioscan_dataloader_all_small_splits(args, world_size=None, rank=None): return_language = True - sequence_pipeline = get_sequence_pipeline() if hasattr(args.model_config, 'dataset') and args.model_config.dataset == "bioscan_5m": train_seen_dataloader = construct_dataloader( args, "seen_keys", length_dict["seen_keys"], - sequence_pipeline, return_language=return_language, labels=None, for_pre_train=False, @@ -766,7 +778,6 @@ def load_bioscan_dataloader_all_small_splits(args, world_size=None, rank=None): args, "train_seen", length_dict["train_seen"], - sequence_pipeline, return_language=return_language, labels=None, for_pre_train=False, @@ -778,7 +789,6 @@ def load_bioscan_dataloader_all_small_splits(args, world_size=None, rank=None): args, "val_seen", length_dict["val_seen"], - sequence_pipeline, return_language=return_language, labels=None, for_pre_train=False, @@ -790,7 +800,6 @@ def load_bioscan_dataloader_all_small_splits(args, world_size=None, rank=None): args, "val_unseen", length_dict["val_unseen"], - sequence_pipeline, return_language=return_language, labels=None, for_pre_train=False, @@ -802,7 +811,6 @@ def load_bioscan_dataloader_all_small_splits(args, world_size=None, rank=None): args, "test_seen", length_dict["test_seen"], - sequence_pipeline, return_language=return_language, labels=None, for_pre_train=False, @@ -814,7 +822,6 @@ def load_bioscan_dataloader_all_small_splits(args, world_size=None, rank=None): args, "test_unseen", length_dict["test_unseen"], - sequence_pipeline, return_language=return_language, labels=None, for_pre_train=False, @@ -826,7 +833,6 @@ def load_bioscan_dataloader_all_small_splits(args, world_size=None, rank=None): args, "seen_keys", length_dict["seen_keys"], - sequence_pipeline, return_language=return_language, labels=None, for_pre_train=False, @@ -839,7 +845,6 @@ def load_bioscan_dataloader_all_small_splits(args, world_size=None, rank=None): args, "unseen_keys", length_dict["unseen_keys"], - sequence_pipeline, return_language=return_language, labels=None, for_pre_train=False, @@ -850,7 +855,6 @@ def load_bioscan_dataloader_all_small_splits(args, world_size=None, rank=None): args, "unseen_keys", length_dict["unseen_keys"], - sequence_pipeline, return_language=return_language, labels=None, for_pre_train=False, @@ -863,7 +867,6 @@ def load_bioscan_dataloader_all_small_splits(args, world_size=None, rank=None): args, "val_unseen_keys", length_dict["val_unseen_keys"], - sequence_pipeline, return_language=return_language, labels=None, for_pre_train=False, @@ -874,7 +877,6 @@ def load_bioscan_dataloader_all_small_splits(args, world_size=None, rank=None): args, "test_unseen_keys", length_dict["test_unseen_keys"], - sequence_pipeline, return_language=return_language, labels=None, for_pre_train=False, @@ -886,7 +888,6 @@ def load_bioscan_dataloader_all_small_splits(args, world_size=None, rank=None): args, "all_keys", length_dict["all_keys"], - sequence_pipeline, return_language=return_language, labels=None, for_pre_train=False, diff --git a/bioscanclip/util/dataset_for_insect_dataset.py b/bioscanclip/util/dataset_for_insect_dataset.py index 931c1eb..08491da 100644 --- a/bioscanclip/util/dataset_for_insect_dataset.py +++ b/bioscanclip/util/dataset_for_insect_dataset.py @@ -5,7 +5,7 @@ import scipy.io as sio import torch from PIL import Image -from bioscanclip.model.dna_encoder import get_sequence_pipeline +from bioscanclip.model.dna_encoder import KmerTokenizerWithAttMask from torch.utils.data import Dataset import torchvision.transforms as transforms from torch.utils.data.distributed import DistributedSampler @@ -173,13 +173,18 @@ def load_insect_dataloader_trainval(args,num_workers=8, shuffle_for_train_seen_k with open(filename, 'r') as file: specie_to_other_labels = json.load(file) - sequence_pipeline = get_sequence_pipeline() + if hasattr(args.model_config, "pre_train_for_barcode_bert") and ( + args.model_config.pre_train_for_barcode_bert == "BIOSCAN-5M" or args.model_config.pre_train_for_barcode_bert == "CANADA-1M"): + dna_tokenizer = AutoTokenizer.from_pretrained("bioscan-ml/BarcodeBERT", trust_remote_code=True) + else: + dna_tokenizer = KmerTokenizerWithAttMask(k=args.barcodebert_setting.old_model_setting.k, + max_len=args.barcodebert_setting.old_model_setting.max_len) trainval_dataset = INSECTDataset( args.insect_data.path_to_att_splits_mat, args.insect_data.path_to_res_101_mat, species_to_others=specie_to_other_labels, split="trainval_loc", image_hdf5_path=args.insect_data.path_to_image_hdf5, - dna_transforms=sequence_pipeline, for_training=True, cl_label=False + dna_tokenizer=dna_tokenizer, for_training=True, cl_label=False ) @@ -193,14 +198,19 @@ def load_insect_dataloader(args, world_size=None, rank=None, num_workers=8, load with open(filename, 'r') as file: specie_to_other_labels = json.load(file) - sequence_pipeline = get_sequence_pipeline() + if hasattr(args.model_config, "pre_train_for_barcode_bert") and ( + args.model_config.pre_train_for_barcode_bert == "BIOSCAN-5M" or args.model_config.pre_train_for_barcode_bert == "CANADA-1M"): + dna_tokenizer = AutoTokenizer.from_pretrained("bioscan-ml/BarcodeBERT", trust_remote_code=True) + else: + dna_tokenizer = KmerTokenizerWithAttMask(k=args.barcodebert_setting.old_model_setting.k, + max_len=args.barcodebert_setting.old_model_setting.max_len) if load_all_in_one: all_dataset = INSECTDataset( args.insect_data.path_to_att_splits_mat, args.insect_data.path_to_res_101_mat, species_to_others=specie_to_other_labels, split="all", image_hdf5_path=args.insect_data.path_to_image_hdf5, - dna_transforms=sequence_pipeline, for_training=False + dna_tokenizer=dna_tokenizer, for_training=False ) all_dataloader = DataLoader(all_dataset, batch_size=args.model_config.batch_size, @@ -212,35 +222,35 @@ def load_insect_dataloader(args, world_size=None, rank=None, num_workers=8, load args.insect_data.path_to_att_splits_mat, args.insect_data.path_to_res_101_mat, species_to_others=specie_to_other_labels, split="train_loc", image_hdf5_path=args.insect_data.path_to_image_hdf5, - dna_transforms=sequence_pipeline, for_training=True + dna_tokenizer=dna_tokenizer, for_training=True ) train_dataset_for_key = INSECTDataset( args.insect_data.path_to_att_splits_mat, args.insect_data.path_to_res_101_mat, species_to_others=specie_to_other_labels, split="train_loc", image_hdf5_path=args.insect_data.path_to_image_hdf5, - dna_transforms=sequence_pipeline, for_training=False + dna_tokenizer=dna_tokenizer, for_training=False ) val_dataset = INSECTDataset( args.insect_data.path_to_att_splits_mat, args.insect_data.path_to_res_101_mat, species_to_others=specie_to_other_labels, split="val_loc", image_hdf5_path=args.insect_data.path_to_image_hdf5, - dna_transforms=sequence_pipeline, for_training=False + dna_tokenizer=dna_tokenizer, for_training=False ) test_seen_dataset = INSECTDataset( args.insect_data.path_to_att_splits_mat, args.insect_data.path_to_res_101_mat, species_to_others=specie_to_other_labels, split="test_seen_loc", image_hdf5_path=args.insect_data.path_to_image_hdf5, - dna_transforms=sequence_pipeline, for_training=False + dna_tokenizer=dna_tokenizer, for_training=False ) test_unseen_dataset = INSECTDataset( args.insect_data.path_to_att_splits_mat, args.insect_data.path_to_res_101_mat, species_to_others=specie_to_other_labels, split="test_unseen_loc", image_hdf5_path=args.insect_data.path_to_image_hdf5, - dna_transforms=sequence_pipeline, for_training=False + dna_tokenizer=dna_tokenizer, for_training=False ) if rank is None: print(rank) diff --git a/bioscanclip/util/util.py b/bioscanclip/util/util.py index 926f0f3..b45f132 100644 --- a/bioscanclip/util/util.py +++ b/bioscanclip/util/util.py @@ -73,28 +73,10 @@ def print_separator(self): print(f"+{separator}+") -class PadSequence(object): - def __init__(self, max_len): - self.max_len = max_len - def __call__(self, dna_sequence): - if len(dna_sequence) > self.max_len: - return dna_sequence[: self.max_len] - else: - return dna_sequence + "N" * (self.max_len - len(dna_sequence)) -class KmerTokenizer(object): - def __init__(self, k, stride=1): - self.k = k - self.stride = stride - def __call__(self, dna_sequence): - tokens = [] - for i in range(0, len(dna_sequence) - self.k + 1, self.stride): - k_mer = dna_sequence[i : i + self.k] - tokens.append(k_mer) - return tokens class NewKmerTokenizer(object): @@ -837,37 +819,6 @@ def remove_module_from_state_dict(state_dict): new_state_dict[key.replace("module.", "")] = value return new_state_dict -def load_kmer_tokenizer(args, k=4): - base_pairs = "ACGT" - tokenize_n_nucleotide = False - special_tokens = ["[MASK]", "[UNK]"] - UNK_TOKEN = "[UNK]" - stride = 1 - max_len = 660 - - k_mer = k - kmers = ["".join(kmer) for kmer in product(base_pairs, repeat=k_mer)] - - if tokenize_n_nucleotide: - prediction_kmers = [] - other_kmers = [] - for kmer in kmers: - if "N" in kmer: - other_kmers.append(kmer) - else: - prediction_kmers.append(kmer) - - kmers = prediction_kmers + other_kmers - - kmer_dict = dict.fromkeys(kmers, 1) - vocab = build_vocab_from_dict(kmer_dict, specials=special_tokens) - vocab.set_default_index(vocab[UNK_TOKEN]) - vocab_size = len(vocab) - tokenizer = KmerTokenizer( - k_mer, vocab, stride=stride, padding=True, max_len=max_len - ) - - return tokenizer class TensorResizeLongEdge(object): def __init__(self, long_edge_size, interpolation_mode='bilinear'): diff --git a/scripts/result/plots/line_plot_for_multiple_experiments_dna_to_dna.py b/scripts/result/plots/line_plot_for_multiple_experiments_dna_to_dna.py index 44f2139..c4f48d4 100644 --- a/scripts/result/plots/line_plot_for_multiple_experiments_dna_to_dna.py +++ b/scripts/result/plots/line_plot_for_multiple_experiments_dna_to_dna.py @@ -38,7 +38,7 @@ -fig, ax = plt.subplots(figsize=(6, 4)) +fig, ax = plt.subplots(figsize=(5.3, 4)) ax.plot(x, baseline_dna_to_dna_seen, 'o-', color=red, linewidth=5) ax.plot(x, baseline_dna_to_dna_unseen, 'o--', color=red, linewidth=5) ax.plot(x, i_d_dna_to_dna_seen, 'o-', color=yellow, linewidth=5) diff --git a/scripts/result/plots/line_plot_for_multiple_experiments_image_to_dna.py b/scripts/result/plots/line_plot_for_multiple_experiments_image_to_dna.py index 0f0c45d..6ade446 100644 --- a/scripts/result/plots/line_plot_for_multiple_experiments_image_to_dna.py +++ b/scripts/result/plots/line_plot_for_multiple_experiments_image_to_dna.py @@ -37,7 +37,7 @@ i_d_t_image_to_dna_unseen = [88.5, 50.1, 20.8, 8.6] -fig, ax = plt.subplots(figsize=(6, 4)) +fig, ax = plt.subplots(figsize=(5.3, 4)) ax.plot(x, baseline_image_to_dna_seen, 'o-', color=red, linewidth=5,) ax.plot(x, baseline_image_to_dna_unseen, 'o--', color=red, linewidth=5) ax.plot(x, i_d_image_to_dna_seen, 'o-', color=yellow, linewidth=5) diff --git a/scripts/result/plots/line_plot_for_multiple_experiments_image_to_image.py b/scripts/result/plots/line_plot_for_multiple_experiments_image_to_image.py index b91f161..3065cc3 100644 --- a/scripts/result/plots/line_plot_for_multiple_experiments_image_to_image.py +++ b/scripts/result/plots/line_plot_for_multiple_experiments_image_to_image.py @@ -6,6 +6,9 @@ red = "#FA7F6F" blue = "#82B0D2" +font_size_a = 22 +font_size_b = 16 + x = np.arange(4) labels = ['order', 'family', 'genus', 'species'] @@ -34,20 +37,20 @@ i_d_t_image_to_dna_unseen = [88.5, 50.1, 20.8, 8.6] -fig, ax = plt.subplots(figsize=(6, 4)) -ax.plot(x, baseline_image_to_image_seen, 'o-', color=red) -ax.plot(x, baseline_image_to_image_unseen, 'o--', color=red) -ax.plot(x, i_d_image_to_image_seen, 'o-', color=yellow) -ax.plot(x, i_d_image_to_image_unseen, 'o--', color=yellow) -ax.plot(x, i_d_t_image_to_image_seen, 'o-', color=blue) -ax.plot(x, i_d_t_image_to_image_unseen, 'o--', color=blue) +fig, ax = plt.subplots(figsize=(5.3, 4)) +ax.plot(x, baseline_image_to_image_seen, 'o-', color=red, linewidth=5) +ax.plot(x, baseline_image_to_image_unseen, 'o--', color=red, linewidth=5) +ax.plot(x, i_d_image_to_image_seen, 'o-', color=yellow, linewidth=5) +ax.plot(x, i_d_image_to_image_unseen, 'o--', color=yellow, linewidth=5) +ax.plot(x, i_d_t_image_to_image_seen, 'o-', color=blue, linewidth=5) +ax.plot(x, i_d_t_image_to_image_unseen, 'o--', color=blue, linewidth=5) ax.set_xticks(x) ax.set_xticklabels(labels) ax.set_ylim(0, 100) -ax.tick_params(axis='both', which='major', labelsize=12) -ax.set_ylabel('Macro-accuracy (%)', fontsize=16) -ax.set_title('Image to Image', fontsize=16) +ax.tick_params(axis='both', which='major', labelsize=font_size_b) +ax.set_ylabel('Macro-accuracy (%)', fontsize=font_size_a) +ax.set_title('Image to Image', fontsize=font_size_a) for y in np.arange(0, 101, 5): if y % 10 == 0: @@ -56,18 +59,18 @@ ax.axhline(y=y, color='grey', linewidth=0.2, linestyle='-') method_handles = [ - Line2D([0], [0], color=red, lw=2, label='No align'), - Line2D([0], [0], color=yellow, lw=2, label='Image + DNA'), - Line2D([0], [0], color=blue, lw=2, label='Image + DNA + Taxonomy') + Line2D([0], [0], color=red, lw=2, label='No align', linewidth=5), + Line2D([0], [0], color=yellow, lw=2, label='Image + DNA', linewidth=5), + Line2D([0], [0], color=blue, lw=2, label='Image + DNA + Taxonomy', linewidth=5) ] style_handles = [ - Line2D([0], [0], color='black', lw=2, linestyle='-', label='Seen'), - Line2D([0], [0], color='black', lw=2, linestyle='--', label='Unseen') + Line2D([0], [0], color='black', lw=2, linestyle='-', label='Seen', linewidth=5), + Line2D([0], [0], color='black', lw=2, linestyle='--', label='Unseen', linewidth=5) ] -legend1 = ax.legend(handles=method_handles, loc='lower left', bbox_to_anchor=(0, 0)) -ax.add_artist(legend1) -legend2 = ax.legend(handles=style_handles, loc='lower left', bbox_to_anchor=(0.48, 0)) +# legend1 = ax.legend(handles=method_handles, loc='lower left', bbox_to_anchor=(0, 0)) +# ax.add_artist(legend1) +# legend2 = ax.legend(handles=style_handles, loc='lower left', bbox_to_anchor=(0.48, 0)) plt.tight_layout() # plt.show() plt.savefig("image_to_image.png", dpi=300, bbox_inches='tight') \ No newline at end of file diff --git a/scripts/train_cl.py b/scripts/train_cl.py index 778bf1a..8cc5be5 100644 --- a/scripts/train_cl.py +++ b/scripts/train_cl.py @@ -148,6 +148,7 @@ def main_process(rank: int, world_size: int, args): args.save_inference = False args.save_ckpt = False + current_datetime = datetime.datetime.now() formatted_datetime = current_datetime.strftime("%Y-%m-%d_%H%M%S") args = copy.deepcopy(args) @@ -285,6 +286,7 @@ def main_process(rank: int, world_size: int, args): for_open_clip=for_open_clip, fix_temperature=fix_temperature, scaler=scaler, enable_autocast=enable_amp) + if (epoch % args.model_config.evaluation_period == 0 or epoch == args.model_config.epochs - 1) and rank == 0 and epoch > eval_skip_epoch: original_model = model.module if hasattr(model, 'module') else model if args.save_ckpt: