diff --git a/.gitignore b/.gitignore
new file mode 100644
index 0000000..868a782
--- /dev/null
+++ b/.gitignore
@@ -0,0 +1,6 @@
+_test_*.*
+__pycache__
+.venv
+.idea
+*.pth
+*.ini
diff --git a/README.md b/README.md
index 97db0b0..116a443 100644
--- a/README.md
+++ b/README.md
@@ -1,2 +1,1073 @@
-# ComfyUI_LayerStyle_Advance
+# ComfyUI Layer Style Advance
+
+[中文说明点这里](./README_CN.MD)
+
+
The nodes detached from [ComfyUI Layer Style](https://github.com/chflame163/ComfyUI_LayerStyle) are mainly those with complex requirements for dependency packages.
+
+
+
+## Example workflow
+
+Some JSON workflow files in the ```workflow``` directory, That's examples of how these nodes can be used in ComfyUI.
+
+## How to install
+
+(Taking ComfyUI official portable package and Aki ComfyUI package as examples, please modify the dependency environment directory for other ComfyUI environments)
+
+### Install plugin
+
+* Recommended use ComfyUI Manager for installation.
+
+* Or open the cmd window in the plugin directory of ComfyUI, like ```ComfyUI\custom_nodes```,type
+
+ ```
+ git clone https://github.com/chflame163/ComfyUI_LayerStyle.git
+ ```
+
+* Or download the zip file and extracted, copy the resulting folder to ```ComfyUI\custom_ Nodes```
+
+### Install dependency packages
+
+* for ComfyUI official portable package, double-click the ```install_requirements.bat``` in the plugin directory, for Aki ComfyUI package double-click on the ```install_requirements_aki.bat``` in the plugin directory, and wait for the installation to complete.
+
+* Or install dependency packages, open the cmd window in the ComfyUI_LayerStyle plugin directory like
+ ```ComfyUI\custom_ Nodes\ComfyUI_LayerStyle``` and enter the following command,
+
+ for ComfyUI official portable package, type:
+
+```
+..\..\..\python_embeded\python.exe -s -m pip install .\whl\docopt-0.6.2-py2.py3-none-any.whl
+..\..\..\python_embeded\python.exe -s -m pip install .\whl\hydra_core-1.3.2-py3-none-any.whl
+..\..\..\python_embeded\python.exe -s -m pip install -r requirements.txt
+.\repair_dependency.bat
+```
+
+ for Aki ComfyUI package, type:
+
+```
+..\..\python\python.exe -s -m pip install .\whl\docopt-0.6.2-py2.py3-none-any.whl
+..\..\python\python.exe -s -m pip install .\whl\hydra_core-1.3.2-py3-none-any.whl
+..\..\python\python.exe -s -m pip install -r requirements.txt
+.\repair_dependency.bat
+```
+
+* Restart ComfyUI.
+
+### Download Model Files
+
+Chinese domestic users from [BaiduNetdisk](https://pan.baidu.com/s/1T_uXMX3OKIWOJLPuLijrgA?pwd=1yye) and other users from [huggingface.co/chflame163/ComfyUI_LayerStyle](https://huggingface.co/chflame163/ComfyUI_LayerStyle/tree/main)
+download all files and copy them to ```ComfyUI\models``` folder. This link provides all the model files required for this plugin.
+Or download the model file according to the instructions of each node.
+
+## Common Issues
+
+If the node cannot load properly or there are errors during use, please check the error message in the ComfyUI terminal window. The following are common errors and their solutions.
+
+### Warning: xxxx.ini not found, use default xxxx..
+
+This warning message indicates that the ini file cannot be found and does not affect usage. If you do not want to see these warnings, please modify all ```*.ini.example``` files in the plugin directory to ```*.ini```.
+
+### ModuleNotFoundError: No module named 'psd_tools'
+
+This error is that the ```psd_tools``` were not installed correctly.
+
+Solution:
+
+* Close ComfyUI and open the terminal window in the plugin directory and execute the following command:
+ ```../../../python_embeded/python.exe -s -m pip install psd_tools```
+ If error occurs during the installation of psd_tool, such as ```ModuleNotFoundError: No module named 'docopt'``` , please download [docopt's whl](https://www.piwheels.org/project/docopt/) and manual install it.
+ execute the following command in terminal window:
+ ```../../../python_embeded/python.exe -s -m pip install path/docopt-0.6.2-py2.py3-none-any.whl``` the ```path``` is path name of whl file.
+
+### Cannot import name 'guidedFilter' from 'cv2.ximgproc'
+
+This error is caused by incorrect version of the ```opencv-contrib-python``` package,or this package is overwriteen by other opencv packages.
+
+### NameError: name 'guidedFilter' is not defined
+
+The reason for the problem is the same as above.
+
+### Cannot import name 'VitMatteImageProcessor' from 'transformers'
+
+This error is caused by the low version of ```transformers``` package.
+
+### insightface Loading very slow
+
+This error is caused by the low version of ```protobuf``` package.
+
+#### For the issues with the above three dependency packages, please double click ```repair_dependency.bat``` (for Official ComfyUI Protable) or ```repair_dependency_aki.bat``` (for ComfyUI-aki-v1.x) in the plugin folder to automatically fix them.
+
+### onnxruntime::python::CreateExecutionProviderInstance CUDA_PATH is set but CUDA wasn't able to be loaded. Please install the correct version of CUDA and cuDNN as mentioned in the GPU requirements page
+
+Solution:
+Reinstall the ```onnxruntime``` dependency package.
+
+### Error loading model xxx: We couldn't connect to huggingface.co ...
+
+Check the network environment. If you cannot access huggingface.co normally in China, try modifying the huggingface_hub package to force the use hf_mirror.
+
+* Find ```constants.py``` in the directory of ```huggingface_hub``` package (usually ```Lib/site packages/huggingface_hub``` in the virtual environment path),
+ Add a line after ```import os```
+
+ ```
+ os.environ['HF_ENDPOINT'] = 'https://hf-mirror.com'
+ ```
+
+### ValueError: Trimap did not contain foreground values (xxxx...)
+
+This error is caused by the mask area being too large or too small when using the ```PyMatting``` method to handle the mask edges.
+
+Solution:
+
+* Please adjust the parameters to change the effective area of the mask. Or use other methods to handle the edges.
+
+### Requests.exceptions.ProxyError: HTTPSConnectionPool(xxxx...)
+
+When this error has occurred, please check the network environment.
+
+### UnboundLocalError: local variable 'clip_processor' referenced before assignment
+### UnboundLocalError: local variable 'text_model' referenced before assignment
+If this error occurs when executing ```JoyCaption2``` node and it has been confirmed that the model file has been placed in the correct directory,
+please check the ```transformers``` dependency package version is at least 4.43.2 or higher.
+If ```transformers``` version is higher than or equal to 4.45.0, and also have error message:
+```
+Error loading models: De️️scriptors cannot be created directly.
+If this call came from a _pb2.py file, your generated code is out of date and must be regenerated with protoc >= 3.19.0.
+......
+```
+Please try downgrading the ```protobuf``` dependency package to 3.20.3, or set environment variables: ```PROTOCOL_BUFFERS_PYTHON_IMPLEMENTATION=python```.
+
+
+
+## Update
+
+**If the dependency package error after updating, please double clicking ```repair_dependency.bat``` (for Official ComfyUI Protable) or ```repair_dependency_aki.bat``` (for ComfyUI-aki-v1.x) in the plugin folder to reinstall the dependency packages.
+
+* Discard the dependencies required for the [ObjectDetector YOLOWorld](#ObjectDetectorYOLOWorld) node from the requirements. txt file. To use this node, please manually install the dependency package.
+* Strip some nodes from [ComfyUI Layer Style](https://github.com/chflame163/ComfyUI_LayerStyle) to this repository.
+
+
+
+## Description
+
+### QWenImage2Prompt
+
+Inference the prompts based on the image. this node is repackage of the [ComfyUI_VLM_nodes](https://github.com/gokayfem/ComfyUI_VLM_nodes)'s ```UForm-Gen2 Qwen Node```, thanks to the original author.
+Download model files from [huggingface](https://huggingface.co/unum-cloud/uform-gen2-qwen-500m) or [Baidu Netdisk](https://pan.baidu.com/s/1oRkUoOKWaxGod_XTJ8NiTA?pwd=d5d2) to ```ComfyUI/models/LLavacheckpoints/files_for_uform_gen2_qwen``` folder.
+
+
+
+Node Options:
+
+* question: Prompt of UForm-Gen-QWen model.
+
+
+### LlamaVision
+Use the Llama 3.2 vision model for local inference. Can be used to generate prompt words. part of the code for this node comes from [ComfyUI-PixtralLlamaMolmoVision](https://github.com/SeanScripts/ComfyUI-PixtralLlamaMolmoVision), thank you to the original author.
+To use this node, the ```transformers``` need upgraded to 4.45.0 or higher.
+Download models from [BaiduNetdisk](https://pan.baidu.com/s/18oHnTrkNMiwKLMcUVrfFjA?pwd=4g81) or [huggingface/SeanScripts](https://huggingface.co/SeanScripts/Llama-3.2-11B-Vision-Instruct-nf4/tree/main) , and copy to ```ComfyUI/models/LLM```.
+
+
+Node Options:
+
+
+* image: Image input.
+* model: Currently, only the "Llama-3.2-11B-Vision-Instruct-nf4" is available.
+* system_prompt: System prompt words for LLM model.
+* user_prompt: User prompt words for LLM model.
+* max_new_tokens: max_new_tokens for LLM model.
+* do_sample: do_sample for LLM model.
+* top-p: top_p for LLM model.
+* top_k: top_k for LLM model.
+* stop_strings: The stop strings.
+* seed: The seed of random number.
+* control_after_generate: Seed change options. If this option is fixed, the generated random number will always be the same.
+* include_prompt_in_output: Does the output contain prompt words.
+* cache_model: Whether to cache the model.
+
+### JoyCaption2
+Use the JoyCaption-alpha-two model for local inference. Can be used to generate prompt words. this node is https://huggingface.co/John6666/joy-caption-alpha-two-cli-mod Implementation in ComfyUI, thank you to the original author.
+Download models form [BaiduNetdisk](https://pan.baidu.com/s/1dOjbUEacUOhzFitAQ3uIeQ?pwd=4ypv) and [BaiduNetdisk](https://pan.baidu.com/s/1mH1SuW45Dy6Wga7aws5siQ?pwd=w6h5) ,
+or [huggingface/Orenguteng](https://huggingface.co/Orenguteng/Llama-3.1-8B-Lexi-Uncensored-V2/tree/main) and [huggingface/unsloth](https://huggingface.co/unsloth/Meta-Llama-3.1-8B-Instruct/tree/main) , then copy to ```ComfyUI/models/LLM```,
+Download models from [BaiduNetdisk](https://pan.baidu.com/s/1pkVymOsDcXqL7IdQJ6lMVw?pwd=v8wp) or [huggingface/google](https://huggingface.co/google/siglip-so400m-patch14-384/tree/main) , and copy to ```ComfyUI/models/clip```,
+Donwload the ```cgrkzexw-599808``` folder from [BaiduNetdisk](https://pan.baidu.com/s/12TDwZAeI68hWT6MgRrrK7Q?pwd=d7dh) or [huggingface/John6666](https://huggingface.co/John6666/joy-caption-alpha-two-cli-mod/tree/main) , and copy to ```ComfyUI/models/Joy_caption```。
+
+
+Node Options:
+
+
+* image: Image input.
+* extra_options: Input the extra_options.
+* llm_model: There are two LLM models to choose, Orenguteng/Llama-3.1-8B-Lexi-Uncensored-V2 and unsloth/Meta-Llama-3.1-8B-Instruct.
+* device: Model loading device. Currently, only CUDA is supported.
+* dtype: Model precision, nf4 and bf16.
+* vlm_lora: Whether to load text_madel.
+* caption_type: Caption type options, including: "Descriptive", "Descriptive (Informal)", "Training Prompt", "MidJourney", "Booru tag list", "Booru-like tag list", "Art Critic", "Product Listing", "Social Media Post".
+* caption_length: The length of caption.
+* user_prompt: User prompt words for LLM model. If there is content here, it will overwrite all the settings for caption_type and extra_options.
+* max_new_tokens: The max_new_token parameter of LLM.
+* do_sample: The do_sample parameter of LLM.
+* top-p: The top_p parameter of LLM.
+* temperature: The temperature parameter of LLM.
+* cache_model: Whether to cache the model.
+
+### JoyCaption2Split
+The node of JoyCaption2 separate model loading and inference, and when multiple JoyCaption2 nodes are used, the model can be shared to improve efficiency.
+
+Node Options:
+
+
+* image: Image input.。
+* joy2_model: The JoyCaption model input.
+* extra_options: Input the extra_options.
+* caption_type: Caption type options, including: "Descriptive", "Descriptive (Informal)", "Training Prompt", "MidJourney", "Booru tag list", "Booru-like tag list", "Art Critic", "Product Listing", "Social Media Post".
+* caption_length: The length of caption.
+* user_prompt: User prompt words for LLM model. If there is content here, it will overwrite all the settings for caption_type and extra_options.
+* max_new_tokens: The max_new_token parameter of LLM.
+* do_sample: The do_sample parameter of LLM.
+* top-p: The top_p parameter of LLM.
+* temperature: The temperature parameter of LLM.
+
+### LoadJoyCaption2Model
+JoyCaption2's model loading node, used in conjunction with JoyCaption2Split.
+
+Node Options:
+
+
+* llm_model: There are two LLM models to choose, Orenguteng/Llama-3.1-8B-Lexi-Uncensored-V2 and unsloth/Meta-Llama-3.1-8B-Instruct.
+* device: Model loading device. Currently, only CUDA is supported.
+* dtype: Model precision, nf4 and bf16.
+* vlm_lora: Whether to load text_madel.
+
+### JoyCaption2ExtraOptions
+The extra_options parameter node of JoyCaption2.
+
+Node Options:
+
+
+* refer_character_name: If there is a person/character in the image you must refer to them as {name}.
+* exclude_people_info: Do NOT include information about people/characters that cannot be changed (like ethnicity, gender, etc), but do still include changeable attributes (like hair style).
+* include_lighting: Include information about lighting.
+* include_camera_angle: Include information about camera angle.
+* include_watermark: Include information about whether there is a watermark or not.
+* include_JPEG_artifacts: Include information about whether there are JPEG artifacts or not.
+* include_exif: If it is a photo you MUST include information about what camera was likely used and details such as aperture, shutter speed, ISO, etc.
+* exclude_sexual: Do NOT include anything sexual; keep it PG.
+* exclude_image_resolution: Do NOT mention the image's resolution.
+* include_aesthetic_quality: You MUST include information about the subjective aesthetic quality of the image from low to very high.
+* include_composition_style: Include information on the image's composition style, such as leading lines, rule of thirds, or symmetry.
+* exclude_text: Do NOT mention any text that is in the image.
+* specify_depth_field: Specify the depth of field and whether the background is in focus or blurred.
+* specify_lighting_sources: If applicable, mention the likely use of artificial or natural lighting sources.
+* do_not_use_ambiguous_language: Do NOT use any ambiguous language.
+* include_nsfw: Include whether the image is sfw, suggestive, or nsfw.
+* only_describe_most_important_elements: ONLY describe the most important elements of the image.
+* character_name: Person/Character Name, if choice ```refer_character_name```.
+
+### PhiPrompt
+
+Use Microsoft Phi 3.5 text and visual models for local inference. Can be used to generate prompt words, process prompt words, or infer prompt words from images. Running this model requires at least 16GB of video memory.
+Download model files from [BaiduNetdisk](https://pan.baidu.com/s/1BdTLdaeGC3trh1U3V-6XTA?pwd=29dh) or [huggingface.co/microsoft/Phi-3.5-vision-instruct](https://huggingface.co/microsoft/Phi-3.5-vision-instruct/tree/main) and [huggingface.co/microsoft/Phi-3.5-mini-instruct](https://huggingface.co/microsoft/Phi-3.5-mini-instruct/tree/main) and copy to ```ComfyUI\models\LLM``` folder.
+
+
+Node Options:
+
+
+* image: Optional input. The input image will serve as the input for Phi-3.5-vision-instruct.
+* model: Selectable to load Phi-3.5-vision-instruct or Phi-3.5-mini-instruct model. The default value of auto will automatically load the corresponding model based on whether there is image input.
+* device: Model loading device. Supports CPU and CUDA.
+* dtype: The model loading accuracy has three options: fp16, bf16, and fp32.
+* cache_model: Whether to cache the model.
+* system_prompt: The system prompt of Phi-3.5-mini-instruct.
+* user_prompt: User prompt words for LLM model.
+* do_sample: The do_Sample parameter of LLM defaults to True.
+* temperature: The temperature parameter of LLM defaults to 0.5.
+* max_new_tokens: The max_new_token parameter of LLM defaults to 512.
+
+### UserPromptGeneratorTxtImg
+
+UserPrompt preset for generating SD text to image prompt words.
+
+Node options:
+
+
+* template: Prompt word template. Currently, only the 'SD txt2img prompt' is available.
+* describe: Prompt word description. Enter a simple description here.
+* limit_word: Maximum length limit for output prompt words. For example, 200 means that the output text will be limited to 200 words.
+
+### UserPromptGeneratorTxtImgWithReference
+
+UserCompt preset for generating SD text to image prompt words based on input content.
+
+Node options:
+
+
+* reference_text: Reference text input. Usually it is a style description of the image.
+* template: Prompt word template. Currently, only the 'SD txt2img prompt' is available.
+* describe: Prompt word description. Enter a simple description here.
+* limit_word: Maximum length limit for output prompt words. For example, 200 means that the output text will be limited to 200 words.
+
+### UserPromptGeneratorReplaceWord
+
+UserPrompt preset used to replace a keyword in text with different content. This is not only a simple replacement, but also a logical sorting of the text based on the context of the prompt words to achieve the rationality of the output content.
+
+Node options:
+
+
+* orig_prompt: Original prompt word input.
+* template: Prompt word template. Currently, only 'prompt replace word' is available.
+* exclude_word: Keywords that need to be excluded.
+* replace_with_word: That word will replace the exclude_word.
+
+### PromptTagger
+
+Inference the prompts based on the image. it can replace key word for the prompt. This node currently uses Google Gemini API as the backend service. Please ensure that the network environment can use Gemini normally.
+Please apply for your API key on [Google AI Studio](https://makersuite.google.com/app/apikey), And fill it in ```api_key.ini```, this file is located in the root directory of the plug-in, and the default name is ```api_key.ini.example```. to use this file for the first time, you need to change the file suffix to ```.ini```. Open it using text editing software, fill in your API key after ```google_api_key=``` and save it.
+
+
+Node options:
+
+
+* api: The Api used. At present, there are two options "gemini-1. 5-flash" and "google-gemini".
+* token_limit: The maximum token limit for generating prompt words.
+* exclude_word: Keywords that need to be excluded.
+* replace_with_word: That word will replace the exclude_word.
+
+### PromptEmbellish
+
+Enter simple prompt words, output polished prompt words, and support inputting images as references, and support Chinese input. This node currently uses Google Gemini API as the backend service. Please ensure that the network environment can use Gemini normally.
+Please apply for your API key on [Google AI Studio](https://makersuite.google.com/app/apikey), And fill it in ```api_key.ini```, this file is located in the root directory of the plug-in, and the default name is ```api_key.ini.example```. to use this file for the first time, you need to change the file suffix to ```.ini```. Open it using text editing software, fill in your API key after ```google_api_key=``` and save it.
+
+
+Node options:
+
+
+* image: Optional, input image as a reference for prompt words.
+* api: The Api used. At present, there are two options "gemini-1. 5-flash" and "google-gemini".
+* token_limit: The maximum token limit for generating prompt words.
+* discribe: Enter a simple description here. supports Chinese text input.
+
+### Florence2Image2Prompt
+
+Use the Florence 2 model to infer prompt words. The code for this node section is from[yiwangsimple/florence_dw](https://github.com/yiwangsimple/florence_dw), thanks to the original author.
+*When using it for the first time, the model will be automatically downloaded. You can also download the model file from [BaiduNetdisk](https://pan.baidu.com/s/1hzw9-QiU1vB8pMbBgofZIA?pwd=mfl3) to ```ComfyUI/models/florence2``` folder.
+
+
+Node Options:
+
+
+* florence2_model: Florence2 model input.
+* image: Image input.
+* task: Select the task for florence2.
+* text_input: Text input for florence2.
+* max_new_tokens: The maximum number of tokens for generating text.
+* num_beams: The number of beam searches that generate text.
+* do_sample: Whether to use text generated sampling.
+* fill_mask: Whether to use text marker mask filling.
+
+
+
+### GetColorTone
+
+Obtain the main color or average color from the image and output RGB values.
+
+
+Node options:
+
+
+* mode: There are two modes to choose from, with the main color and average color.
+
+Output type:
+
+* RGB color in HEX: The RGB color described by hexadecimal RGB format, like '#FA3D86'.
+* HSV color in list: The HSV color described by python's list data format.
+
+### GetColorToneV2
+
+V2 upgrade of GetColorTone. You can specify the dominant or average color to get the body or background.
+
+
+The following changes have been made on the basis of GetColorTong:
+
+
+* color_of: Provides 4 options, mask, entire, background, and subject, to select the color of the mask area, entire picture, background, or subject, respectively.
+* remove_background_method: There are two methods of background recognition: BiRefNet and RMBG V1.4.
+* invert_mask: Whether to reverse the mask.
+* mask_grow: Mask expansion. For subject, a larger value brings the obtained color closer to the color at the center of the body.
+
+Output:
+
+* image: Solid color picture output, the size is the same as the input picture.
+* mask: Mask output.
+
+
+### ImageRewardFilter
+
+
+Rating bulk pictures and outputting top-ranked pictures. it used [ImageReward] (https://github.com/THUDM/ImageReward) for image scoring, thanks to the original authors.
+
+
+Node options:
+
+* prompt: Optional input. Entering prompt here will be used as a basis to determine how well it matches the picture.
+* output_nun: Number of pictures outputted. This value should be less than the picture batch.
+
+Outputs:
+
+* images: Bulk pictures output from high to low in order of rating.
+* obsolete_images: Knockout pictures. Also output in order of rating from high to low.
+
+
+### LaMa
+
+
+Erase objects from the image based on the mask. this node is repackage of [IOPaint](https://www.iopaint.com), powered by state-of-the-art AI models, thanks to the original author.
+It is have [LaMa](https://github.com/advimman/lama), [LDM](https://github.com/CompVis/latent-diffusion), [ZITS](https://github.com/DQiaole/ZITS_inpainting),[MAT](https://github.com/fenglinglwb/MAT), [FcF](https://github.com/SHI-Labs/FcF-Inpainting), [Manga](https://github.com/msxie92/MangaInpainting) models and the SPREAD method to erase. Please refer to the original link for the introduction of each model.
+Please download the model files from [lama models(BaiduNetdisk)](https://pan.baidu.com/s/1m7La2ELsSKaIFhQ57qg1XQ?pwd=jn10) or [lama models(Google Drive)](https://drive.google.com/drive/folders/1Aq0a4sybb3SRxi7j1e1_ZbBRjaWDdP9e?usp=sharing) to ```ComfyUI/models/lama``` folder.
+
+Node optons:
+
+
+* lama_model: Choose a model or method.
+* device: After correctly installing Torch and Nvidia CUDA drivers, using cuda will significantly improve running speed.
+* invert_mask: Whether to reverse the mask.
+* grow: Positive values expand outward, while negative values contract inward.
+* blur: Blur the edge.
+
+
+
+### ImageAutoCrop
+
+
+Automatically cutout and crop the image according to the mask. it can specify the background color, aspect ratio, and size for output image. this node is designed to generate the image materials for training models.
+*Please refer to the model installation methods for [SegmentAnythingUltra](#SegmentAnythingUltra) and [RemBgUltra](#RemBgUltra).
+
+Node options:
+
+
+* background_color4: The background color.
+* aspect_ratio: Here are several common frame ratios provided. alternatively, you can choose "original" to keep original ratio or customize the ratio using "custom".
+* proportional_width: Proportional width. if the aspect ratio option is not "custom", this setting will be ignored.
+* proportional_height: Proportional height. if the aspect ratio option is not "custom", this setting will be ignored.
+* scale_by_longest_side: Allow scaling by long edge size.
+* longest_side: When the scale_by_longest_side is set to True, this will be used this value to the long edge of the image. when the original_size have input, this setting will be ignored.
+* detect: Detection method, min_bounding_rect is the minimum bounding rectangle, max_inscribed_rect is the maximum inscribed rectangle.
+* border_reserve: Keep the border. expand the cutting range beyond the detected mask body area.
+* ultra_detail_range: Mask edge ultra fine processing range, 0 is not processed, which can save generation time.
+* matting_method: The method of generate masks. There are two methods available: Segment Anything and RMBG 1.4. RMBG 1.4 runs faster.
+* sam_model: Select the SAM model used by Segment Anything here.
+* grounding_dino_model: Select the Grounding_Dino model used by Segment Anything here.
+* sam_threshold: The threshold for Segment Anything.
+* sam_prompt: The prompt for Segment Anything.
+
+Output:
+cropped_image: Crop and replace the background image.
+box_preview: Crop position preview.
+cropped_mask: Cropped mask.
+
+### ImageAutoCropV2
+
+The V2 upgrad version of ```ImageAutoCrop```, it has made the following changes based on the previous version:
+
+
+* Add optional input for mask. when there is a mask input, use that input directly to skip the built-in mask generation.
+* Add ```fill_background```. When set to False, the background will not be processed and any parts beyond the frame will not be included in the output range.
+* ```aspect_ratio``` adds the ```original``` option.
+* scale_by: Allow scaling by specified dimensions for longest, shortest, width, or height.
+* scale_by_length: The value here is used as ```scale_by``` to specify the length of the edge.
+
+### ImageAutoCropV3
+
+Automatically crop the image to the specified size. You can input a mask to preserve the specified area of the mask. This node is designed to generate image materials for training the model.
+
+Node Options:
+
+
+* image: The input image.
+* mask: Optional input mask. The masking part will be preserved within the range of the cutting aspect ratio.
+* aspect_ratio: The aspect ratio of the output. Here are common frame ratios provided, with "custom" being the custom ratio and "original" being the original frame ratio.
+* proportional_width: Proportionally wide. If the aspect_ratio option is not 'custom', this setting will be ignored.
+* proportional_height: High proportion. If the aspect_ratio option is not 'custom', this setting will be ignored.
+* method: Scaling sampling methods include Lanczos, Bicubic, Hamming, Bilinear, Box, and Nearest.
+* scale_to_side: Allow scaling to be specified by long side, short side, width, height, or total pixels.
+* scale_to_length: The value here is used as the scale_to-side to specify the length of the edge or the total number of pixels (kilo pixels).
+* round_to_multiple: Multiply to the nearest whole. For example, if set to 8, the width and height will be forcibly set to multiples of 8.
+
+Outputs:
+cropped_image: The cropped image.
+box_preview: Preview of cutting position.
+
+
+
+### SaveImagePlus
+
+
+Enhanced save image node. You can customize the directory where the picture is saved, add a timestamp to the file name, select the save format, set the image compression rate, set whether to save the workflow, and optionally add invisible watermarks to the picture. (Add information in a way that is invisible to the naked eye, and use the ```ShowBlindWaterMark``` node to decode the watermark). Optionally output the json file of the workflow.
+
+Node Options:
+
+
+* iamge: The input image.
+* custom_path*: User-defined directory, enter the directory name in the correct format. If empty, it is saved in the default output directory of ComfyUI.
+* filename_prefix*: The prefix of file name.
+* timestamp: Timestamp the file name, opting for date, time to seconds, and time to milliseconds.
+* format: The format of image save. Currently available in ```png``` and ```jpg```. Note that only png format is supported for RGBA mode pictures.
+* quality: Image quality, the value range 10-100, the higher the value, the better the picture quality, the volume of the file also correspondingly increases.
+* meta_data: Whether to save metadata to png file, that is workflow information. Set this to false if you do not want the workflow to be leaked.
+* blind_watermark: The text entered here (does not support multilingualism) will be converted into a QR code and saved as an invisible watermark. Use ```ShowBlindWaterMark``` node can decode watermarks. Note that pictures with watermarks are recommended to be saved in png format, and lower-quality jpg format will cause watermark information to be lost.
+* save_workflow_as_json: Whether the output workflow is a json file at the same time (the output json is in the same directory as the picture).
+* preview: Preview switch.
+
+* Enter```%date``` for the current date (YY-mm-dd) and ```%time``` for the current time (HH-MM-SS). You can enter ```/``` for subdirectories. For example, ```%date/name_%tiem``` will output the image to the ```YY-mm-dd``` folder, with ```name_HH-MM-SS``` as the file name prefix.
+
+
+
+### AddBlindWaterMark
+
+
+Add an invisible watermark to a picture. Add the watermark image in a way that is invisible to the naked eye, and use the ```ShowBlindWaterMark``` node to decode the watermark.
+
+Node Options:
+
+
+* iamge: The input image.
+* watermark_image: Watermark image. The image entered here will automatically be converted to a square black and white image as a watermark. It is recommended to use a QR code as a watermark.
+
+### ShowBlindWaterMark
+
+Decoding the invisible watermark added to the ```AddBlindWaterMark``` and ```SaveImagePlus``` nodes.
+
+
+### CreateQRCode
+
+Generate a square QR code picture.
+
+Node Options:
+
+
+* size: The side length of image.
+* border: The size of the border around the QR code, the larger the value, the wider the border.
+* text: Enter the text content of the QR code here, and multi-language is not supported.
+
+### DecodeQRCode
+
+Decoding the QR code.
+
+Node Options:
+
+
+* image: The input QR code image.
+* pre_blur: Pre-blurring, you can try to adjust this value for QR codes that are difficult to identify.
+
+### LoadPSD
+
+
+
+Load the PSD format file and export the layers.
+Note that this node requires the installation of the ```psd_tools``` dependency package, If error occurs during the installation of psd_tool, such as ```ModuleNotFoundError: No module named 'docopt'``` , please download [docopt's whl](https://www.piwheels.org/project/docopt/) and manual install it.
+
+Node Options:
+
+
+* image: Here is a list of *.psd files under ```ComfyUI/input```, where previously loaded psd images can be selected.
+* file_path: The complete path and file name of the psd file.
+* include_hidden_layer: whether include hidden layers.
+* find_layer_by: The method for finding layers can be selected by layer key number or layer name. Layer groups are treated as one layer.
+* layer_index: The layer key number, where 0 is the bottom layer, is incremented sequentially. If include_hiddenlayer is set to false, hidden layers are not counted. Set to -1 to output the top layer.
+* layer_name: Layer name. Note that capitalization and punctuation must match exactly.
+
+Outputs:
+flat_image: PSD preview image.
+layer_iamge: Find the layer output.
+all_layers: Batch images containing all layers.
+
+### SD3NegativeConditioning
+
+
+Encapsulate the four nodes of Negative Condition in SD3 into a separate node.
+
+Node Options:
+
+
+* zero_out_start: Set the ConditioningSetTimestepRange start value for Negative ConditioningZeroOut, which is the same as the ConditioningSetTimestepRange end value for Negative.
+
+
+
+### SegmentAnythingUltra
+
+Improvements to [ComfyUI Segment Anything](https://github.com/storyicon/comfyui_segment_anything), thanks to the original author.
+
+*Please refer to the installation of ComfyUI Segment Anything to install the model. If ComfyUI Segment Anything has been correctly installed, you can skip this step.
+
+* From [here](https://huggingface.co/bert-base-uncased/tree/main) download the config.json,model.safetensors,tokenizer_config.json,tokenizer.json and vocab.txt 5 files to ```ComfyUI/models/bert-base-uncased``` folder.
+* Download [GroundingDINO_SwinT_OGC config file](https://huggingface.co/ShilongLiu/GroundingDINO/resolve/main/GroundingDINO_SwinT_OGC.cfg.py), [GroundingDINO_SwinT_OGC model](https://huggingface.co/ShilongLiu/GroundingDINO/resolve/main/groundingdino_swint_ogc.pth),
+ [GroundingDINO_SwinB config file](https://huggingface.co/ShilongLiu/GroundingDINO/resolve/main/GroundingDINO_SwinB.cfg.py), [GroundingDINO_SwinB model](https://huggingface.co/ShilongLiu/GroundingDINO/resolve/main/groundingdino_swinb_cogcoor.pth) to ```ComfyUI/models/grounding-dino``` folder.
+* Download [sam_vit_h](https://dl.fbaipublicfiles.com/segment_anything/sam_vit_h_4b8939.pth),[sam_vit_l](https://dl.fbaipublicfiles.com/segment_anything/sam_vit_l_0b3195.pth),
+ [sam_vit_b](https://dl.fbaipublicfiles.com/segment_anything/sam_vit_b_01ec64.pth), [sam_hq_vit_h](https://huggingface.co/lkeab/hq-sam/resolve/main/sam_hq_vit_h.pth),
+ [sam_hq_vit_l](https://huggingface.co/lkeab/hq-sam/resolve/main/sam_hq_vit_l.pth), [sam_hq_vit_b](https://huggingface.co/lkeab/hq-sam/resolve/main/sam_hq_vit_b.pth),
+ [mobile_sam](https://github.com/ChaoningZhang/MobileSAM/blob/master/weights/mobile_sam.pt) to ```ComfyUI/models/sams``` folder.
+ *Or download them from [GroundingDino models on BaiduNetdisk](https://pan.baidu.com/s/1P7WQDuaqSYazlSQX8SJjxw?pwd=24ki) and [SAM models on BaiduNetdisk](https://pan.baidu.com/s/1n7JrHb2vzV2K2z3ktqpNxg?pwd=yoqh) .
+ 
+ 
+
+Node options:
+
+
+* sam_model: Select the SAM model.
+* ground_dino_model: Select the Grounding DINO model.
+* threshold: The threshold of SAM.
+* detail_range: Edge detail range.
+* black_point: Edge black sampling threshold.
+* white_point: Edge white sampling threshold.
+* process_detail: Set to false here will skip edge processing to save runtime.
+* prompt: Input for SAM's prompt.
+* cache_model: Set whether to cache the model.
+
+### SegmentAnythingUltraV2
+
+The V2 upgraded version of SegmentAnythingUltra has added the VITMatte edge processing method.(Note: Images larger than 2K in size using this method will consume huge memory)
+
+
+On the basis of SegmentAnythingUltra, the following changes have been made:
+
+
+* detail_method: Edge processing methods. provides VITMatte, VITMatte(local), PyMatting, GuidedFilter. If the model has been downloaded after the first use of VITMatte, you can use VITMatte (local) afterwards.
+* detail_erode: Mask the erosion range inward from the edge. the larger the value, the larger the range of inward repair.
+* detail_dilate: The edge of the mask expands outward. the larger the value, the wider the range of outward repair.
+* device: Set whether the VitMatte to use cuda.
+* max_megapixels: Set the maximum size for VitMate operations.
+
+### SAM2Ultra
+
+This node is modified from [kijai/ComfyUI-segment-anything-2](https://github.com/kijai/ComfyUI-segment-anything-2). Thank to [kijai](https://github.com/kijai) for making significant contributions to the Comfyui community.
+SAM2 Ultra node only support single image. If you need to process multiple images, please first convert the image batch to image list.
+*Download models from [BaiduNetdisk](https://pan.baidu.com/s/1xaQYBA6ktxvAxm310HXweQ?pwd=auki) or [huggingface.co/Kijai/sam2-safetensors](https://huggingface.co/Kijai/sam2-safetensors/tree/main) and copy to ```ComfyUI/models/sam2``` folder.
+
+
+
+Node Options:
+
+
+* image: The image to segment.
+* bboxes: Input recognition box data.
+* sam2_model: Select the SAM2 model.
+* presicion: Model's persicion. can be selected from fp16, bf16, and fp32.
+* bbox_select: Select the input box data. There are three options: "all" to select all, "first" to select the box with the highest confidence, and "by_index" to specify the index of the box.
+* select_index: This option is valid when bbox_delect is 'by_index'. 0 is the first one. Multiple values can be entered, separated by any non numeric character, including but not limited to commas, periods, semicolons, spaces or letters, and even Chinese.
+* cache_model: Whether to cache the model. After caching the model, it will save time for model loading.
+* detail_method: Edge processing methods. provides VITMatte, VITMatte(local), PyMatting, GuidedFilter. If the model has been downloaded after the first use of VITMatte, you can use VITMatte (local) afterwards.
+* detail_erode: Mask the erosion range inward from the edge. the larger the value, the larger the range of inward repair.
+* detail_dilate: The edge of the mask expands outward. the larger the value, the wider the range of outward repair.
+* black_point: Edge black sampling threshold.
+* white_point: Edge white sampling threshold.
+* process_detail: Set to false here will skip edge processing to save runtime.
+* device: Set whether the VitMatte to use cuda.
+* max_megapixels: Set the maximum size for VitMate operations.
+
+### SAM2VideoUltra
+
+SAM2 Video Ultra node support processing multiple frames of images or video sequences. Please define the recognition box data in the first frame of the sequence to ensure correct recognition.
+
+https://github.com/user-attachments/assets/4726b8bf-9b98-4630-8f54-cb7ed7a3d2c5
+
+https://github.com/user-attachments/assets/b2a45c96-4be1-4470-8ceb-addaf301b0cb
+
+Node Options:
+
+
+* image: The image to segment.
+* bboxes: Optional input of recognition bbox data. ```bboxes``` and ```first_frame_mask``` must have least one input. If first_frame_mask inputed, bbboxes will be ignored.
+* first_frame_mask: Optional input of the first frame mask. The mask will be used as the first frame recognition object. ```bboxes``` and ```first_frame_mask``` must have least one input. If first_frame_mask inputed, bbboxes will be ignored.
+* pre_mask: Optional input mask, which will serve as a propagation focus range limitation and help improve recognition accuracy.
+* sam2_model: Select the SAM2 model.
+* presicion: Model's persicion. can be selected from fp16 and bf16.
+* cache_model: Whether to cache the model. After caching the model, it will save time for model loading.
+* individual_object: When set to True, it will focus on identifying a single object. When set to False, attempts will be made to generate recognition boxes for multiple objects.
+* mask_preview_color: Display the color of non masked areas in the preview output.
+* detail_method: Edge processing methods. Only VITMatte method can be used.
+* detail_erode: Mask the erosion range inward from the edge. the larger the value, the larger the range of inward repair.
+* detail_dilate: The edge of the mask expands outward. the larger the value, the wider the range of outward repair.
+* black_point: Edge black sampling threshold.
+* white_point: Edge white sampling threshold.
+* process_detail: Set to false here will skip edge processing to save runtime.
+* device: Only cuda can be used.
+* max_megapixels: Set the maximum size for VitMate operations.A larger size will result in finer mask edges, but it will lead to a significant decrease in computation speed.
+
+### ObjectDetectorFL2
+
+Use the Florence2 model to identify objects in images and output recognition box data.
+*Download models from [BaiduNetdisk](https://pan.baidu.com/s/1hzw9-QiU1vB8pMbBgofZIA?pwd=mfl3) and copy to ```ComfyUI/models/florence2``` folder.
+
+Node Options:
+
+
+* image: The image to segment.
+* florence2_model: Florence2 model, it from [LoadFlorence2Model](#LoadFlorence2Model) node.
+* prompt: Describe the object that needs to be identified.
+* sort_method: The selection box sorting method has 4 options: "left_to_right", "top_to_bottom", "big_to_small" and "confidence".
+* bbox_select: Select the input box data. There are three options: "all" to select all, "first" to select the box with the highest confidence, and "by_index" to specify the index of the box.
+* select_index: This option is valid when bbox_delect is 'by_index'. 0 is the first one. Multiple values can be entered, separated by any non numeric character, including but not limited to commas, periods, semicolons, spaces or letters, and even Chinese.
+
+### ObjectDetectorYOLOWorld(Obsoleted. If you want to continue using it, you need to manually install the dependency package)
+Due to potential installation issues with dependency packages, this node has been obsoleted. To use, please manually install the following dependency packages:
+```
+pip install inference-cli>=0.13.0
+pip install inference-gpu[yolo-world]>=0.13.0
+```
+
+Use the YOLO-World model to identify objects in images and output recognition box data.
+*Download models from [BaiduNetdisk](https://pan.baidu.com/s/1QpjajeTA37vEAU2OQnbDcQ?pwd=nqsk) or [GoogleDrive](https://drive.google.com/drive/folders/1nrsfq4S-yk9ewJgwrhXAoNVqIFLZ1at7?usp=sharing) and copy to ```ComfyUI/models/yolo-world``` folder.
+
+Node Options:
+
+
+* image: The image to segment.
+* confidence_threshold: The threshold of confidence.
+* nms_iou_threshold: The threshold of Non-Maximum Suppression.
+* prompt: Describe the object that needs to be identified.
+* sort_method: The selection box sorting method has 4 options: "left_to_right", "top_to_bottom", "big_to_small" and "confidence".
+* bbox_select: Select the input box data. There are three options: "all" to select all, "first" to select the box with the highest confidence, and "by_index" to specify the index of the box.
+* select_index: This option is valid when bbox_delect is 'by_index'. 0 is the first one. Multiple values can be entered, separated by any non numeric character, including but not limited to commas, periods, semicolons, spaces or letters, and even Chinese.
+
+### ObjectDetectorYOLO8
+
+Use the YOLO-8 model to identify objects in images and output recognition box data.
+*Download models from [GoogleDrive](https://drive.google.com/drive/folders/1I5TISO2G1ArSkKJu1O9b4Uvj3DVgn5d2) or [BaiduNetdisk](https://pan.baidu.com/s/1pEY6sjABQaPs6QtpK0q6XA?pwd=grqe) and copy to ```ComfyUI/models/yolo``` folder.
+
+Node Options:
+
+
+* image: The image to segment.
+* yolo_model: Choose the yolo model.
+* sort_method: The selection box sorting method has 4 options: "left_to_right", "top_to_bottom", "big_to_small" and "confidence".
+* bbox_select: Select the input box data. There are three options: "all" to select all, "first" to select the box with the highest confidence, and "by_index" to specify the index of the box.
+* select_index: This option is valid when bbox_delect is 'by_index'. 0 is the first one. Multiple values can be entered, separated by any non numeric character, including but not limited to commas, periods, semicolons, spaces or letters, and even Chinese.
+
+### ObjectDetectorMask
+
+Use mask as recognition box data. All areas surrounded by white areas on the mask will be recognized as an object. Multiple enclosed areas will be identified separately.
+
+Node Options:
+
+
+* object_mask: The mask input.
+* sort_method: The selection box sorting method has 4 options: "left_to_right", "top_to_bottom", "big_to_small" and "confidence".
+* bbox_select: Select the input box data. There are three options: "all" to select all, "first" to select the box with the highest confidence, and "by_index" to specify the index of the box.
+* select_index: This option is valid when bbox_delect is 'by_index'. 0 is the first one. Multiple values can be entered, separated by any non numeric character, including but not limited to commas, periods, semicolons, spaces or letters, and even Chinese.
+
+### BBoxJoin
+
+Merge recognition box data.
+
+Node Options:
+
+
+* bboxes_1: Required input. The first set of identification boxes.
+* bboxes_2: Optional input. The second set of identification boxes.
+* bboxes_3: Optional input. The third set of identification boxes.
+* bboxes_4: Optional input. The fourth set of identification boxes.
+
+### DrawBBoxMask
+
+Draw the recognition BBoxes data output by the Object Detector node as a mask.
+
+
+Node Options:
+
+
+* image: Image input. It must be consistent with the image recognized by the Object Detector node.
+* bboxes: Input recognition BBoxes data.
+* grow_top: Each BBox expands upwards as a percentage of its height, positive values indicate upward expansion and negative values indicate downward expansion.
+* grow_bottom: Each BBox expands downwards as a percentage of its height, positive values indicating downward expansion and negative values indicating upward expansion.
+* grow_left: Each BBox expands to the left as a percentage of its width, positive values expand to the left and negative values expand to the right.
+* grow_right: Each BBox expands to the right as a percentage of its width, positive values indicate expansion to the right and negative values indicate expansion to the left.
+
+### EVF-SAMUltra
+
+This node is implementation of [EVF-SAM](https://github.com/hustvl/EVF-SAM) in ComfyUI.
+*Please download model files from [BaiduNetdisk](https://pan.baidu.com/s/1EvaxgKcCxUpMbYKzLnEx9w?pwd=69bn) or [huggingface/EVF-SAM2](https://huggingface.co/YxZhang/evf-sam2/tree/main), [huggingface/EVF-SAM](https://huggingface.co/YxZhang/evf-sam/tree/main) to ```ComfyUI/models/EVF-SAM``` folder(save the models in their respective subdirectories).
+
+
+Node Options:
+
+
+* image: The input image.
+* model: Select the model. Currently, there are options for evf-sam2 and evf sam.
+* presicion: Model accuracy can be selected from fp16, bf16, and fp32.
+* load_in_bit: Load the model with positional accuracy. You can choose from full, 8, and 4.
+* pormpt: Prompt words used for segmentation.
+* detail_method: Edge processing methods. provides VITMatte, VITMatte(local), PyMatting, GuidedFilter. If the model has been downloaded after the first use of VITMatte, you can use VITMatte (local) afterwards.
+* detail_erode: Mask the erosion range inward from the edge. the larger the value, the larger the range of inward repair.
+* detail_dilate: The edge of the mask expands outward. the larger the value, the wider the range of outward repair.
+* black_point: Edge black sampling threshold.
+* white_point: Edge white sampling threshold.
+* process_detail: Set to false here will skip edge processing to save runtime.
+* device: Set whether the VitMatte to use cuda.
+* max_megapixels: Set the maximum size for VitMate operations.
+
+### Florence2Ultra
+
+Using the segmentation function of the Florence2 model, while also having ultra-high edge details.
+The code for this node section is from [spacepxl/ComfyUI-Florence-2](https://github.com/spacepxl/ComfyUI-Florence-2), thanks to the original author.
+*Download the model files from [BaiduNetdisk](https://pan.baidu.com/s/1hzw9-QiU1vB8pMbBgofZIA?pwd=mfl3) to ```ComfyUI/models/florence2``` folder.
+
+
+
+Node Options:
+
+
+* florence2_model: Florence2 model input.
+* image: Image input.
+* task: Select the task for florence2.
+* text_input: Text input for florence2.
+* detail_method: Edge processing methods. provides VITMatte, VITMatte(local), PyMatting, GuidedFilter. If the model has been downloaded after the first use of VITMatte, you can use VITMatte (local) afterwards.
+* detail_erode: Mask the erosion range inward from the edge. the larger the value, the larger the range of inward repair.
+* detail_dilate: The edge of the mask expands outward. the larger the value, the wider the range of outward repair.
+* black_point: Edge black sampling threshold.
+* white_point: Edge white sampling threshold.
+* process_detail: Set to false here will skip edge processing to save runtime.
+* device: Set whether the VitMatte to use cuda.
+* max_megapixels: Set the maximum size for VitMate operations.
+
+### LoadFlorence2Model
+
+Florence2 model loader.
+*When using it for the first time, the model will be automatically downloaded.
+
+
+At present, there are base, base-ft, large, large-ft, DocVQA, SD3-Captioner and base-PromptGen models to choose from.
+
+
+
+### BiRefNetUltra
+
+Using the BiRefNet model to remove background has better recognition ability and ultra-high edge details.
+The code for the model part of this node comes from Viper's [ComfyUI-BiRefNet](https://github.com/viperyl/ComfyUI-BiRefNet),thanks to the original author.
+
+*From [https://huggingface.co/ViperYX/BiRefNet](https://huggingface.co/ViperYX/BiRefNet/tree/main) or [BaiduNetdisk](https://pan.baidu.com/s/1GxtuNDTIHkuu4FR4uGAT-g?pwd=t2cf) download the ```BiRefNet-ep480.pth```,```pvt_v2_b2.pth```,```pvt_v2_b5.pth```,```swin_base_patch4_window12_384_22kto1k.pth```, ```swin_large_patch4_window12_384_22kto1k.pth``` 5 files to ```ComfyUI/models/BiRefNet``` folder.
+
+
+
+Node options:
+
+
+* detail_method: Edge processing methods. provides VITMatte, VITMatte(local), PyMatting, GuidedFilter. If the model has been downloaded after the first use of VITMatte, you can use VITMatte (local) afterwards.
+* detail_erode: Mask the erosion range inward from the edge. the larger the value, the larger the range of inward repair.
+* detail_dilate: The edge of the mask expands outward. the larger the value, the wider the range of outward repair.
+* black_point: Edge black sampling threshold.
+* white_point: Edge white sampling threshold.
+* process_detail: Set to false here will skip edge processing to save runtime.
+* device: Set whether the VitMatte to use cuda.
+* max_megapixels: Set the maximum size for VitMate operations.
+
+### BiRefNetUltraV2
+
+This node supports the use of the latest BiRefNet model.
+*Download model file from [BaiduNetdisk](https://pan.baidu.com/s/12z3qUuqag3nqpN2NJ5pSzg?pwd=ek65) or [GoogleDrive](https://drive.google.com/drive/folders/1s2Xe0cjq-2ctnJBR24563yMSCOu4CcxM) named ```BiRefNet-general-epoch_244.pth``` to ```ComfyUI/Models/BiRefNet/pth``` folder. You can also download more BiRefNet models and put them here.
+
+
+
+Node Options:
+
+
+* image: The input image.
+* birefnet_model: The BiRefNet model is input and it is output from the LoadBiRefNetModel node.
+* detail_method: Edge processing methods. provides VITMatte, VITMatte(local), PyMatting, GuidedFilter. If the model has been downloaded after the first use of VITMatte, you can use VITMatte (local) afterwards.
+* detail_erode: Mask the erosion range inward from the edge. the larger the value, the larger the range of inward repair.
+* detail_dilate: The edge of the mask expands outward. the larger the value, the wider the range of outward repair.
+* black_point: Edge black sampling threshold.
+* white_point: Edge white sampling threshold.
+* process_detail: Due to the excellent edge processing of BiRefNet, it is set to False by default here.
+* device: Set whether the VitMatte to use cuda.
+* max_megapixels: Set the maximum size for VitMate operations.
+
+### LoadBiRefNetModel
+
+Load the BiRefNet model.
+
+Node Options:
+
+
+* model: Select the model. List the files in the ```CoomfyUI/models/BiRefNet/pth``` folder for selection.
+
+
+### LoadBiRefNetModelV2
+This node is a PR submitted by [jimlee2048](https://github.com/jimlee2048) and supports loading RMBG-2.0 models.
+
+Download model files from [huggingface](https://huggingface.co/briaai/RMBG-2.0/tree/main) or [百度网盘](https://pan.baidu.com/s/1viIXlZnpTYTKkm2F-QMj_w?pwd=axr9) and copy to ```ComfyUI/models/BiRefNet/RMBG-2.0``` folder.
+
+Node Options:
+
+
+* model: Select the model. There are two options, ```BiRefNet-General``` and ```RMBG-2.0```.
+
+
+### TransparentBackgroundUltra
+
+Using the transparent-background model to remove background has better recognition ability and speed, while also having ultra-high edge details.
+
+*From [googledrive](https://drive.google.com/drive/folders/10KBDY19egb8qEQBv34cqIVSwd38bUAa9?usp=sharing) or [BaiduNetdisk](https://pan.baidu.com/s/10JO0uKzTxJaIkhN_J7RSyw?pwd=v0b0) download all files to ```ComfyUI/models/transparent-background``` folder.
+
+
+
+Node Options:
+
+
+* model: Select the model.
+* detail_method: Edge processing methods. provides VITMatte, VITMatte(local), PyMatting, GuidedFilter. If the model has been downloaded after the first use of VITMatte, you can use VITMatte (local) afterwards.
+* detail_erode: Mask the erosion range inward from the edge. the larger the value, the larger the range of inward repair.
+* detail_dilate: The edge of the mask expands outward. the larger the value, the wider the range of outward repair.
+* black_point: Edge black sampling threshold.
+* white_point: Edge white sampling threshold.
+* process_detail: Set to false here will skip edge processing to save runtime.
+* device: Set whether the VitMatte to use cuda.
+* max_megapixels: Set the maximum size for VitMate operations.
+
+### PersonMaskUltra
+
+Generate masks for portrait's face, hair, body skin, clothing, or accessories. Compared to the previous A Person Mask Generator node, this node has ultra-high edge details.
+The model code for this node comes from [a-person-mask-generator](https://github.com/djbielejeski/a-person-mask-generator), edge processing code from [ComfyUI-Image-Filters](https://github.com/spacepxl/ComfyUI-Image-Filters),thanks to the original author.
+*Download model files from [BaiduNetdisk](https://pan.baidu.com/s/13zqZtBt89ueCyFufzUlcDg?pwd=jh5g) to ```ComfyUI/models/mediapipe``` folder.
+
+
+
+Node options:
+
+
+* face: Face recognition.
+* hair: Hair recognition.
+* body: Body skin recognition.
+* clothes: Clothing recognition.
+* accessories: Identification of accessories (such as backpacks).
+* background: Background recognition.
+* confidence: Recognition threshold, lower values will output more mask ranges.
+* detail_range: Edge detail range.
+* black_point: Edge black sampling threshold.
+* white_point: Edge white sampling threshold.
+* process_detail: Set to false here will skip edge processing to save runtime.
+
+### PersonMaskUltraV2
+
+The V2 upgraded version of PersonMaskUltra has added the VITMatte edge processing method.(Note: Images larger than 2K in size using this method will consume huge memory)
+
+On the basis of PersonMaskUltra, the following changes have been made:
+
+
+* detail_method: Edge processing methods. provides VITMatte, VITMatte(local), PyMatting, GuidedFilter. If the model has been downloaded after the first use of VITMatte, you can use VITMatte (local) afterwards.
+* detail_erode: Mask the erosion range inward from the edge. the larger the value, the larger the range of inward repair.
+* detail_dilate: The edge of the mask expands outward. the larger the value, the wider the range of outward repair.
+* device: Set whether the VitMatte to use cuda.
+* max_megapixels: Set the maximum size for VitMate operations.
+
+
+### HumanPartsUltra
+
+Used for generate human body parts masks, it is based on the warrper of [metal3d/ComfyUI_Human_Parts](https://github.com/metal3d/ComfyUI_Human_Parts), thank the original author.
+This node has added ultra-fine edge processing based on the original work. Download model file from [BaiduNetdisk](https://pan.baidu.com/s/1-6uwH6RB0FhIVfa3qO7hhQ?pwd=d862) or [huggingface](https://huggingface.co/Metal3d/deeplabv3p-resnet50-human/tree/main) and copy to ```ComfyUI\models\onnx\human-parts``` folder.
+
+
+Node Options:
+
+
+* image: The input image.
+* face: Recognize face switch.
+* hair: Recognize hair switch.
+* galsses: Recognize glasses switch.
+* top_clothes: Recognize top clothes switch.
+* bottom_clothes: Recognize bottom clothes switch.
+* torso_skin: Recognize torso skin switch.
+* left_arm: Recognize left arm switch.
+* right_arm: Recognize right arm switch.
+* left_leg: Recognize left leg switch.
+* right_leg: Recognize right leg switch.
+* left_foot: Recognize left foot switch.
+* right_foot: Recognize right foot switch.
+* detail_method: Edge processing methods. provides VITMatte, VITMatte(local), PyMatting, GuidedFilter. If the model has been downloaded after the first use of VITMatte, you can use VITMatte (local) afterwards.
+* detail_erode: Mask the erosion range inward from the edge. the larger the value, the larger the range of inward repair.
+* detail_dilate: The edge of the mask expands outward. the larger the value, the wider the range of outward repair.
+* black_point: Edge black sampling threshold.
+* white_point: Edge white sampling threshold.
+* process_detail: Set to false here will skip edge processing to save runtime.
+* device: Set whether the VitMatte to use cuda.
+* max_megapixels: Set the maximum size for VitMate operations.
+
+
+
+### YoloV8Detect
+
+Use the YoloV8 model to detect faces, hand box areas, or character segmentation. Supports the output of the selected number of channels.
+Download the model files from [GoogleDrive](https://drive.google.com/drive/folders/1I5TISO2G1ArSkKJu1O9b4Uvj3DVgn5d2) or [BaiduNetdisk](https://pan.baidu.com/s/1pEY6sjABQaPs6QtpK0q6XA?pwd=grqe) to ```ComfyUI/models/yolo``` folder.
+
+
+
+Node Options:
+
+
+* yolo_model: Yolo model selection. the model with ```seg``` name can output segmented masks, otherwise they can only output box masks.
+* mask_merge: Select the merged mask. ```all``` is to merge all mask outputs. The selected number is how many masks to output, sorted by recognition confidence to merge the output.
+
+Outputs:
+
+* mask: The output mask.
+* yolo_plot_image: Preview of yolo recognition results.
+* yolo_masks: For all masks identified by yolo, each individual mask is output as a mask.
+
+### MediapipeFacialSegment
+
+Use the Mediapipe model to detect facial features, segment left and right eyebrows, eyes, lips, and tooth.
+*Download the model files from [BaiduNetdisk](https://pan.baidu.com/s/13zqZtBt89ueCyFufzUlcDg?pwd=jh5g) to ```ComfyUI/models/mediapipe``` folder.
+
+
+
+Node Options:
+
+
+* left_eye: Recognition switch of left eye.
+* left_eyebrow: Recognition switch of left eyebrow.
+* right_eye: Recognition switch of right eye.
+* right_eyebrow: Recognition switch of right eyebrow.
+* lips: Recognition switch of lips.
+* tooth: Recognition switch of tooth.
+
+
+### MaskByDifferent
+
+Calculate the differences between two images and output them as mask.
+
+
+Node options:
+
+
+* gain: The gain of difference calculate. higher value will result in a more significant slight difference.
+* fix_gap: Fix the internal gaps of the mask. higher value will repair larger gaps.
+* fix_threshold: The threshold for fix_gap.
+* main_subject_detect: Setting this to True will enable subject detection, ignoring differences outside of the subject.
+
+
+## Annotation for notes
+
+1 The layer_image, layer_mask and the background_image(if have input), These three items must be of the same size.
+
+2 The mask not a mandatory input item. the alpha channel of the image is used by default. If the image input does not include an alpha channel, the entire image's alpha channel will be automatically created. if have masks input simultaneously, the alpha channel will be overwrite by the mask.
+
+3 The Blend Mode include **normal, multply, screen, add, subtract, difference, darker, color_burn, color_dodge, linear_burn, linear_dodge, overlay, soft_light, hard_light, vivid_light, pin_light, linear_light, and hard_mix.** all of 19 blend modes in total.
+
+*Preview of the blend mode
+
+3 The BlendModeV2 include **normal, dissolve, darken, multiply, color burn, linear burn, darker color, lighten, screen, color dodge, linear dodge(add), lighter color, dodge, overlay, soft light, hard light, vivid light, linear light, pin light, hard mix, difference, exclusion, subtract, divide, hue, saturation, color, luminosity, grain extract, grain merge** all of 30 blend modes in total.
+Part of the code for BlendMode V2 is from [Virtuoso Nodes for ComfyUI](https://github.com/chrisfreilich/virtuoso-nodes). Thanks to the original authors.
+
+*Preview of the Blend Mode V2
+
+4 The RGB color described by hexadecimal RGB format, like '#FA3D86'.
+
+5 The layer_image and layer_mask must be of the same size.
+
+## Stars
+
+[](https://star-history.com/#chflame163/ComfyUI_LayerStyle_Advance&Date)
+
+# statement
+
+LayerStyle Advance nodes follows the MIT license, Some of its functional code comes from other open-source projects. Thanks to the original author. If used for commercial purposes, please refer to the original project license to authorization agreement.
diff --git a/README_CN.MD b/README_CN.MD
new file mode 100644
index 0000000..a2004b4
--- /dev/null
+++ b/README_CN.MD
@@ -0,0 +1,980 @@
+# ComfyUI Layer Style Advance
+
+从ComfyUI Layer Style 剥离出来的节点,主要是一些对依赖包要求较为复杂的节点。
+
+
+## 工作流用示例
+在workflow目录下有json格式的工作流示例文件,示范了如何在ComfyUI中使用这些节点。
+
+
+## 安装方法
+(以ComfyUI官方便携包和秋叶整合包为例,其他ComfyUI环境请修改依赖环境目录)
+
+### 安装插件
+* 推荐使用 ComfyUI Manager 安装。
+* 或者在CompyUI插件目录(例如“CompyUI\custom_nodes\”)中打开cmd窗口,键入
+```
+git clone https://github.com/chflame163/ComfyUI_LayerStyle_Advance.git
+```
+
+* 或者下载解压zip文件,将得到的文件夹复制到 ```ComfyUI\custom_nodes\```。
+
+### 安装依赖包
+
+* 官方便携包请双击运行插件目录下的```install_requirements.bat```,秋叶整合包请双击运行插件目录下的```install_requirements_aki.bat```,然后等待安装完成。
+
+* 或者在资源管理器```ComfyUI\custom_nodes\ComfyUI_LayerStyle_Advance``` 插件目录位置打开cmd窗口,
+
+ 官方便携包输入以下命令:
+
+```
+..\..\..\python_embeded\python.exe -s -m pip install .\whl\docopt-0.6.2-py2.py3-none-any.whl
+..\..\..\python_embeded\python.exe -s -m pip install .\whl\hydra_core-1.3.2-py3-none-any.whl
+..\..\..\python_embeded\python.exe -s -m pip install -r requirements.txt
+.\repair_dependency.bat
+```
+ 秋叶整合包输入以下命令:
+
+```
+..\..\python\python.exe -s -m pip install .\whl\docopt-0.6.2-py2.py3-none-any.whl
+..\..\python\python.exe -s -m pip install .\whl\hydra_core-1.3.2-py3-none-any.whl
+..\..\python\python.exe -s -m pip install -r requirements.txt
+.\repair_dependency.bat
+```
+* 重新打开ComfyUI。
+
+### 下载模型
+国内用户请从[百度网盘](https://pan.baidu.com/s/1T_uXMX3OKIWOJLPuLijrgA?pwd=1yye), 海外用户请从[huggingface](https://huggingface.co/chflame163/ComfyUI_LayerStyle/tree/main),
+下载全部模型文件并复制到```ComfyUI\models```文件夹。这个链接提供了本插件需要的所有的模型文件。
+或者按各个节点的说明下载模型文件。
+
+## 常见问题
+如果节点不能正常加载,或者使用中出现错误,请在ComfyUI终端窗口查看报错信息。以下是常见的错误及解决方法。
+
+### Warning: xxxx.ini not found, use default xxxx..
+这个警告信息是找不到ini文件的提示,不影响使用。如果不想看到这些警告,请修改插件目录下所有的 ```*.ini.example``` 文件名为```*.ini```。
+
+### ModuleNotFoundError: No module named 'psd_tools'
+这个错误是```psd_tools```没有正确安装。
+
+解决方法:
+* 关闭ComfyUI,在插件目录下打开终端窗口,执行以下命令:
+```../../../python_embeded/python.exe -s -m pip install psd_tools```
+如果安装psd_tool中出现```ModuleNotFoundError: No module named 'docopt'```错误,请下载[docopt的whl](https://www.piwheels.org/project/docopt/)手动安装。在终端执行以下命令:
+```../../../python_embeded/python.exe -s -m pip install path/docopt-0.6.2-py2.py3-none-any.whl``` path为whl文件的路径名。
+
+### Cannot import name 'guidedFilter' from 'cv2.ximgproc'
+这个错误是```opencv-contrib-python```没有正确安装,或者安装后又安装了其他opencv包导致。
+
+### NameError: name 'guidedFilter' is not defined
+问题原因同上。
+
+### Cannot import name 'VitMatteImageProcessor' from 'transformers'
+这个错误是由于```transformers``` 版本过低造成的
+
+### insightface 加载缓慢
+这是由于```protobuf``` 版本过低造成的。
+
+#### 以上3个依赖包的问题,请双击运行插件目录下的```repair_dependency.bat```(官方便携包)或者```repair_dependency_aki.bat```(秋叶整合包)自动修复。
+
+### onnxruntime::python::CreateExecutionProviderInstance CUDA_PATH is set but CUDA wasn't able to be loaded. Please install the correct version of CUDA and cuDNN as mentioned in the GPU requirements page
+解决方法:
+请重新安装```onnxruntime```依赖包
+
+### Error loading model xxx: We couldn't connect to huggingface.co ...
+请检查网络环境。如果在中国不能正常访问huggingface.co,请尝试修改huggingface_hub包强制使用hf_mirror镜像。
+* 在```huggingface_hub```包的目录(通常在虚拟环境内的```Lib/site-packages/huggingface_hub```)中找到```constants.py```,
+在```import os```之后增加一行
+```
+os.environ['HF_ENDPOINT'] = 'https://hf-mirror.com'
+```
+
+### ValueError: Trimap did not contain foreground values (xxxx...)
+这个错误是由于使用PyMatting方法处理遮罩边缘时,遮罩面积过大或者过小引起的。
+
+解决方法:
+* 请调整参数,改变遮罩有效面积。或者换用其他的方法处理边缘。
+
+### Requests.exceptions.ProxyError: HTTPSConnectionPool(xxxx...)
+出现这个错误,请检查网络环境。
+
+### UnboundLocalError: local variable 'clip_processor' referenced before assignment
+### UnboundLocalError: local variable 'text_model' referenced before assignment
+如果执行JoyCaption2节点时出现这个报错,同时已确定模型文件已放在正确的目录,请检查```transformers```依赖包版本至少在4.43.2以上。
+如果```transformers```依赖包版本大于等于4.45.0, 并同时有报错信息:
+```
+Error loading models: De️️scriptors cannot be created directly.
+If this call came from a _pb2.py file, your generated code is out of date and must be regenerated with protoc >= 3.19.0.
+......
+```
+请尝试降级```protobuf```依赖包到3.20.3, 或者设置环境变量:```PROTOCOL_BUFFERS_PYTHON_IMPLEMENTATION=python```。
+
+## 如何找到本节点组
+* 在ComfyUI画布点击右键 - Add Node, 找到 "😺dzNodes"。
+
+
+* 或者在ComfyUI画布双击, 在搜索框输入"layer"。
+
+
+
+## 更新说明
+**如果本插件更新后出现依赖包错误,请双击运行插件目录下的```install_requirements.bat```(官方便携包),或 ```install_requirements_aki.bat```(秋叶整合包) 重新安装依赖包。
+
+* 从requirements.txt 中废弃 [ObjectDetector YOLOWorld](#ObjectDetectorYOLOWorld) 节点所需的依赖。如需使用此节点,请手动安装依赖包。
+* 从ComfyUI Layer Style 剥离部分节点至本仓库。
+
+
+## 节点说明
+
+### QWenImage2Prompt
+根据图片反推提示词。这个节点是[ComfyUI_VLM_nodes](https://github.com/gokayfem/ComfyUI_VLM_nodes)中的```UForm-Gen2 Qwen Node```节点的重新封装,感谢原作者。
+请从[huggingface](https://huggingface.co/unum-cloud/uform-gen2-qwen-500m)或者[百度网盘](https://pan.baidu.com/s/1oRkUoOKWaxGod_XTJ8NiTA?pwd=d5d2)下载模型到```ComfyUI/models/LLavacheckpoints/files_for_uform_gen2_qwen```文件夹。
+
+
+
+节点选项说明:
+* question: 对UForm-Gen-QWen模型的提示词。
+
+### LlamaVision
+使用Llama 3.2 vision 模型进行本地推理。可以用于生成提示词。本节点部分代码来自[ComfyUI-PixtralLlamaMolmoVision](https://github.com/SeanScripts/ComfyUI-PixtralLlamaMolmoVision),感谢原作者。
+运行这个节点需要transformers升级到4.45.0以上。
+请从 [百度网盘](https://pan.baidu.com/s/18oHnTrkNMiwKLMcUVrfFjA?pwd=4g81) 或 [huggingface/SeanScripts](https://huggingface.co/SeanScripts/Llama-3.2-11B-Vision-Instruct-nf4/tree/main)下载整个文件夹,并复制到ComfyUI/models/LLM。
+
+
+
+节点选项说明:
+
+
+* image: 图片输入。
+* model: 目前仅有"Llama-3.2-11B-Vision-Instruct-nf4"这一个模型可用。
+* system_prompt: LLM模型的系统提示词。
+* user_prompt: LLM模型的用户提示词。
+* max_new_tokens: LLM的max_new_tokens参数。
+* do_sample: LLM的do_sample参数。
+* top-p: LLM的top_p参数。
+* top_k: LLM的top_k参数。
+* stop_strings: 截止字符串。
+* seed: 随机种子。
+* control_after_generate: 种子变化选项。
+* include_prompt_in_output: 输出是否包含提示词。
+* cache_model: 是否缓存模型。
+
+### JoyCaption2
+使用JoyCaption-alpha-two模型生成提示词。本节点是 https://huggingface.co/John6666/joy-caption-alpha-two-cli-mod 在ComfyUI中的实现,感谢原作者。
+请从 [百度网盘](https://pan.baidu.com/s/1dOjbUEacUOhzFitAQ3uIeQ?pwd=4ypv) 以及 [百度网盘](https://pan.baidu.com/s/1mH1SuW45Dy6Wga7aws5siQ?pwd=w6h5) ,
+或者 [huggingface/Orenguteng](https://huggingface.co/Orenguteng/Llama-3.1-8B-Lexi-Uncensored-V2/tree/main) 以及 [huggingface/unsloth](https://huggingface.co/unsloth/Meta-Llama-3.1-8B-Instruct/tree/main) 下载整个文件夹,并复制到ComfyUI/models/LLM,
+从 [百度网盘](https://pan.baidu.com/s/1pkVymOsDcXqL7IdQJ6lMVw?pwd=v8wp) 或者 [huggingface/google](https://huggingface.co/google/siglip-so400m-patch14-384/tree/main) 下载整个文件夹,并复制到ComfyUI/models/clip,
+从 [百度网盘](https://pan.baidu.com/s/12TDwZAeI68hWT6MgRrrK7Q?pwd=d7dh) 或者 [huggingface/John6666](https://huggingface.co/John6666/joy-caption-alpha-two-cli-mod/tree/main)下载 ```cgrkzexw-599808``` 文件夹,并复制到ComfyUI/models/Joy_caption。
+
+
+节点选项说明:
+
+
+* image: 图片输入。
+* extra_options: extra_options参数输入。
+* llm_model: 目前有 Orenguteng/Llama-3.1-8B-Lexi-Uncensored-V2 和 unsloth/Meta-Llama-3.1-8B-Instruct 两种LLM模型可选择。
+* device: 模型加载设备。目前仅支持cuda。
+* dtype: 模型加载精度,有nf4 和 bf16 两个选项。
+* vlm_lora: 是否加载text_model。
+* caption_type: caption类型选项, 包括"Descriptive"(正式语气描述), "Descriptive (Informal)"(非正式语气描述), "Training Prompt"(SD训练描述), "MidJourney"(MJ风格描述), "Booru tag list"(标签列表), "Booru-like tag list"(类标签列表), "Art Critic"(艺术评论), "Product Listing"(产品列表), "Social Media Post"(社交媒体风格)。
+* caption_length: 描述长度。
+* user_prompt: LLM模型的用户提示词。如果这里有内容将覆盖caption_type和extra_options的所有设置。
+* max_new_tokens: LLM的max_new_tokens参数。
+* do_sample: LLM的do_sample参数。
+* top-p: LLM的top_p参数。
+* temperature: LLM的temperature参数。
+* cache_model: 是否缓存模型。
+
+### JoyCaption2Split
+JoyCaption2 的分离式节点,将模型加载与推理分离,使用多个JoyCaption2节点时可共用模型提高效率。
+
+节点选项说明:
+
+
+* image: 图片输入。
+* joy2_model: JoyCaption模型输入。
+* extra_options: extra_options参数输入。
+* caption_type: caption类型选项, 包括"Descriptive"(正式语气描述), "Descriptive (Informal)"(非正式语气描述), "Training Prompt"(SD训练描述), "MidJourney"(MJ风格描述), "Booru tag list"(标签列表), "Booru-like tag list"(类标签列表), "Art Critic"(艺术评论), "Product Listing"(产品列表), "Social Media Post"(社交媒体风格)。
+* caption_length: 描述长度。
+* user_prompt: LLM模型的用户提示词。如果这里有内容将覆盖caption_type和extra_options的所有设置。
+* max_new_tokens: LLM的max_new_tokens参数。
+* do_sample: LLM的do_sample参数。
+* top-p: LLM的top_p参数。
+* temperature: LLM的temperature参数。
+
+### LoadJoyCaption2Model
+JoyCaption2 的模型加载节点,与JoyCaption2Split配合使用。
+
+节点选项说明:
+
+
+* llm_model: 目前有 Orenguteng/Llama-3.1-8B-Lexi-Uncensored-V2 和 unsloth/Meta-Llama-3.1-8B-Instruct 两种LLM模型可选择。
+* device: 模型加载设备。目前仅支持cuda。
+* dtype: 模型加载精度,有nf4 和 bf16 两个选项。
+* vlm_lora: 是否加载text_model。
+
+
+### JoyCaption2ExtraOptions
+JoyCaption2的extra_options参数节点。
+
+节点选项说明:
+
+
+* refer_character_name: 如果图像中有人物/角色,必须将其称为{name}
+* exclude_people_info: 不要包含有关无法更改的人物/角色的信息(例如种族、性别等),但仍包含可更改的属性(例如发型)。
+* include_lighting: 包括照明信息。
+* include_camera_angle: 包括摄影机角度信息。
+* include_watermark: 包括是否有水印信息。
+* include_JPEG_artifacts: 包括是否存在 JPEG 伪影信息。
+* include_exif: 如果是照片,包含相机的信息以及光圈、快门速度、ISO等信息。
+* exclude_sexual: 不要包含任何与性有关的内容,保持PG。
+* exclude_image_resolution: 不要包含图像分辨率信息。
+* include_aesthetic_quality: 包含图像美学(从低到非常高)信息。
+* include_composition_style: 包括有关图像构图风格的信息,例如引导线、三分法或对称性。
+* exclude_text: 不要包含任何文字信息。
+* specify_depth_field: 包含景深以及背景模糊信息。
+* specify_lighting_sources: 如果可以判别人造或自然光源,则包含在内。
+* do_not_use_ambiguous_language: 不要使用任何含糊不清的言辞。
+* include_nsfw: 包含NSFW或性暗示信息。
+* only_describe_most_important_elements: 只描述最重要的元素。
+* character_name: 如果选择了```refer_character_name```,则使用此处的名字。
+
+### PhiPrompt
+使用Micrisoft Phi 3.5文字及视觉模型进行本地推理。可以用于生成提示词,加工提示词或者反推图片的提示词。运行这个模型需要至少16GB的显存。
+请从[百度网盘](https://pan.baidu.com/s/1BdTLdaeGC3trh1U3V-6XTA?pwd=29dh) 或者 [huggingface.co/microsoft/Phi-3.5-vision-instruct](https://huggingface.co/microsoft/Phi-3.5-vision-instruct/tree/main) 和 [huggingface.co/microsoft/Phi-3.5-mini-instruct](https://huggingface.co/microsoft/Phi-3.5-mini-instruct/tree/main) 下载全部模型文件并放到 ```ComfyUI\models\LLM``` 文件夹。
+
+
+节点选项说明:
+
+
+* image: 可选输入。输入的图片将作为Phi-3.5-vision-instruct的输入。
+* model: 可选择加载的Phi-3.5-vision-instruct模型,或者Phi-3.5-mini-instruct模型。默认值auto将根据是否有图片输入自动加载对应模型。
+* device: 模型加载设备。支持cpu和cuda。
+* dtype: 模型加载精度,有fp16、bf16和fp32三个选项。
+* cache_model: 是否缓存模型。
+* system_prompt: Phi-3.5-mini-instruct的系统提示词。
+* user_prompt: LLM模型的用户提示词。
+* do_sample: LLM的do_sample参数,默认为True。
+* temperature: LLM的temperature参数,默认为0.5。
+* max_new_tokens: LLM的max_new_tokens参数,默认为512。
+
+
+### UserPromptGeneratorTxtImg
+用于生成SD文本到图片提示词的UserPrompt预设。
+
+节点选项说明:
+
+
+* template: 提示词模板。目前仅有“SD txt2img prompt”可用。
+* describe: 提示词描述。在这里输入简单的描述。
+* limit_word: 输出的提示词最大长度限制。例如200即表示输出文本将被限制在200个词以内。
+
+### UserPromptGeneratorTxtImgWithReference
+用于参考输入的内容生成SD文本到图片提示词的UserPrompt预设。
+
+节点选项说明:
+
+
+* reference_text: 参考文本输入。通常是图片的风格描述。
+* template: 提示词模板。目前仅有“SD txt2img prompt”可用。
+* describe: 提示词描述。在这里输入简单的描述。
+* limit_word: 输出的提示词最大长度限制。例如200即表示输出文本将被限制在200个词以内。
+
+
+### UserPromptGeneratorReplaceWord
+用于将文本中的某个关键词替换为不同内容的UserPrompt预设。这不仅是简单的替换,还可以根据提示词上下文进行文字逻辑梳理以达到输出内容的合理性。
+
+节点选项说明:
+
+
+* orig_prompt: 原始提示词输入。
+* template: 提示词模板。目前仅有“prompt replace word”可用。
+* exclude_word: 需要排除的关键词。
+* replace_with_word: 替换exclude_word的关键词。
+
+### PromptTagger
+根据图片反推提示词,可以设置替换词。这个节点目前使用Google Gemini API作为后端服务,请确保网络环境可以正常使用Gemini。
+请在[Google AI Studio](https://makersuite.google.com/app/apikey)申请你的API key, 并将其填到```api_key.ini```, 这个文件位于插件根目录下, 默认名字是```api_key.ini.example```, 初次使用这个文件需将文件后缀改为.ini。用文本编辑软件打开,在```google_api_key=```后面填入你的API key并保存。
+
+
+节点选项说明:
+
+
+* api: 使用的Api。有"gemini-1.5-flash"和"google-gemini"两个选项。
+* token_limit: 生成提示词的最大token限制。
+* exclude_word: 需要排除的关键词。
+* replace_with_word: 替换exclude_word的关键词。
+
+### PromptEmbellish
+输入简单的提示词,输出经过润色的提示词,支持输入图片作为参考,支持中文输入。这个节点目前使用Google Gemini API作为后端服务,请确保网络环境可以正常使用Gemini。
+请在[Google AI Studio](https://makersuite.google.com/app/apikey)申请你的API key, 并将其填到```api_key.ini```, 这个文件位于插件根目录下, 默认名字是```api_key.ini.example```, 初次使用这个文件需将文件后缀改为.ini。用文本编辑软件打开,在```google_api_key=```后面填入你的API key并保存。
+
+
+节点选项说明:
+
+
+* image: 可选项,输入图像作为提示词参考。
+* api: 使用的Api。有"gemini-1.5-flash"和"google-gemini"两个选项。
+* token_limit: 生成提示词的最大token限制。
+* discribe: 在这里输入简单的描述。支持中文。
+
+### Florence2Image2Prompt
+使用florence2模型反推提示词。本节点部分的代码来自[yiwangsimple/florence_dw](https://github.com/yiwangsimple/florence_dw),感谢原作者。
+*首次使用时将自动下载模型,请在可以访问huggingface.co的网络环境下使用。您也可以从[百度网盘](https://pan.baidu.com/s/1hzw9-QiU1vB8pMbBgofZIA?pwd=mfl3)下载模型文件并复制到```ComfyUI/models/florence2```文件夹。
+
+
+
+节点选项说明:
+
+* florence2_model: Florence2模型输入。
+* image: 图片输入。
+* task: 选择florence2任务。
+* text_input: florence2任务文本输入。
+* max_new_tokens: 生成文本的最大token数量。
+* num_beams: 生成文本的beam search数量。
+* do_sample: 是否使用文本生成采样。
+* fill_mask: 是否使用文本标记掩码填充。
+
+### GetColorTone
+从图片中获取主颜色或平均色。
+
+
+节点选项说明:
+
+* mode: 模式,有两种可选择,主颜色main_color和平均色average。
+
+输出:
+* RGB color in HEX: 使用16进制RGB字符串格式描述,例如 '#FA3D86'。
+* HSV color in list: HSV颜色值,使用list格式描述。
+
+### GetColorToneV2
+GetColorTone的V2升级版。可以指定获取主体或背景的主色或平均色。
+
+
+
+在GetColorTong基础上做了如下改变:
+
+* color_of: 提供4个选项,mask, entire, background和subject, 分别表示选择遮罩区域,整个图片,背景,或主体的颜色。
+* remove_background_method: 背景识别的方法, 有BiRefNet和RMBG V1.4两种可以选择。
+* invert_mask: 是否反转遮罩。
+* mask_grow: 遮罩扩张。对于subject, 更大的值使获得的颜色更接近主体中心的颜色。
+
+输出:
+* image: 纯色图片输出, 尺寸与输入的图片相同。
+* mask: 遮罩输出。
+
+
+### ImageRewardFilter
+
+对批量图片评分并输出排名靠前的图片。这个节点使用了[ImageReward](https://github.com/THUDM/ImageReward)作为图片评分,感谢原作者。
+
+
+节点选项说明:
+* prompt: 可选输入。将prompt在此输入将作为依据判定其与图片的符合程度。
+* output_nun: 输出的图片数量。此数值应小于图片批量。
+
+输出:
+* images: 按评分顺序从高到低输出的批量图片。
+* obsolete_images: 淘汰的图片。同样按评分顺序从高到低输出。
+
+
+### LaMa
+
+根据图像遮罩擦除物体。本节点是对[IOPaint](https://www.iopaint.com)的封装,由 SOTA AI 模型提供支持, 感谢原作者。
+提供[LaMa](https://github.com/advimman/lama), [LDM](https://github.com/CompVis/latent-diffusion), [ZITS](https://github.com/DQiaole/ZITS_inpainting),[MAT](https://github.com/fenglinglwb/MAT), [FcF](https://github.com/SHI-Labs/FcF-Inpainting), [Manga](https://github.com/msxie92/MangaInpainting) 模型以及 SPREAD 擦除方法。请查看链接了解各个模型的介绍。
+请下载模型文件 [lama models(百度网盘)](https://pan.baidu.com/s/1m7La2ELsSKaIFhQ57qg1XQ?pwd=jn10) 或者 [lama models(Google Drive)](https://drive.google.com/drive/folders/1Aq0a4sybb3SRxi7j1e1_ZbBRjaWDdP9e?usp=sharing), 将文件放到```ComfyUI/models/lama```
+
+节点选项说明:
+
+* lama_model: 选择模型或方法。
+* device: 在正确安装torch和Nvidia CUDA驱动程序后,使用cuda将明显提高运行速度。
+* invert_mask: 是否反转遮罩。
+* grow: 遮罩扩张幅度。正值是向外扩张,负值是向内收缩。
+* blur: 遮罩模糊幅度。
+
+
+### ImageAutoCrop
+
+自动抠图并按照遮罩裁切图片。可指定生成图片的背景颜色、长宽比和大小。这个节点是为生成训练模型的图片素材而设计的。
+*请参照 [SegmentAnythingUltra](#SegmentAnythingUltra) 和 [RemBgUltra](#RemBgUltra) 节点的模型安装方法安装模型。
+
+
+节点选项说明:
+
+* background_color4: 背景颜色。
+* aspect_ratio: 输出的宽高比。这里提供了常见的画幅比例, "custom"为自定义比例。
+* proportional_width: 比例宽。如果aspect_ratio选项不是"custom",此处设置将被忽略。
+* proportional_height: 比例高。如果aspect_ratio选项不是"custom",此处设置将被忽略。
+* scale_by_longest_side: 允许按长边尺寸缩放。
+* longest_side: scale_by_longest_side被设置为True时,此项将作为是图像长边的长度。
+* detect: 探测方法,min_bounding_rect是最小外接矩形, max_inscribed_rect是最大内接矩形。
+* border_reserve: 保留边框。在探测到的遮罩主体区域之外扩展裁切范围。
+* ultra_detail_range: 遮罩边缘超精细处理范围,0为不处理,可以节省生成时间。
+* matting_method: 生成遮罩的方法。有Segment Anything和 RMBG 1.4两种方法。RMBG 1.4运行速度更快。
+* sam_model: 此处选择Segment Anything所使用的sam模型。
+* grounding_dino_model: 此处选择Segment Anything所使用的grounding_dino模型。
+* sam_threshold: Segment Anything的阈值。
+* sam_prompt: Segment Anything的提示词。
+
+输出:
+cropped_image: 裁切并更换背景后的图像。
+box_preview: 裁切位置预览。
+cropped_mask: 裁切后的遮罩。
+
+### ImageAutoCropV2
+
+```ImageAutoCrop```的V2升级版,在之前基础上做了如下改变:
+
+
+* 增加```mask```可选输入。当有mask输入时,直接使用该输入跳过内置遮罩生成。
+* 增加```fill_background```, 当此项设置为False时将不处理背景,并且超出画幅的部分不纳入输出范围。
+* ```aspect_ratio```增加```original```(原始画面宽高比)选项。
+* scale_by: 允许按长边、短边、宽度或高度指定尺寸缩放。
+* scale_by_length: 这里的数值作为scale_by指定边的长度。
+
+### ImageAutoCropV3
+自动裁切图片到指定的尺寸。可输入mask以保留遮罩指定的区域。这个节点是为生成训练模型的图片素材而设计的。
+
+
+节点选项说明:
+
+* image: 输入的图像。
+* mask: 可选输入遮罩。遮罩部分将在裁切长宽比例范围内得到保留。
+* aspect_ratio: 输出的宽高比。这里提供了常见的画幅比例, "custom"为自定义比例, "original"为原始画面比例。
+* proportional_width: 比例宽。如果aspect_ratio选项不是"custom",此处设置将被忽略。
+* proportional_height: 比例高。如果aspect_ratio选项不是"custom",此处设置将被忽略。
+* method: 缩放的采样方法,包括lanczos、bicubic、hamming、bilinear、box和nearest。
+* scale_to_side: 允许按长边、短边、宽度、高度或总像素指定尺寸缩放。
+* scale_to_length: 这里的数值作为scale_to_side指定边的长度, 或者总像素数量(kilo pixels)。
+* round_to_multiple: 倍数取整。例如设置为8,宽和高将强制设置为8的倍数。
+
+输出:
+cropped_image: 裁切后的图像。
+box_preview: 裁切位置预览。
+
+
+### SaveImagePlus
+
+增强版的保存图片节点。可自定义保存图片的目录,文件名增加时间戳,选择保存格式,设置图片压缩率,设置是否保存工作流,以及可选给图片添加隐形水印(以肉眼无法觉察的方式添加信息,使用配套的```ShowBlindWaterMark```节点可以解码水印)。可选择是否同时输出工作流的json文件。
+
+节点选项说明:
+
+* iamge: 输入的图片。
+* custom_path*: 用户自定义目录,请按正确的格式输入目录名。如果为空则保存在ComfyUI默认的output目录。
+* filename_prefix*:文件名前缀。。
+* timestamp: 为文件名加上时间戳,可选择日期、时间到秒和时间到毫秒。
+* format:图片保存格式。目前提供png和jpg两种。注意RGBA模式的图片仅支持png格式。
+* quality:图片质量,数值范围10-100,数值越高,图片质量越好,文件的体积也对应增大。
+* meta_data:是否保存元数据即工作流信息到png文件。如果不希望泄露工作流,请把这里设置为false。
+* blind_watermark:这里输入的文字(不支持多语言)将被转换为二维码作为隐形水印保存,使用```ShowBlindWaterMark```节点可以解码水印。注意有水印的图片建议保存为png格式,质量较低的jpg格式将导致水印信息丢失。
+* save_workflow_as_json: 是否同时输出工作流为json文件(输出的json与图片在同一目录)。
+* preview: 预览开关。
+
+*输入```%date```表示当前日期(YY-mm-dd),```%time```表示当前时间(HH-MM-SS)。可以输入```/```表示子目录。例如```%date/name_%time``` 将输出图片到```YY-mm-dd```文件夹下,以```name_HH-MM-SS```为文件名前缀。
+
+
+### AddBlindWaterMark
+
+给图片添加隐形水印。以肉眼无法觉察的方式添加水印图片,使用```ShowBlindWaterMark```节点可以解码水印。
+
+节点选项说明:
+
+* iamge: 输入的图片。
+* watermark_image: 水印图片。这里输入的图片将自动转为正方形的黑白图片作为水印。建议使用二维码作为水印。
+
+
+### ShowBlindWaterMark
+对```AddBlindWaterMark``` 和 ```SaveImagePlus``` 节点添加的隐形水印解码。
+
+
+
+### CreateQRCode
+生成一个正方形的二维码图片。
+
+节点选项说明:
+
+* size: 生成图片的边长。
+* border: 二维码四周边框的大小,数值越大,边框越宽。
+* text: 这里输入二维码文字内容,不支持多语言。
+
+### DecodeQRCode
+解码二维码。
+
+节点选项说明:
+
+* image: 输入二维码图片。
+* pre_blur: 预模糊,对难以识别的二维码可以尝试调整此数值。
+
+### LoadPSD
+
+
+加载PSD格式文件,并导出图层。
+注意这个节点需要安装psd_tools依赖包,如果安装psd_tool中出现```ModuleNotFoundError: No module named 'docopt'```错误,请下载[docopt的whl](https://www.piwheels.org/project/docopt/)手动安装。
+
+节点选项说明:
+
+* image: 这里列出了```ComfyUI/input```下的*.psd文件,之前加载过的psd图片可以从这里选择。
+* file_path: psd文件的完整路径以及文件名。
+* include_hidden_layer: 是否包括隐藏图层。
+* find_layer_by: 查找图层的方法,可选择按图层索引编号或者图层名称查找。图层组被作为一个图层对待。
+* layer_index: 图层索引编号,0是最下面的图层,依次递增。如果include_hidden_layer设置为false,隐藏的图层不计入。设为-1则输出最上层的图层。
+* layer_name: 图层名称。注意大小写和标点符号必须完全匹配。
+
+输出:
+flat_image: psd预览图。
+layer_iamge: 查找的图层输出。
+all_layers: 包含全部图层的批量图片。
+
+### SD3NegativeConditioning
+
+把SD3的Negative Conditioning 的4个节点封装为一个单独节点。
+
+节点选项说明:
+
+* zero_out_start: 设置Negative ConditioningZeroOut的ConditioningSetTimestepRange start值, 此数值与Negative的ConditioningSetTimestepRange end值相同。
+
+
+
+### SegmentAnythingUltra
+对[ComfyUI Segment Anything](https://github.com/storyicon/comfyui_segment_anything)的改进,使遮罩有更具细节的边缘,感谢原作者。
+*请参照ComfyUI Segment Anything的安装方法安装模型。如果已经正确安装了ComfyUI Segment Anything,可跳过此步骤。
+* 从 [这里](https://huggingface.co/bert-base-uncased/tree/main) 下载 config.json,model.safetensors,tokenizer_config.json,tokenizer.json 和 vocab.txt 5个文件到 ```ComfyUI/models/bert-base-uncased```文件夹。
+* 下载 [GroundingDINO_SwinT_OGC config file](https://huggingface.co/ShilongLiu/GroundingDINO/resolve/main/GroundingDINO_SwinT_OGC.cfg.py), [GroundingDINO_SwinT_OGC model](https://huggingface.co/ShilongLiu/GroundingDINO/resolve/main/groundingdino_swint_ogc.pth),
+[GroundingDINO_SwinB config file](https://huggingface.co/ShilongLiu/GroundingDINO/resolve/main/GroundingDINO_SwinB.cfg.py), [GroundingDINO_SwinB model](https://huggingface.co/ShilongLiu/GroundingDINO/resolve/main/groundingdino_swinb_cogcoor.pth) 到 ```ComfyUI/models/grounding-dino```文件夹。
+* 下载 [sam_vit_h](https://dl.fbaipublicfiles.com/segment_anything/sam_vit_h_4b8939.pth),[sam_vit_l](https://dl.fbaipublicfiles.com/segment_anything/sam_vit_l_0b3195.pth),
+[sam_vit_b](https://dl.fbaipublicfiles.com/segment_anything/sam_vit_b_01ec64.pth), [sam_hq_vit_h](https://huggingface.co/lkeab/hq-sam/resolve/main/sam_hq_vit_h.pth),
+[sam_hq_vit_l](https://huggingface.co/lkeab/hq-sam/resolve/main/sam_hq_vit_l.pth), [sam_hq_vit_b](https://huggingface.co/lkeab/hq-sam/resolve/main/sam_hq_vit_b.pth),
+[mobile_sam](https://github.com/ChaoningZhang/MobileSAM/blob/master/weights/mobile_sam.pt) 这几个文件到```ComfyUI/models/sams```文件夹。
+*或者从[GroundingDino模型百度网盘](https://pan.baidu.com/s/1P7WQDuaqSYazlSQX8SJjxw?pwd=24ki) 和 [SAM模型百度网盘](https://pan.baidu.com/s/1n7JrHb2vzV2K2z3ktqpNxg?pwd=yoqh) 下载它们。
+
+
+
+
+节点选项说明:
+
+* sam_model: 选择SAM模型。
+* ground_dino_model: 选择Grounding DINO模型。
+* threshold: SAM阈值。
+* detail_range: 边缘细节范围。
+* black_point: 边缘黑色采样阈值。
+* white_point: 边缘白色采样阈值。
+* process_detail: 此处设为False将跳过边缘处理以节省运行时间。
+* prompt: SAM的prompt输入。
+* cache_model: 是否缓存模型。
+
+### SegmentAnythingUltraV2
+SegmentAnythingUltra的V2升级版,增加了VITMatte边缘处理方法。
+
+
+在SegmentAnythingUltra的基础上做了如下改变:
+
+* detail_method: 边缘处理方法。提供了VITMatte, VITMatte(local), PyMatting, GuidedFilter。
+* detail_erode: 遮罩边缘向内侵蚀范围。数值越大,向内修复的范围越大。
+* detail_dilate: 遮罩边缘向外扩张范围。数值越大,向外修复的范围越大。
+* device: 设置是否使用cuda。
+* max_megapixels: 设置vitmatte运算的最大尺寸。
+
+### SAM2Ultra
+本节点是[kijai/ComfyUI-segment-anything-2](https://github.com/kijai/ComfyUI-segment-anything-2)的改造版本。感谢[kijai](https://github.com/kijai)为ComfyUI社区做出的巨大贡献。
+SAM2 Ultra 节点仅支持单张图片,如果需要处理多张图片,请先将image batch 转换为 image list。
+*请从[百度网盘](https://pan.baidu.com/s/1xaQYBA6ktxvAxm310HXweQ?pwd=auki) 或者 [huggingface.co/Kijai/sam2-safetensors](https://huggingface.co/Kijai/sam2-safetensors/tree/main)下载全部模型文件并复制到```ComfyUI/models/sam2```文件夹。
+
+
+
+节点选项说明:
+
+
+* image: 图片输入。
+* bboxes: 识别框数据输入。
+* sam2_model: 选择SAM2模型。
+* presicion: 模型精度,可选择fp16, bf16 和 fp32。
+* bbox_select: 选择输入的框数据。有3个选项:"all"为全部选择,"first"为选择置信度最高的框,"by_index"可以指定框的索引。
+* select_index: 当bbox_select为"by_index"时,此选项有效。0为第一张。可以输入多个值,中间用任意非数字字符分隔,包括不仅限于逗号,句号,分号,空格或者字母,甚至中文。
+* cache_model: 是否缓存模型。缓存模型后将节省模型加载的时间。
+* detail_method: 边缘处理方法。提供了VITMatte, VITMatte(local), PyMatting, GuidedFilter。如果首次使用VITMatte后模型已经下载,之后可以使用VITMatte(local)。
+* detail_erode: 遮罩边缘向内侵蚀范围。数值越大,向内修复的范围越大。
+* detail_dilate: 遮罩边缘向外扩张范围。数值越大,向外修复的范围越大。
+* black_point: 边缘黑色采样阈值。
+* white_point: 边缘白色采样阈值。
+* process_detail: 此处设为False将跳过边缘处理以节省运行时间。
+* device: 设置是否使用cuda。
+* max_megapixels: 设置vitmatte运算的最大尺寸。
+
+### SAM2VideoUltra
+SAM2 Video Ultra 节点支持处理多张图片或视频序列帧。请在序列的第一帧定义识别框数据以保证正确识别。
+
+https://github.com/user-attachments/assets/4726b8bf-9b98-4630-8f54-cb7ed7a3d2c5
+
+https://github.com/user-attachments/assets/b2a45c96-4be1-4470-8ceb-addaf301b0cb
+
+节点选项说明:
+
+
+* image: 图片输入。
+* bboxes: 可选输入,识别框数据输入。bboxes 和 first_frame_mask 二者必须输入其中之一。如果有first_frame_mask输入,bboxes将被忽略。
+* first_frame_mask: 可选输入遮罩,这里的遮罩将作为首帧识别对象。bboxes 和 first_frame_mask 二者必须输入其中之一。如果有first_frame_mask输入,bboxes将被忽略。
+* pre_mask: 可选输入遮罩,这里的遮罩将作为传播关注范围限制,有助于提高识别准确度。
+* sam2_model: 选择SAM2模型。
+* presicion: 模型精度,可选择fp16, bf16。
+* cache_model: 是否缓存模型。缓存模型后将节省模型加载的时间。
+* individual_object: 当设置为 True时,将专注于识别单一对象。设置为False时,将尝试为多个对象生成识别框。
+* mask_preview_color: 在预览输出中显示非遮罩区域的颜色。
+* detail_method: 边缘处理方法。仅VITMatte可用。
+* detail_erode: 遮罩边缘向内侵蚀范围。数值越大,向内修复的范围越大。
+* detail_dilate: 遮罩边缘向外扩张范围。数值越大,向外修复的范围越大。
+* black_point: 边缘黑色采样阈值。
+* white_point: 边缘白色采样阈值。
+* process_detail: 此处设为False将跳过边缘处理以节省运行时间。
+* device: 本节点限制仅使用cuda。
+* max_megapixels: 设置vitmatte运算的最大尺寸。更大的尺寸将获得更精细的遮罩边缘,但会导致运算速度明显下降。
+
+### ObjectDetectorFL2
+使用Florence2模型识别图片中的对象,并输出识别框数据。
+*请从 [百度网盘](https://pan.baidu.com/s/1hzw9-QiU1vB8pMbBgofZIA?pwd=mfl3)下载模型文件并复制到```ComfyUI/models/florence2```文件夹。
+
+节点选项说明:
+
+
+* image: 图片输入。
+* florence2_model: Florence2模型。从[Florence2模型加载器](#LoadFlorence2Model)输入。
+* prompt: 描述需要识别的对象。
+* sort_method: 选择框排序方法, 有4个选项:"left_to_right"为从左到右排序,"top_to_bottom"为从上到下排序,"big_to_small"为从大到小排序,"confidence"为按置信度排序。
+* bbox_select: 选择输入的框数据。有3个选项:"all"为全部选择,"first"为选择置信度最高的框,"by_index"可以指定框的索引。
+* select_index: 当bbox_select为"by_index"时,此选项有效。0为第一张。可以输入多个值,中间用任意非数字字符分隔,包括不仅限于逗号,句号,分号,空格或者字母,甚至中文。
+
+### ObjectDetectorYOLOWorld (已废弃,如继续使用需要手动安装依赖包)
+
+由于依赖包安装易出问题,已废弃此节点。如需使用,请手动安装下列依赖包:
+```
+pip install inference-cli>=0.13.0
+pip install inference-gpu[yolo-world]>=0.13.0
+```
+
+使用YOLO World模型识别图片中的对象,并输出识别框数据。
+*请从 [百度网盘](https://pan.baidu.com/s/1QpjajeTA37vEAU2OQnbDcQ?pwd=nqsk) 或[GoogleDrive](https://drive.google.com/drive/folders/1nrsfq4S-yk9ewJgwrhXAoNVqIFLZ1at7?usp=sharing)下载模型文件并复制到```ComfyUI/models/yolo-world```文件夹。
+
+
+节点选项说明:
+
+
+* image: 图片输入。
+* confidence_threshold: 置信度阈值。
+* nms_iou_threshold: 非极大值抑制阈值。
+* prompt: 描述需要识别的对象。
+* sort_method: 选择框排序方法, 有4个选项:"left_to_right"为从左到右排序,"top_to_bottom"为从上到下排序,"big_to_small"为从大到小排序,"confidence"为按置信度排序。
+* bbox_select: 选择输入的框数据。有3个选项:"all"为全部选择,"first"为选择置信度最高的框,"by_index"可以指定框的索引。
+* select_index: 当bbox_select为"by_index"时,此选项有效。0为第一张。可以输入多个值,中间用任意非数字字符分隔,包括不仅限于逗号,句号,分号,空格或者字母,甚至中文。
+
+### ObjectDetectorYOLO8
+使用YOLO 8模型识别图片中的对象,并输出识别框数据。
+*请在 [GoogleDrive](https://drive.google.com/drive/folders/1I5TISO2G1ArSkKJu1O9b4Uvj3DVgn5d2) 或者 [百度网盘](https://pan.baidu.com/s/1pEY6sjABQaPs6QtpK0q6XA?pwd=grqe) 下载模型文件并放到 ```ComfyUI/models/yolo``` 文件夹。
+
+节点选项说明:
+
+* image: 图片输入。
+* yolo_model: 选择yolo模型。
+* sort_method: 选择框排序方法, 有4个选项:"left_to_right"为从左到右排序,"top_to_bottom"为从上到下排序,"big_to_small"为从大到小排序,"confidence"为按置信度排序。
+* bbox_select: 选择输入的框数据。有3个选项:"all"为全部选择,"first"为选择置信度最高的框,"by_index"可以指定框的索引。
+* select_index: 当bbox_select为"by_index"时,此选项有效。0为第一张。可以输入多个值,中间用任意非数字字符分隔,包括不仅限于逗号,句号,分号,空格或者字母,甚至中文。
+
+### ObjectDetectorMask
+使用遮罩作为识别框数据。遮罩上所有被白色区域包围的区域,将被识别为一个对象。多个封闭区域将各自识别。
+
+节点选项说明:
+
+* object_mask: 遮罩输入。
+* sort_method: 选择框排序方法, 有4个选项:"left_to_right"为从左到右排序,"top_to_bottom"为从上到下排序,"big_to_small"为从大到小排序,"confidence"为默认排序。
+* bbox_select: 选择输入的框数据。有3个选项:"all"为全部选择,"first"为选择置信度最高的框,"by_index"可以指定框的索引。
+* select_index: 当bbox_select为"by_index"时,此选项有效。0为第一张。可以输入多个值,中间用任意非数字字符分隔,包括不仅限于逗号,句号,分号,空格或者字母,甚至中文。
+
+### BBoxJoin
+合并识别框数据。
+
+节点选项说明:
+
+* bboxes_1: 必选输入。第一组识别框。
+* bboxes_2: 可选输入。第二组识别框。
+* bboxes_3: 可选输入。第三组识别框。
+* bboxes_4: 可选输入。第四组识别框。
+
+### DrawBBoxMask
+将ObjectDetector节点输出的识别框数据绘制为遮罩。
+
+
+节点选项说明:
+
+* image: 图片输入。必须与ObjectDetector节点识别的图片一致。
+* bboxes: 识别框数据输入。
+* grow_top: 每个识别框向上扩展范围,为识别框高度的百分比。正值为向上扩展,负值为向下扩展。
+* grow_bottom: 每个识别框向下扩展范围,为识别框高度的百分比,正值为向下扩展,负值为向上扩展。
+* grow_left: 每个识别框向左扩展范围,为识别框宽度的百分比。正值为向左扩展,负值为向右扩展。
+* grow_right: 每个识别框向右扩展范围,为识别框宽度的百分比。正值为向右扩展,负值为向左扩展。
+
+
+### EVF-SAMUltra
+本节点是[EVF-SAM](https://github.com/hustvl/EVF-SAM)在ComfyUI中的实现。
+*请从[百度网盘](https://pan.baidu.com/s/1EvaxgKcCxUpMbYKzLnEx9w?pwd=69bn) 或者 [huggingface/EVF-SAM2](https://huggingface.co/YxZhang/evf-sam2/tree/main), [huggingface/EVF-SAM](https://huggingface.co/YxZhang/evf-sam/tree/main) 下载全部模型文件并复制到```ComfyUI/models/EVF-SAM```文件夹(请将模型保存在各自子目录中)。
+
+
+
+节点选项说明:
+
+
+* image: 图片输入。
+* model: 选择模型。目前有 evf-sam2 和 evf-sam 可选。
+* presicion: 模型精度,可选择fp16, bf16 和 fp32。
+* load_in_bit: 按位精度加载模型。可选择full, 8 和 4。
+* pormpt: 用于分割的提示词。
+* detail_method: 边缘处理方法。提供了VITMatte, VITMatte(local), PyMatting, GuidedFilter。如果首次使用VITMatte后模型已经下载,之后可以使用VITMatte(local)。
+* detail_erode: 遮罩边缘向内侵蚀范围。数值越大,向内修复的范围越大。
+* detail_dilate: 遮罩边缘向外扩张范围。数值越大,向外修复的范围越大。
+* black_point: 边缘黑色采样阈值。
+* white_point: 边缘白色采样阈值。
+* process_detail: 此处设为False将跳过边缘处理以节省运行时间。
+* device: 设置是否使用cuda。
+* max_megapixels: 设置vitmatte运算的最大尺寸。
+
+
+### Florence2Ultra
+使用 Florence2 模型的分割功能,同时具有超高的边缘细节。
+本节点部分的代码来自[spacepxl/ComfyUI-Florence-2](https://github.com/spacepxl/ComfyUI-Florence-2),感谢原作者。
+*请从 [百度网盘](https://pan.baidu.com/s/1hzw9-QiU1vB8pMbBgofZIA?pwd=mfl3)下载模型文件并复制到```ComfyUI/models/florence2```文件夹。
+
+
+
+节点选项说明:
+
+* florence2_model: Florence2模型输入。
+* image: 图片输入。
+* task: 选择florence2任务。
+* text_input: florence2任务文本输入。
+* detail_method: 边缘处理方法。提供了VITMatte, VITMatte(local), PyMatting, GuidedFilter。如果首次使用VITMatte后模型已经下载,之后可以使用VITMatte(local)。
+* detail_erode: 遮罩边缘向内侵蚀范围。数值越大,向内修复的范围越大。
+* detail_dilate: 遮罩边缘向外扩张范围。数值越大,向外修复的范围越大。
+* black_point: 边缘黑色采样阈值。
+* white_point: 边缘白色采样阈值。
+* process_detail: 此处设为False将跳过边缘处理以节省运行时间。
+* device: 设置是否使用cuda。
+* max_megapixels: 设置vitmatte运算的最大尺寸。
+
+### LoadFlorence2Model
+Florence2 模型加载器。
+
+目前有 base, base-ft, large, large-ft, DocVQA, SD3-Captioner 和 base-PromptGen模型可以选择。
+
+
+### BiRefNetUltra
+使用BiRefNet模型去除背景,有更好的识别能力,同时具有超高的边缘细节。
+本节点模型部分的代码来自vipery的[ComfyUI-BiRefNet](https://github.com/viperyl/ComfyUI-BiRefNet),感谢原作者。
+
+*从[https://huggingface.co/ViperYX/BiRefNet](https://huggingface.co/ViperYX/BiRefNet/tree/main) 或者 [百度网盘](https://pan.baidu.com/s/1GxtuNDTIHkuu4FR4uGAT-g?pwd=t2cf) 下载```BiRefNet-ep480.pth```,```pvt_v2_b2.pth```,```pvt_v2_b5.pth```,```swin_base_patch4_window12_384_22kto1k.pth```, ```swin_large_patch4_window12_384_22kto1k.pth```5个文件至```ComfyUI/models/BiRefNet```文件夹。
+
+
+
+节点选项说明:
+
+* detail_method: 边缘处理方法。提供了VITMatte, VITMatte(local), PyMatting, GuidedFilter。如果首次使用VITMatte后模型已经下载,之后可以使用VITMatte(local)。
+* detail_erode: 遮罩边缘向内侵蚀范围。数值越大,向内修复的范围越大。
+* detail_dilate: 遮罩边缘向外扩张范围。数值越大,向外修复的范围越大。
+* black_point: 边缘黑色采样阈值。
+* white_point: 边缘白色采样阈值。
+* process_detail: 此处设为False将跳过边缘处理以节省运行时间。
+* device: 设置是否使用cuda。
+* max_megapixels: 设置vitmatte运算的最大尺寸。
+
+
+### BiRefNetUltraV2
+本节点支持使用最新的BiRefNet模型。
+*从[百度网盘](https://pan.baidu.com/s/12z3qUuqag3nqpN2NJ5pSzg?pwd=ek65) 或 [GoogleDrive](https://drive.google.com/drive/folders/1s2Xe0cjq-2ctnJBR24563yMSCOu4CcxM) 下载 ```BiRefNet-general-epoch_244.pth``` 到 ```ComfyUI/Models/BiRefNet/pth``` 文件夹。也可以下载更多的BiRefNet模型放到这里。
+
+
+
+节点选项说明:
+
+
+* image: 图片输入。
+* birefnet_model: BiRefNet模型输入,模型从LoadBiRefNetModel节点输出。
+* detail_method: 边缘处理方法。提供了VITMatte, VITMatte(local), PyMatting, GuidedFilter。如果首次使用VITMatte后模型已经下载,之后可以使用VITMatte(local)。
+* detail_erode: 遮罩边缘向内侵蚀范围。数值越大,向内修复的范围越大。
+* detail_dilate: 遮罩边缘向外扩张范围。数值越大,向外修复的范围越大。
+* black_point: 边缘黑色采样阈值。
+* white_point: 边缘白色采样阈值。
+* process_detail: 由于BiRefNet的边缘处理已经非常不错,此处默认设为False。
+* device: 设置是否使用cuda。
+* max_megapixels: 设置vitmatte运算的最大尺寸。
+
+
+### LoadBiRefNetModel
+加载BiRefNet模型。
+
+
+节点选项说明:
+
+
+* model: 选择模型。列出 ```CoomfyUI/models/BiRefNet/pth``` 文件夹下的文件供选择。
+
+### LoadBiRefNetModelV2
+本节点是[jimlee2048](https://github.com/jimlee2048)提交的PR,支持加载RMBG-2.0模型。
+从 [huggingface](https://huggingface.co/briaai/RMBG-2.0/tree/main) 或 [百度网盘](https://pan.baidu.com/s/1viIXlZnpTYTKkm2F-QMj_w?pwd=axr9) 下载全部文件并复制到```ComfyUI/models/BiRefNet/RMBG-2.0```文件夹。
+
+节点选项说明:
+
+
+* model: 选择模型。有两个选项: ```BiRefNet-General``` 和 ```RMBG-2.0```。
+
+
+
+### TransparentBackgroundUltra
+使用transparent-background模型去除背景,有更好的识别能力和识别速度,同时具有超高的边缘细节。
+
+*从 [googledrive](https://drive.google.com/drive/folders/10KBDY19egb8qEQBv34cqIVSwd38bUAa9?usp=sharing) 或 [百度网盘](https://pan.baidu.com/s/10JO0uKzTxJaIkhN_J7RSyw?pwd=v0b0) 下载全部文件至```ComfyUI/models/transparent-background```文件夹。
+
+
+
+节点选项说明:
+
+* model: 选择模型。
+* detail_method: 边缘处理方法。提供了VITMatte, VITMatte(local), PyMatting, GuidedFilter。如果首次使用VITMatte后模型已经下载,之后可以使用VITMatte(local)。
+* detail_erode: 遮罩边缘向内侵蚀范围。数值越大,向内修复的范围越大。
+* detail_dilate: 遮罩边缘向外扩张范围。数值越大,向外修复的范围越大。
+* black_point: 边缘黑色采样阈值。
+* white_point: 边缘白色采样阈值。
+* process_detail: 此处设为False将跳过边缘处理以节省运行时间。
+* device: 设置是否使用cuda。
+* max_megapixels: 设置vitmatte运算的最大尺寸。
+
+### PersonMaskUltra
+为人物生成脸、头发、身体皮肤、衣服或配饰的遮罩。与之前的A Person Mask Generator节点相比,这个节点具有超高的边缘细节。
+本节点的模型代码来自[a-person-mask-generator](https://github.com/djbielejeski/a-person-mask-generator),边缘处理代码来自spacepxl的[ComfyUI-Image-Filters](https://github.com/spacepxl/ComfyUI-Image-Filters),感谢原作者。
+*从[百度网盘](https://pan.baidu.com/s/13zqZtBt89ueCyFufzUlcDg?pwd=jh5g) 下载模型文件并放到```ComfyUI/models/mediapipe```文件夹。
+
+
+
+节点选项说明:
+
+* face: 脸部识别。
+* hair: 头发识别。
+* body: 身体皮肤识别。
+* clothes: 衣服识别。
+* accessories: 配饰(例如背包)识别。
+* background: 背景识别。
+* confidence: 识别阈值,更低的值将输出更多的遮罩范围。
+* detail_range: 边缘细节范围。
+* black_point: 边缘黑色采样阈值。
+* white_point: 边缘白色采样阈值。
+* process_detail: 此处设为False将跳过边缘处理以节省运行时间。
+
+### PersonMaskUltraV2
+PersonMaskUltra的V2升级版,增加了VITMatte边缘处理方法。
+
+在PersonMaskUltra的基础上做了如下改变:
+
+* detail_method: 边缘处理方法。提供了VITMatte, VITMatte(local), PyMatting, GuidedFilter。如果首次使用VITMatte后模型已经下载,之后可以使用VITMatte(local)。
+* detail_erode: 遮罩边缘向内侵蚀范围。数值越大,向内修复的范围越大。
+* detail_dilate: 遮罩边缘向外扩张范围。数值越大,向外修复的范围越大。
+* device: 设置是否使用cuda。
+* max_megapixels: 设置vitmatte运算的最大尺寸。
+
+
+### HumanPartsUltra
+用于分割人体肢体,是基于[metal3d/ComfyUI_Human_Parts](https://github.com/metal3d/ComfyUI_Human_Parts) 的重新封装,感谢原作者。
+本节点在原作基础上增加了超精细边缘处理。请从[百度网盘](https://pan.baidu.com/s/1-6uwH6RB0FhIVfa3qO7hhQ?pwd=d862) 或 [huggingface](https://huggingface.co/Metal3d/deeplabv3p-resnet50-human/tree/main) 下载模型文件并复制到 ```ComfyUI\models\onnx\human-parts``` 文件夹。
+
+
+节点选项说明:
+
+
+* image: 图片输入。
+* face: 是否识别人脸。
+* hair: 是否识别头发。
+* galsses: 是否识别眼镜。
+* top_clothes: 是否识别上装。
+* bottom_clothes: 是否识别下装。
+* torso_skin: 是否识别躯干皮肤。
+* left_arm: 是否识别左手臂。
+* right_arm: 是否识别右手臂。
+* left_leg: 是否识别左腿。
+* right_leg: 是否识别右腿。
+* left_foot: 是否识别左脚。
+* right_foot: 是否识别右脚。
+* detail_method: 边缘处理方法。提供了VITMatte, VITMatte(local), PyMatting, GuidedFilter。如果首次使用VITMatte后模型已经下载,之后可以使用VITMatte(local)。
+* detail_erode: 遮罩边缘向内侵蚀范围。数值越大,向内修复的范围越大。
+* detail_dilate: 遮罩边缘向外扩张范围。数值越大,向外修复的范围越大。
+* black_point: 边缘黑色采样阈值。
+* white_point: 边缘白色采样阈值。
+* process_detail: 此处设为False将跳过边缘处理以节省运行时间。
+* device: 设置是否使用cuda。
+* max_megapixels: 设置vitmatte运算的最大尺寸。
+
+
+### YoloV8Detect
+使用YoloV8模型检测人脸、手部box区域,或者人物分割。支持输出所选择数量的通道。
+请在 [GoogleDrive](https://drive.google.com/drive/folders/1I5TISO2G1ArSkKJu1O9b4Uvj3DVgn5d2) 或者 [百度网盘](https://pan.baidu.com/s/1pEY6sjABQaPs6QtpK0q6XA?pwd=grqe) 下载模型文件并放到 ```ComfyUI/models/yolo``` 文件夹。
+
+
+
+节点选项说明:
+
+* yolo_model: yolo模型选择。带有```seg```名字的模型可以输出分割的mask, 否则只能输出box区域的遮罩。
+* mask_merge: 选择合并的遮罩。```all```是合并全部遮罩输出。选数值是输出多少个遮罩,按识别置信度排序合并输出。
+
+输出:
+* mask: 输出的遮罩。
+* yolo_plot_image: yolo识别结果预览图。
+* yolo_masks: yolo识别出来的所有遮罩,每个单独的遮罩输出为一个mask。
+
+
+### MediapipeFacialSegment
+使用Mediapipe模型检测人脸五官,分割左右眉、眼睛、嘴唇和牙齿。
+*从[百度网盘](https://pan.baidu.com/s/13zqZtBt89ueCyFufzUlcDg?pwd=jh5g) 下载模型文件并放到```ComfyUI/models/mediapipe```文件夹。
+
+
+
+节点选项说明:
+
+* left_eye: 左眼识别开关。
+* left_eyebrow: 左眉识别开关。
+* right_eye: 右眼识别开关。
+* right_eyebrow: 右眉识别开关。
+* lips: 嘴唇识别开关。
+* tooth: 牙齿识别开关。
+
+
+### MaskByDifferent
+计算两张图像不同之处,并输出为遮罩。
+
+
+节点选项说明:
+
+* gain: 计算增益。调高此值,微弱的差异将更显著的呈现。
+* fix_gap: 修补遮罩内部缝隙。更高的值将修补更大的缝隙。
+* fix_threshold: 修补阈值。
+* main_subject_detect: 此项设为True将开启主体侦测,忽略主体之外的差异。
+
+
+
+## 节点注解
+1 image、mask和background_image(如果有输入)这三项必须是相同的尺寸。
+
+2 mask不是必须的输入项,默认使用image的alpha通道,如果image输入不包含alpha通道将自动创建整个图像的alpha通道。如果输入mask,原本的alpha通道将被mask覆盖。
+
+3 混合模式 包括normal、multply、screen、add、subtract、difference、darker、lighter、color_burn、color_dodge、linear_burn、linear_dodge、overlay、soft_light、hard_light、vivid_light、pin_light、linear_light、hard_mix, 共19种混合模式。
+
+*混合模式预览
+
+
+3 混合模式V2 包括nomal, dissolve, darken, multiply, color burn, linear burn, darker color, lighten, screen, color dodge, linear dodge(add), lighter color, dodge, overlay, soft light, hard light, vivid light, linear light, pin light, hard mix, difference, exclusion, subtract, divide, hue, saturation, color, luminosity, grain extract, grain merge共30种模式。
+混合模式V2的部分代码来自[Virtuoso Nodes for ComfyUI](https://github.com/chrisfreilich/virtuoso-nodes)的```Blend Modes```节点。感谢原作者。
+
+*混合模式V2版预览
+
+4 颜色使用16进制RGB字符串格式描述,例如 '#FA3D86'。
+
+5 image和mask这两项必须是相同的尺寸。
+
+## Star 记录
+
+[](https://star-history.com/#chflame163/ComfyUI_LayerStyle_Advance&Date)
+
+## 声明
+LayerStyle Advance节点遵照MIT开源协议,有部分功能代码和模型来自其他开源项目,感谢原作者。如果作为商业用途,请查阅原项目授权协议使用。
diff --git a/__init__.py b/__init__.py
new file mode 100644
index 0000000..cec1e3d
--- /dev/null
+++ b/__init__.py
@@ -0,0 +1,53 @@
+import importlib.util
+import os
+import sys
+import json
+
+NODE_CLASS_MAPPINGS = {}
+NODE_DISPLAY_NAME_MAPPINGS = {}
+
+python = sys.executable
+
+def get_ext_dir(subpath=None, mkdir=False):
+ dir = os.path.dirname(__file__)
+ if subpath is not None:
+ dir = os.path.join(dir, subpath)
+
+ dir = os.path.abspath(dir)
+
+ if mkdir and not os.path.exists(dir):
+ os.makedirs(dir)
+ return dir
+
+def serialize(obj):
+ if isinstance(obj, (str, int, float, bool, list, dict, type(None))):
+ return obj
+ return str(obj) # 转为字符串
+
+
+py = get_ext_dir("py")
+files = os.listdir(py)
+all_nodes = {}
+for file in files:
+ if not file.endswith(".py"):
+ continue
+ name = os.path.splitext(file)[0]
+ imported_module = importlib.import_module(".py.{}".format(name), __name__)
+ try:
+ NODE_CLASS_MAPPINGS = {**NODE_CLASS_MAPPINGS, **imported_module.NODE_CLASS_MAPPINGS}
+ NODE_DISPLAY_NAME_MAPPINGS = {**NODE_DISPLAY_NAME_MAPPINGS, **imported_module.NODE_DISPLAY_NAME_MAPPINGS}
+ serialized_CLASS_MAPPINGS = {k: serialize(v) for k, v in imported_module.NODE_CLASS_MAPPINGS.items()}
+ serialized_DISPLAY_NAME_MAPPINGS = {k: serialize(v) for k, v in imported_module.NODE_DISPLAY_NAME_MAPPINGS.items()}
+ all_nodes[file]={"NODE_CLASS_MAPPINGS": serialized_CLASS_MAPPINGS, "NODE_DISPLAY_NAME_MAPPINGS": serialized_DISPLAY_NAME_MAPPINGS}
+ except:
+ pass
+
+
+# 保存为文件
+with open("all_nodes.json", "w", encoding="utf-8") as f:
+ json.dump(all_nodes, f, ensure_ascii=False, indent=4)
+
+
+WEB_DIRECTORY = "./js"
+
+__all__ = ["NODE_CLASS_MAPPINGS", "NODE_DISPLAY_NAME_MAPPINGS", "WEB_DIRECTORY"]
diff --git a/api_key.ini.example b/api_key.ini.example
new file mode 100644
index 0000000..3ed1f43
--- /dev/null
+++ b/api_key.ini.example
@@ -0,0 +1,2 @@
+# LayerStyle api_key
+google_api_key=
\ No newline at end of file
diff --git a/custom_size.ini.example b/custom_size.ini.example
new file mode 100644
index 0000000..fc8de0c
--- /dev/null
+++ b/custom_size.ini.example
@@ -0,0 +1,10 @@
+# LayerStyle Custom_size
+1024 x 1024
+768 x 512
+512 x 768
+1280 x 720
+720 x 1280
+1344 x 768
+768 x 1344
+1536 x 640
+640 x 1536
diff --git a/font/Alibaba-PuHuiTi-Heavy.ttf b/font/Alibaba-PuHuiTi-Heavy.ttf
new file mode 100644
index 0000000..7eb047f
Binary files /dev/null and b/font/Alibaba-PuHuiTi-Heavy.ttf differ
diff --git a/image/add_blind_watermark_node.jpg b/image/add_blind_watermark_node.jpg
new file mode 100644
index 0000000..f2c95fb
Binary files /dev/null and b/image/add_blind_watermark_node.jpg differ
diff --git a/image/bbox_join_node.jpg b/image/bbox_join_node.jpg
new file mode 100644
index 0000000..c80cbb1
Binary files /dev/null and b/image/bbox_join_node.jpg differ
diff --git a/image/birefnet_ultra_example.jpg b/image/birefnet_ultra_example.jpg
new file mode 100644
index 0000000..7ec3f83
Binary files /dev/null and b/image/birefnet_ultra_example.jpg differ
diff --git a/image/birefnet_ultra_node.jpg b/image/birefnet_ultra_node.jpg
new file mode 100644
index 0000000..b5f69f9
Binary files /dev/null and b/image/birefnet_ultra_node.jpg differ
diff --git a/image/birefnet_ultra_v2_example.jpg b/image/birefnet_ultra_v2_example.jpg
new file mode 100644
index 0000000..c424fe7
Binary files /dev/null and b/image/birefnet_ultra_v2_example.jpg differ
diff --git a/image/birefnet_ultra_v2_node.jpg b/image/birefnet_ultra_v2_node.jpg
new file mode 100644
index 0000000..4c95748
Binary files /dev/null and b/image/birefnet_ultra_v2_node.jpg differ
diff --git a/image/blend_mode_result.jpg b/image/blend_mode_result.jpg
new file mode 100644
index 0000000..732d7b8
Binary files /dev/null and b/image/blend_mode_result.jpg differ
diff --git a/image/blend_mode_v2_example.jpg b/image/blend_mode_v2_example.jpg
new file mode 100644
index 0000000..b2fbaad
Binary files /dev/null and b/image/blend_mode_v2_example.jpg differ
diff --git a/image/create_qrcode_node.jpg b/image/create_qrcode_node.jpg
new file mode 100644
index 0000000..87d1a17
Binary files /dev/null and b/image/create_qrcode_node.jpg differ
diff --git a/image/data_nodes_example.jpg b/image/data_nodes_example.jpg
new file mode 100644
index 0000000..6995c09
Binary files /dev/null and b/image/data_nodes_example.jpg differ
diff --git a/image/decode_qrcode_node.jpg b/image/decode_qrcode_node.jpg
new file mode 100644
index 0000000..569814d
Binary files /dev/null and b/image/decode_qrcode_node.jpg differ
diff --git a/image/draw_bbox_mask_example.jpg b/image/draw_bbox_mask_example.jpg
new file mode 100644
index 0000000..a478112
Binary files /dev/null and b/image/draw_bbox_mask_example.jpg differ
diff --git a/image/draw_bbox_mask_node.jpg b/image/draw_bbox_mask_node.jpg
new file mode 100644
index 0000000..cf008a3
Binary files /dev/null and b/image/draw_bbox_mask_node.jpg differ
diff --git a/image/evf_sam_ultra_example.jpg b/image/evf_sam_ultra_example.jpg
new file mode 100644
index 0000000..ed747af
Binary files /dev/null and b/image/evf_sam_ultra_example.jpg differ
diff --git a/image/evf_sam_ultra_node.jpg b/image/evf_sam_ultra_node.jpg
new file mode 100644
index 0000000..e45626d
Binary files /dev/null and b/image/evf_sam_ultra_node.jpg differ
diff --git a/image/florence2_image2prompt_example.jpg b/image/florence2_image2prompt_example.jpg
new file mode 100644
index 0000000..eedca22
Binary files /dev/null and b/image/florence2_image2prompt_example.jpg differ
diff --git a/image/florence2_image2prompt_node.jpg b/image/florence2_image2prompt_node.jpg
new file mode 100644
index 0000000..99b5b50
Binary files /dev/null and b/image/florence2_image2prompt_node.jpg differ
diff --git a/image/florence2_ultra_example.jpg b/image/florence2_ultra_example.jpg
new file mode 100644
index 0000000..737c9cf
Binary files /dev/null and b/image/florence2_ultra_example.jpg differ
diff --git a/image/florence2_ultra_node.jpg b/image/florence2_ultra_node.jpg
new file mode 100644
index 0000000..06e9e3d
Binary files /dev/null and b/image/florence2_ultra_node.jpg differ
diff --git a/image/get_color_tone_example.jpg b/image/get_color_tone_example.jpg
new file mode 100644
index 0000000..a3dbaad
Binary files /dev/null and b/image/get_color_tone_example.jpg differ
diff --git a/image/get_color_tone_node.jpg b/image/get_color_tone_node.jpg
new file mode 100644
index 0000000..197dc95
Binary files /dev/null and b/image/get_color_tone_node.jpg differ
diff --git a/image/get_color_tone_v2_example.jpg b/image/get_color_tone_v2_example.jpg
new file mode 100644
index 0000000..15ffef1
Binary files /dev/null and b/image/get_color_tone_v2_example.jpg differ
diff --git a/image/get_color_tone_v2_example2.jpg b/image/get_color_tone_v2_example2.jpg
new file mode 100644
index 0000000..606ebc9
Binary files /dev/null and b/image/get_color_tone_v2_example2.jpg differ
diff --git a/image/get_color_tone_v2_node.jpg b/image/get_color_tone_v2_node.jpg
new file mode 100644
index 0000000..5f656a9
Binary files /dev/null and b/image/get_color_tone_v2_node.jpg differ
diff --git a/image/human_parts_node.jpg b/image/human_parts_node.jpg
new file mode 100644
index 0000000..efcc868
Binary files /dev/null and b/image/human_parts_node.jpg differ
diff --git a/image/human_parts_ultra_example.jpg b/image/human_parts_ultra_example.jpg
new file mode 100644
index 0000000..578122f
Binary files /dev/null and b/image/human_parts_ultra_example.jpg differ
diff --git a/image/image_auto_crop_example.jpg b/image/image_auto_crop_example.jpg
new file mode 100644
index 0000000..59377b9
Binary files /dev/null and b/image/image_auto_crop_example.jpg differ
diff --git a/image/image_auto_crop_node.jpg b/image/image_auto_crop_node.jpg
new file mode 100644
index 0000000..d3d2f83
Binary files /dev/null and b/image/image_auto_crop_node.jpg differ
diff --git a/image/image_auto_crop_v2_node.jpg b/image/image_auto_crop_v2_node.jpg
new file mode 100644
index 0000000..8420c78
Binary files /dev/null and b/image/image_auto_crop_v2_node.jpg differ
diff --git a/image/image_auto_crop_v3_node.jpg b/image/image_auto_crop_v3_node.jpg
new file mode 100644
index 0000000..4bd1683
Binary files /dev/null and b/image/image_auto_crop_v3_node.jpg differ
diff --git a/image/image_reward_filter_example.jpg b/image/image_reward_filter_example.jpg
new file mode 100644
index 0000000..1024b3e
Binary files /dev/null and b/image/image_reward_filter_example.jpg differ
diff --git a/image/image_reward_filter_node.jpg b/image/image_reward_filter_node.jpg
new file mode 100644
index 0000000..8de28f8
Binary files /dev/null and b/image/image_reward_filter_node.jpg differ
diff --git a/image/joycaption2_example.jpg b/image/joycaption2_example.jpg
new file mode 100644
index 0000000..433c5fb
Binary files /dev/null and b/image/joycaption2_example.jpg differ
diff --git a/image/joycaption2_extra_options_node.jpg b/image/joycaption2_extra_options_node.jpg
new file mode 100644
index 0000000..879cfe5
Binary files /dev/null and b/image/joycaption2_extra_options_node.jpg differ
diff --git a/image/joycaption2_node.jpg b/image/joycaption2_node.jpg
new file mode 100644
index 0000000..fc817b3
Binary files /dev/null and b/image/joycaption2_node.jpg differ
diff --git a/image/joycaption2_split_node.jpg b/image/joycaption2_split_node.jpg
new file mode 100644
index 0000000..4eccf6f
Binary files /dev/null and b/image/joycaption2_split_node.jpg differ
diff --git a/image/lama_example.jpg b/image/lama_example.jpg
new file mode 100644
index 0000000..68f9e5e
Binary files /dev/null and b/image/lama_example.jpg differ
diff --git a/image/lama_node.jpg b/image/lama_node.jpg
new file mode 100644
index 0000000..a870081
Binary files /dev/null and b/image/lama_node.jpg differ
diff --git a/image/light_leak_example.jpg b/image/light_leak_example.jpg
new file mode 100644
index 0000000..bff5535
Binary files /dev/null and b/image/light_leak_example.jpg differ
diff --git a/image/light_leak_node.jpg b/image/light_leak_node.jpg
new file mode 100644
index 0000000..597e1dd
Binary files /dev/null and b/image/light_leak_node.jpg differ
diff --git a/image/llama_vision_example.jpg b/image/llama_vision_example.jpg
new file mode 100644
index 0000000..141f563
Binary files /dev/null and b/image/llama_vision_example.jpg differ
diff --git a/image/llama_vision_node.jpg b/image/llama_vision_node.jpg
new file mode 100644
index 0000000..bb78416
Binary files /dev/null and b/image/llama_vision_node.jpg differ
diff --git a/image/load_ben_model_node.jpg b/image/load_ben_model_node.jpg
new file mode 100644
index 0000000..95d0dc7
Binary files /dev/null and b/image/load_ben_model_node.jpg differ
diff --git a/image/load_birefnet_model_node.jpg b/image/load_birefnet_model_node.jpg
new file mode 100644
index 0000000..d9500f2
Binary files /dev/null and b/image/load_birefnet_model_node.jpg differ
diff --git a/image/load_birefnet_model_v2_node.jpg b/image/load_birefnet_model_v2_node.jpg
new file mode 100644
index 0000000..1b63f0f
Binary files /dev/null and b/image/load_birefnet_model_v2_node.jpg differ
diff --git a/image/load_florence2_model_node.jpg b/image/load_florence2_model_node.jpg
new file mode 100644
index 0000000..a6b6ffb
Binary files /dev/null and b/image/load_florence2_model_node.jpg differ
diff --git a/image/load_image_example.jpg b/image/load_image_example.jpg
new file mode 100644
index 0000000..c132eb1
Binary files /dev/null and b/image/load_image_example.jpg differ
diff --git a/image/load_image_example_psd_file.jpg b/image/load_image_example_psd_file.jpg
new file mode 100644
index 0000000..3bf4464
Binary files /dev/null and b/image/load_image_example_psd_file.jpg differ
diff --git a/image/load_image_node.jpg b/image/load_image_node.jpg
new file mode 100644
index 0000000..b12cc2b
Binary files /dev/null and b/image/load_image_node.jpg differ
diff --git a/image/load_joycaption2_model_node.jpg b/image/load_joycaption2_model_node.jpg
new file mode 100644
index 0000000..7b6ae05
Binary files /dev/null and b/image/load_joycaption2_model_node.jpg differ
diff --git a/image/mask_by_different_example.jpg b/image/mask_by_different_example.jpg
new file mode 100644
index 0000000..f114ce4
Binary files /dev/null and b/image/mask_by_different_example.jpg differ
diff --git a/image/mask_by_different_node.jpg b/image/mask_by_different_node.jpg
new file mode 100644
index 0000000..489db5d
Binary files /dev/null and b/image/mask_by_different_node.jpg differ
diff --git a/image/mediapipe_facial_segment_example.jpg b/image/mediapipe_facial_segment_example.jpg
new file mode 100644
index 0000000..46561cb
Binary files /dev/null and b/image/mediapipe_facial_segment_example.jpg differ
diff --git a/image/mediapipe_facial_segment_node.jpg b/image/mediapipe_facial_segment_node.jpg
new file mode 100644
index 0000000..8f5b2d6
Binary files /dev/null and b/image/mediapipe_facial_segment_node.jpg differ
diff --git a/image/node-menu.jpg b/image/node-menu.jpg
new file mode 100644
index 0000000..8edcf7f
Binary files /dev/null and b/image/node-menu.jpg differ
diff --git a/image/node-search.jpg b/image/node-search.jpg
new file mode 100644
index 0000000..47ac217
Binary files /dev/null and b/image/node-search.jpg differ
diff --git a/image/object_detector_fl2_node.jpg b/image/object_detector_fl2_node.jpg
new file mode 100644
index 0000000..b8e28fe
Binary files /dev/null and b/image/object_detector_fl2_node.jpg differ
diff --git a/image/object_detector_mask_node.jpg b/image/object_detector_mask_node.jpg
new file mode 100644
index 0000000..bb03ce5
Binary files /dev/null and b/image/object_detector_mask_node.jpg differ
diff --git a/image/object_detector_yolo8_node.jpg b/image/object_detector_yolo8_node.jpg
new file mode 100644
index 0000000..5a9ce4f
Binary files /dev/null and b/image/object_detector_yolo8_node.jpg differ
diff --git a/image/object_detector_yolo_world_node.jpg b/image/object_detector_yolo_world_node.jpg
new file mode 100644
index 0000000..a2ef02a
Binary files /dev/null and b/image/object_detector_yolo_world_node.jpg differ
diff --git a/image/outer_glow_example.jpg b/image/outer_glow_example.jpg
new file mode 100644
index 0000000..2955587
Binary files /dev/null and b/image/outer_glow_example.jpg differ
diff --git a/image/person_mask_ultra_example.jpg b/image/person_mask_ultra_example.jpg
new file mode 100644
index 0000000..a5f7645
Binary files /dev/null and b/image/person_mask_ultra_example.jpg differ
diff --git a/image/person_mask_ultra_node.jpg b/image/person_mask_ultra_node.jpg
new file mode 100644
index 0000000..4f4773e
Binary files /dev/null and b/image/person_mask_ultra_node.jpg differ
diff --git a/image/person_mask_ultra_v2_node.jpg b/image/person_mask_ultra_v2_node.jpg
new file mode 100644
index 0000000..ec2ce54
Binary files /dev/null and b/image/person_mask_ultra_v2_node.jpg differ
diff --git a/image/phi_prompt_example.jpg b/image/phi_prompt_example.jpg
new file mode 100644
index 0000000..5d420ed
Binary files /dev/null and b/image/phi_prompt_example.jpg differ
diff --git a/image/phi_prompt_node.jpg b/image/phi_prompt_node.jpg
new file mode 100644
index 0000000..5325b9e
Binary files /dev/null and b/image/phi_prompt_node.jpg differ
diff --git a/image/prompt_embellish_example.jpg b/image/prompt_embellish_example.jpg
new file mode 100644
index 0000000..e65d18f
Binary files /dev/null and b/image/prompt_embellish_example.jpg differ
diff --git a/image/prompt_embellish_node.jpg b/image/prompt_embellish_node.jpg
new file mode 100644
index 0000000..b8ca5dc
Binary files /dev/null and b/image/prompt_embellish_node.jpg differ
diff --git a/image/prompt_tagger_example.jpg b/image/prompt_tagger_example.jpg
new file mode 100644
index 0000000..1dd9727
Binary files /dev/null and b/image/prompt_tagger_example.jpg differ
diff --git a/image/prompt_tagger_example1.jpg b/image/prompt_tagger_example1.jpg
new file mode 100644
index 0000000..07d1d9f
Binary files /dev/null and b/image/prompt_tagger_example1.jpg differ
diff --git a/image/prompt_tagger_node.jpg b/image/prompt_tagger_node.jpg
new file mode 100644
index 0000000..dde7564
Binary files /dev/null and b/image/prompt_tagger_node.jpg differ
diff --git a/image/qwen_image2prompt_example.jpg b/image/qwen_image2prompt_example.jpg
new file mode 100644
index 0000000..99f84bb
Binary files /dev/null and b/image/qwen_image2prompt_example.jpg differ
diff --git a/image/sam2_example.jpg b/image/sam2_example.jpg
new file mode 100644
index 0000000..9d47627
Binary files /dev/null and b/image/sam2_example.jpg differ
diff --git a/image/sam2_ultra_node.jpg b/image/sam2_ultra_node.jpg
new file mode 100644
index 0000000..0dfbaf8
Binary files /dev/null and b/image/sam2_ultra_node.jpg differ
diff --git a/image/sam2_video_ultra_node.jpg b/image/sam2_video_ultra_node.jpg
new file mode 100644
index 0000000..a6aaf5f
Binary files /dev/null and b/image/sam2_video_ultra_node.jpg differ
diff --git a/image/saveimage_plus_example.jpg b/image/saveimage_plus_example.jpg
new file mode 100644
index 0000000..2427d57
Binary files /dev/null and b/image/saveimage_plus_example.jpg differ
diff --git a/image/saveimage_plus_node.jpg b/image/saveimage_plus_node.jpg
new file mode 100644
index 0000000..2ae6d18
Binary files /dev/null and b/image/saveimage_plus_node.jpg differ
diff --git a/image/sd3_negative_conditioning_example.jpg b/image/sd3_negative_conditioning_example.jpg
new file mode 100644
index 0000000..299718b
Binary files /dev/null and b/image/sd3_negative_conditioning_example.jpg differ
diff --git a/image/sd3_negative_conditioning_node.jpg b/image/sd3_negative_conditioning_node.jpg
new file mode 100644
index 0000000..e19548e
Binary files /dev/null and b/image/sd3_negative_conditioning_node.jpg differ
diff --git a/image/sd3_negative_conditioning_node_note.jpg b/image/sd3_negative_conditioning_node_note.jpg
new file mode 100644
index 0000000..913d07a
Binary files /dev/null and b/image/sd3_negative_conditioning_node_note.jpg differ
diff --git a/image/segment_anything_ultra_compare.jpg b/image/segment_anything_ultra_compare.jpg
new file mode 100644
index 0000000..550df2c
Binary files /dev/null and b/image/segment_anything_ultra_compare.jpg differ
diff --git a/image/segment_anything_ultra_example.jpg b/image/segment_anything_ultra_example.jpg
new file mode 100644
index 0000000..e816dc2
Binary files /dev/null and b/image/segment_anything_ultra_example.jpg differ
diff --git a/image/segment_anything_ultra_node.jpg b/image/segment_anything_ultra_node.jpg
new file mode 100644
index 0000000..9a9e7fc
Binary files /dev/null and b/image/segment_anything_ultra_node.jpg differ
diff --git a/image/segment_anything_ultra_v2_node.jpg b/image/segment_anything_ultra_v2_node.jpg
new file mode 100644
index 0000000..ea31e86
Binary files /dev/null and b/image/segment_anything_ultra_v2_node.jpg differ
diff --git a/image/show_blind_watermark_node.jpg b/image/show_blind_watermark_node.jpg
new file mode 100644
index 0000000..166ffa5
Binary files /dev/null and b/image/show_blind_watermark_node.jpg differ
diff --git a/image/transparent_background_ultra_example.jpg b/image/transparent_background_ultra_example.jpg
new file mode 100644
index 0000000..b906e5a
Binary files /dev/null and b/image/transparent_background_ultra_example.jpg differ
diff --git a/image/transparent_background_ultra_node.jpg b/image/transparent_background_ultra_node.jpg
new file mode 100644
index 0000000..c979480
Binary files /dev/null and b/image/transparent_background_ultra_node.jpg differ
diff --git a/image/ultra_v2_nodes_example.jpg b/image/ultra_v2_nodes_example.jpg
new file mode 100644
index 0000000..d9b06c4
Binary files /dev/null and b/image/ultra_v2_nodes_example.jpg differ
diff --git a/image/userprompt_generator_replace_word_node.jpg b/image/userprompt_generator_replace_word_node.jpg
new file mode 100644
index 0000000..a592e23
Binary files /dev/null and b/image/userprompt_generator_replace_word_node.jpg differ
diff --git a/image/userprompt_generator_txt2img_node.jpg b/image/userprompt_generator_txt2img_node.jpg
new file mode 100644
index 0000000..6463c77
Binary files /dev/null and b/image/userprompt_generator_txt2img_node.jpg differ
diff --git a/image/userprompt_generator_txt2img_with_reference_node.jpg b/image/userprompt_generator_txt2img_with_reference_node.jpg
new file mode 100644
index 0000000..88c10c9
Binary files /dev/null and b/image/userprompt_generator_txt2img_with_reference_node.jpg differ
diff --git a/image/water_color_example.jpg b/image/water_color_example.jpg
new file mode 100644
index 0000000..3642972
Binary files /dev/null and b/image/water_color_example.jpg differ
diff --git a/image/water_color_node.jpg b/image/water_color_node.jpg
new file mode 100644
index 0000000..8a66baa
Binary files /dev/null and b/image/water_color_node.jpg differ
diff --git a/image/watermark_example.jpg b/image/watermark_example.jpg
new file mode 100644
index 0000000..e3b4c09
Binary files /dev/null and b/image/watermark_example.jpg differ
diff --git a/image/yolov8_detect_example.jpg b/image/yolov8_detect_example.jpg
new file mode 100644
index 0000000..1dccdec
Binary files /dev/null and b/image/yolov8_detect_example.jpg differ
diff --git a/image/yolov8_detect_node.jpg b/image/yolov8_detect_node.jpg
new file mode 100644
index 0000000..ea1a6b3
Binary files /dev/null and b/image/yolov8_detect_node.jpg differ
diff --git a/install_requirements.bat b/install_requirements.bat
new file mode 100644
index 0000000..d86aa03
--- /dev/null
+++ b/install_requirements.bat
@@ -0,0 +1,31 @@
+@echo off
+
+set "python_exec=..\..\..\python_embeded\python.exe"
+set "repair_dependency_txt=%~dp0\repair_dependency_list.txt"
+set "requirements_txt=%~dp0\requirements.txt"
+
+echo Installing with ComfyUI Portable
+echo .
+echo Install whl...
+%python_exec% -s -m pip install ./whl/docopt-0.6.2-py2.py3-none-any.whl
+%python_exec% -s -m pip install ./whl/hydra_core-1.3.2-py3-none-any.whl
+
+echo .
+echo Install requirement.txt...
+
+for /f "delims=" %%i in (%requirements_txt%) do (
+ %python_exec% -s -m pip install "%%i"
+ )
+
+echo .
+echo Fixing Dependency Package...
+%python_exec% -s -m pip uninstall -y onnxruntime
+%python_exec% -s -m pip uninstall -y opencv-python opencv-contrib-python opencv-python-headless opencv-contrib-python-headless
+for /f "delims=" %%i in (%repair_dependency_txt%) do (
+ %python_exec% -s -m pip install "%%i"
+ )
+
+echo .
+echo Install Finish!
+pause
+
diff --git a/install_requirements_aki.bat b/install_requirements_aki.bat
new file mode 100644
index 0000000..f7b9e33
--- /dev/null
+++ b/install_requirements_aki.bat
@@ -0,0 +1,30 @@
+@echo off
+
+set "python_exec=..\..\python\python.exe"
+set "repair_dependency_txt=%~dp0\repair_dependency_list.txt"
+set "requirements_txt=%~dp0\requirements.txt"
+
+echo Installing with ComfyUI Portable
+echo .
+echo Install whl...
+%python_exec% -s -m pip install ./whl/docopt-0.6.2-py2.py3-none-any.whl
+%python_exec% -s -m pip install ./whl/hydra_core-1.3.2-py3-none-any.whl
+
+echo .
+echo Install requirement.txt...
+for /f "delims=" %%i in (%requirements_txt%) do (
+ %python_exec% -s -m pip install "%%i"
+ )
+
+echo .
+echo Fixing Dependency Package...
+%python_exec% -s -m pip uninstall -y onnxruntime
+%python_exec% -s -m pip uninstall -y opencv-python opencv-contrib-python opencv-python-headless opencv-contrib-python-headless
+for /f "delims=" %%i in (%repair_dependency_txt%) do (
+ %python_exec% -s -m pip install "%%i"
+ )
+
+echo .
+echo Install Finish!
+pause
+
diff --git a/js/dz_node_palette.js b/js/dz_node_palette.js
new file mode 100644
index 0000000..6356500
--- /dev/null
+++ b/js/dz_node_palette.js
@@ -0,0 +1,42 @@
+import { app } from "../../scripts/app.js";
+
+
+app.registerExtension({
+ name: "ColorOverlay",
+ async nodeCreated(node) {
+ // 判断是否为layer节点
+ if(!node.comfyClass.startsWith("Layer")) {
+ return;
+ }
+
+ if(node.comfyClass.startsWith("LayerStyle:")) {
+ node.color = "rgba(20, 95, 121, 0.7)";
+// node.bgcolor = "rgba(50, 241, 255, 0.15)";
+ }
+
+ if(node.comfyClass.startsWith("LayerColor:")) {
+ node.color = "rgba(27, 89, 123, 0.7)";
+// node.bgcolor = "rgba(43, 209, 255, 0.15)";
+ }
+
+ if(node.comfyClass.startsWith("LayerMask:")) {
+ node.color = "rgba(27, 80, 119, 0.7)";
+// node.bgcolor = "rgba(4, 174, 255, 0.15)";
+ }
+
+ if(node.comfyClass.startsWith("LayerUtility:")) {
+ node.color = "rgba(38, 73, 116, 0.7)";
+// node.bgcolor = "rgba(23, 113, 255, 0.15)";
+ }
+
+ if(node.comfyClass.startsWith("LayerFilter:")) {
+ node.color = "rgba(34, 67, 111, 0.7)";
+// node.bgcolor = "rgba(19, 85, 255, 0.15)";
+ }
+
+
+// if(node.comfyClass === "LayerStyle: ColorOverlay"){
+// node.setSize([600, 120]);
+// }
+ }
+});
\ No newline at end of file
diff --git a/py/BiRefNet_legacy/__init__.py b/py/BiRefNet_legacy/__init__.py
new file mode 100644
index 0000000..e69de29
diff --git a/py/BiRefNet_legacy/backbones/__init__.py b/py/BiRefNet_legacy/backbones/__init__.py
new file mode 100644
index 0000000..e69de29
diff --git a/py/BiRefNet_legacy/backbones/build_backbone.py b/py/BiRefNet_legacy/backbones/build_backbone.py
new file mode 100644
index 0000000..e637c5e
--- /dev/null
+++ b/py/BiRefNet_legacy/backbones/build_backbone.py
@@ -0,0 +1,45 @@
+import torch
+import torch.nn as nn
+from collections import OrderedDict
+from torchvision.models import vgg16, vgg16_bn, VGG16_Weights, VGG16_BN_Weights, resnet50, ResNet50_Weights
+from BiRefNet_legacy.backbones.pvt_v2 import pvt_v2_b2, pvt_v2_b5
+from BiRefNet_legacy.backbones.swin_v1 import swin_v1_t, swin_v1_s, swin_v1_b, swin_v1_l
+from ..config import Config
+
+
+config = Config()
+
+def build_backbone(bb_name, pretrained=True, params_settings=''):
+ if bb_name == 'vgg16':
+ bb_net = list(vgg16(pretrained=VGG16_Weights.DEFAULT if pretrained else None).children())[0]
+ bb = nn.Sequential(OrderedDict({'conv1': bb_net[:4], 'conv2': bb_net[4:9], 'conv3': bb_net[9:16], 'conv4': bb_net[16:23]}))
+ elif bb_name == 'vgg16bn':
+ bb_net = list(vgg16_bn(pretrained=VGG16_BN_Weights.DEFAULT if pretrained else None).children())[0]
+ bb = nn.Sequential(OrderedDict({'conv1': bb_net[:6], 'conv2': bb_net[6:13], 'conv3': bb_net[13:23], 'conv4': bb_net[23:33]}))
+ elif bb_name == 'resnet50':
+ bb_net = list(resnet50(pretrained=ResNet50_Weights.DEFAULT if pretrained else None).children())
+ bb = nn.Sequential(OrderedDict({'conv1': nn.Sequential(*bb_net[0:3]), 'conv2': bb_net[4], 'conv3': bb_net[5], 'conv4': bb_net[6]}))
+ else:
+ bb = eval('{}({})'.format(bb_name, params_settings))
+ if pretrained:
+ bb = load_weights(bb, bb_name)
+ return bb
+
+def load_weights(model, model_name):
+ # save_model = torch.load(config.weights[model_name])
+ save_model = torch.load(config.weights[model_name], map_location=torch.device('cpu'))
+ model_dict = model.state_dict()
+ state_dict = {k: v if v.size() == model_dict[k].size() else model_dict[k] for k, v in save_model.items() if k in model_dict.keys()}
+ # to ignore the weights with mismatched size when I modify the backbone itself.
+ if not state_dict:
+ save_model_keys = list(save_model.keys())
+ sub_item = save_model_keys[0] if len(save_model_keys) == 1 else None
+ state_dict = {k: v if v.size() == model_dict[k].size() else model_dict[k] for k, v in save_model[sub_item].items() if k in model_dict.keys()}
+ if not state_dict or not sub_item:
+ print('Weights are not successully loaded. Check the state dict of weights file.')
+ return None
+ else:
+ print('Found correct weights in the "{}" item of loaded state_dict.'.format(sub_item))
+ model_dict.update(state_dict)
+ model.load_state_dict(model_dict)
+ return model
diff --git a/py/BiRefNet_legacy/backbones/pvt_v2.py b/py/BiRefNet_legacy/backbones/pvt_v2.py
new file mode 100644
index 0000000..7164f85
--- /dev/null
+++ b/py/BiRefNet_legacy/backbones/pvt_v2.py
@@ -0,0 +1,434 @@
+import torch
+import torch.nn as nn
+from functools import partial
+
+from timm.models.layers import DropPath, to_2tuple, trunc_normal_
+from timm.models import register_model
+
+import math
+
+from ..config import Config
+config = Config()
+
+class Mlp(nn.Module):
+ def __init__(self, in_features, hidden_features=None, out_features=None, act_layer=nn.GELU, drop=0.):
+ super().__init__()
+ out_features = out_features or in_features
+ hidden_features = hidden_features or in_features
+ self.fc1 = nn.Linear(in_features, hidden_features)
+ self.dwconv = DWConv(hidden_features)
+ self.act = act_layer()
+ self.fc2 = nn.Linear(hidden_features, out_features)
+ self.drop = nn.Dropout(drop)
+
+ self.apply(self._init_weights)
+
+ def _init_weights(self, m):
+ if isinstance(m, nn.Linear):
+ trunc_normal_(m.weight, std=.02)
+ if isinstance(m, nn.Linear) and m.bias is not None:
+ nn.init.constant_(m.bias, 0)
+ elif isinstance(m, nn.LayerNorm):
+ nn.init.constant_(m.bias, 0)
+ nn.init.constant_(m.weight, 1.0)
+ elif isinstance(m, nn.Conv2d):
+ fan_out = m.kernel_size[0] * m.kernel_size[1] * m.out_channels
+ fan_out //= m.groups
+ m.weight.data.normal_(0, math.sqrt(2.0 / fan_out))
+ if m.bias is not None:
+ m.bias.data.zero_()
+
+ def forward(self, x, H, W):
+ x = self.fc1(x)
+ x = self.dwconv(x, H, W)
+ x = self.act(x)
+ x = self.drop(x)
+ x = self.fc2(x)
+ x = self.drop(x)
+ return x
+
+
+class Attention(nn.Module):
+ def __init__(self, dim, num_heads=8, qkv_bias=False, qk_scale=None, attn_drop=0., proj_drop=0., sr_ratio=1):
+ super().__init__()
+ assert dim % num_heads == 0, f"dim {dim} should be divided by num_heads {num_heads}."
+
+ self.dim = dim
+ self.num_heads = num_heads
+ head_dim = dim // num_heads
+ self.scale = qk_scale or head_dim ** -0.5
+
+ self.q = nn.Linear(dim, dim, bias=qkv_bias)
+ self.kv = nn.Linear(dim, dim * 2, bias=qkv_bias)
+ self.attn_drop_prob = attn_drop
+ self.attn_drop = nn.Dropout(attn_drop)
+ self.proj = nn.Linear(dim, dim)
+ self.proj_drop = nn.Dropout(proj_drop)
+
+ self.sr_ratio = sr_ratio
+ if sr_ratio > 1:
+ self.sr = nn.Conv2d(dim, dim, kernel_size=sr_ratio, stride=sr_ratio)
+ self.norm = nn.LayerNorm(dim)
+
+ self.apply(self._init_weights)
+
+ def _init_weights(self, m):
+ if isinstance(m, nn.Linear):
+ trunc_normal_(m.weight, std=.02)
+ if isinstance(m, nn.Linear) and m.bias is not None:
+ nn.init.constant_(m.bias, 0)
+ elif isinstance(m, nn.LayerNorm):
+ nn.init.constant_(m.bias, 0)
+ nn.init.constant_(m.weight, 1.0)
+ elif isinstance(m, nn.Conv2d):
+ fan_out = m.kernel_size[0] * m.kernel_size[1] * m.out_channels
+ fan_out //= m.groups
+ m.weight.data.normal_(0, math.sqrt(2.0 / fan_out))
+ if m.bias is not None:
+ m.bias.data.zero_()
+
+ def forward(self, x, H, W):
+ B, N, C = x.shape
+ q = self.q(x).reshape(B, N, self.num_heads, C // self.num_heads).permute(0, 2, 1, 3)
+
+ if self.sr_ratio > 1:
+ x_ = x.permute(0, 2, 1).reshape(B, C, H, W)
+ x_ = self.sr(x_).reshape(B, C, -1).permute(0, 2, 1)
+ x_ = self.norm(x_)
+ kv = self.kv(x_).reshape(B, -1, 2, self.num_heads, C // self.num_heads).permute(2, 0, 3, 1, 4)
+ else:
+ kv = self.kv(x).reshape(B, -1, 2, self.num_heads, C // self.num_heads).permute(2, 0, 3, 1, 4)
+ k, v = kv[0], kv[1]
+
+ if config.SDPA_enabled:
+ x = torch.nn.functional.scaled_dot_product_attention(
+ q, k, v,
+ attn_mask=None, dropout_p=self.attn_drop_prob, is_causal=False
+ ).transpose(1, 2).reshape(B, N, C)
+ else:
+ attn = (q @ k.transpose(-2, -1)) * self.scale
+ attn = attn.softmax(dim=-1)
+ attn = self.attn_drop(attn)
+
+ x = (attn @ v).transpose(1, 2).reshape(B, N, C)
+ x = self.proj(x)
+ x = self.proj_drop(x)
+
+ return x
+
+
+class Block(nn.Module):
+
+ def __init__(self, dim, num_heads, mlp_ratio=4., qkv_bias=False, qk_scale=None, drop=0., attn_drop=0.,
+ drop_path=0., act_layer=nn.GELU, norm_layer=nn.LayerNorm, sr_ratio=1):
+ super().__init__()
+ self.norm1 = norm_layer(dim)
+ self.attn = Attention(
+ dim,
+ num_heads=num_heads, qkv_bias=qkv_bias, qk_scale=qk_scale,
+ attn_drop=attn_drop, proj_drop=drop, sr_ratio=sr_ratio)
+ # NOTE: drop path for stochastic depth, we shall see if this is better than dropout here
+ self.drop_path = DropPath(drop_path) if drop_path > 0. else nn.Identity()
+ self.norm2 = norm_layer(dim)
+ mlp_hidden_dim = int(dim * mlp_ratio)
+ self.mlp = Mlp(in_features=dim, hidden_features=mlp_hidden_dim, act_layer=act_layer, drop=drop)
+
+ self.apply(self._init_weights)
+
+ def _init_weights(self, m):
+ if isinstance(m, nn.Linear):
+ trunc_normal_(m.weight, std=.02)
+ if isinstance(m, nn.Linear) and m.bias is not None:
+ nn.init.constant_(m.bias, 0)
+ elif isinstance(m, nn.LayerNorm):
+ nn.init.constant_(m.bias, 0)
+ nn.init.constant_(m.weight, 1.0)
+ elif isinstance(m, nn.Conv2d):
+ fan_out = m.kernel_size[0] * m.kernel_size[1] * m.out_channels
+ fan_out //= m.groups
+ m.weight.data.normal_(0, math.sqrt(2.0 / fan_out))
+ if m.bias is not None:
+ m.bias.data.zero_()
+
+ def forward(self, x, H, W):
+ x = x + self.drop_path(self.attn(self.norm1(x), H, W))
+ x = x + self.drop_path(self.mlp(self.norm2(x), H, W))
+
+ return x
+
+
+class OverlapPatchEmbed(nn.Module):
+ """ Image to Patch Embedding
+ """
+
+ def __init__(self, img_size=224, patch_size=7, stride=4, in_channels=3, embed_dim=768):
+ super().__init__()
+ img_size = to_2tuple(img_size)
+ patch_size = to_2tuple(patch_size)
+
+ self.img_size = img_size
+ self.patch_size = patch_size
+ self.H, self.W = img_size[0] // patch_size[0], img_size[1] // patch_size[1]
+ self.num_patches = self.H * self.W
+ self.proj = nn.Conv2d(in_channels, embed_dim, kernel_size=patch_size, stride=stride,
+ padding=(patch_size[0] // 2, patch_size[1] // 2))
+ self.norm = nn.LayerNorm(embed_dim)
+
+ self.apply(self._init_weights)
+
+ def _init_weights(self, m):
+ if isinstance(m, nn.Linear):
+ trunc_normal_(m.weight, std=.02)
+ if isinstance(m, nn.Linear) and m.bias is not None:
+ nn.init.constant_(m.bias, 0)
+ elif isinstance(m, nn.LayerNorm):
+ nn.init.constant_(m.bias, 0)
+ nn.init.constant_(m.weight, 1.0)
+ elif isinstance(m, nn.Conv2d):
+ fan_out = m.kernel_size[0] * m.kernel_size[1] * m.out_channels
+ fan_out //= m.groups
+ m.weight.data.normal_(0, math.sqrt(2.0 / fan_out))
+ if m.bias is not None:
+ m.bias.data.zero_()
+
+ def forward(self, x):
+ x = self.proj(x)
+ _, _, H, W = x.shape
+ x = x.flatten(2).transpose(1, 2)
+ x = self.norm(x)
+
+ return x, H, W
+
+
+class PyramidVisionTransformerImpr(nn.Module):
+ def __init__(self, img_size=224, patch_size=16, in_channels=3, num_classes=1000, embed_dims=[64, 128, 256, 512],
+ num_heads=[1, 2, 4, 8], mlp_ratios=[4, 4, 4, 4], qkv_bias=False, qk_scale=None, drop_rate=0.,
+ attn_drop_rate=0., drop_path_rate=0., norm_layer=nn.LayerNorm,
+ depths=[3, 4, 6, 3], sr_ratios=[8, 4, 2, 1]):
+ super().__init__()
+ self.num_classes = num_classes
+ self.depths = depths
+
+ # patch_embed
+ self.patch_embed1 = OverlapPatchEmbed(img_size=img_size, patch_size=7, stride=4, in_channels=in_channels,
+ embed_dim=embed_dims[0])
+ self.patch_embed2 = OverlapPatchEmbed(img_size=img_size // 4, patch_size=3, stride=2, in_channels=embed_dims[0],
+ embed_dim=embed_dims[1])
+ self.patch_embed3 = OverlapPatchEmbed(img_size=img_size // 8, patch_size=3, stride=2, in_channels=embed_dims[1],
+ embed_dim=embed_dims[2])
+ self.patch_embed4 = OverlapPatchEmbed(img_size=img_size // 16, patch_size=3, stride=2, in_channels=embed_dims[2],
+ embed_dim=embed_dims[3])
+
+ # transformer encoder
+ dpr = [x.item() for x in torch.linspace(0, drop_path_rate, sum(depths))] # stochastic depth decay rule
+ cur = 0
+ self.block1 = nn.ModuleList([Block(
+ dim=embed_dims[0], num_heads=num_heads[0], mlp_ratio=mlp_ratios[0], qkv_bias=qkv_bias, qk_scale=qk_scale,
+ drop=drop_rate, attn_drop=attn_drop_rate, drop_path=dpr[cur + i], norm_layer=norm_layer,
+ sr_ratio=sr_ratios[0])
+ for i in range(depths[0])])
+ self.norm1 = norm_layer(embed_dims[0])
+
+ cur += depths[0]
+ self.block2 = nn.ModuleList([Block(
+ dim=embed_dims[1], num_heads=num_heads[1], mlp_ratio=mlp_ratios[1], qkv_bias=qkv_bias, qk_scale=qk_scale,
+ drop=drop_rate, attn_drop=attn_drop_rate, drop_path=dpr[cur + i], norm_layer=norm_layer,
+ sr_ratio=sr_ratios[1])
+ for i in range(depths[1])])
+ self.norm2 = norm_layer(embed_dims[1])
+
+ cur += depths[1]
+ self.block3 = nn.ModuleList([Block(
+ dim=embed_dims[2], num_heads=num_heads[2], mlp_ratio=mlp_ratios[2], qkv_bias=qkv_bias, qk_scale=qk_scale,
+ drop=drop_rate, attn_drop=attn_drop_rate, drop_path=dpr[cur + i], norm_layer=norm_layer,
+ sr_ratio=sr_ratios[2])
+ for i in range(depths[2])])
+ self.norm3 = norm_layer(embed_dims[2])
+
+ cur += depths[2]
+ self.block4 = nn.ModuleList([Block(
+ dim=embed_dims[3], num_heads=num_heads[3], mlp_ratio=mlp_ratios[3], qkv_bias=qkv_bias, qk_scale=qk_scale,
+ drop=drop_rate, attn_drop=attn_drop_rate, drop_path=dpr[cur + i], norm_layer=norm_layer,
+ sr_ratio=sr_ratios[3])
+ for i in range(depths[3])])
+ self.norm4 = norm_layer(embed_dims[3])
+
+ # classification head
+ # self.head = nn.Linear(embed_dims[3], num_classes) if num_classes > 0 else nn.Identity()
+
+ self.apply(self._init_weights)
+
+ def _init_weights(self, m):
+ if isinstance(m, nn.Linear):
+ trunc_normal_(m.weight, std=.02)
+ if isinstance(m, nn.Linear) and m.bias is not None:
+ nn.init.constant_(m.bias, 0)
+ elif isinstance(m, nn.LayerNorm):
+ nn.init.constant_(m.bias, 0)
+ nn.init.constant_(m.weight, 1.0)
+ elif isinstance(m, nn.Conv2d):
+ fan_out = m.kernel_size[0] * m.kernel_size[1] * m.out_channels
+ fan_out //= m.groups
+ m.weight.data.normal_(0, math.sqrt(2.0 / fan_out))
+ if m.bias is not None:
+ m.bias.data.zero_()
+
+ def init_weights(self, pretrained=None):
+ if isinstance(pretrained, str):
+ logger = 1
+ #load_checkpoint(self, pretrained, map_location='cpu', strict=False, logger=logger)
+
+ def reset_drop_path(self, drop_path_rate):
+ dpr = [x.item() for x in torch.linspace(0, drop_path_rate, sum(self.depths))]
+ cur = 0
+ for i in range(self.depths[0]):
+ self.block1[i].drop_path.drop_prob = dpr[cur + i]
+
+ cur += self.depths[0]
+ for i in range(self.depths[1]):
+ self.block2[i].drop_path.drop_prob = dpr[cur + i]
+
+ cur += self.depths[1]
+ for i in range(self.depths[2]):
+ self.block3[i].drop_path.drop_prob = dpr[cur + i]
+
+ cur += self.depths[2]
+ for i in range(self.depths[3]):
+ self.block4[i].drop_path.drop_prob = dpr[cur + i]
+
+ def freeze_patch_emb(self):
+ self.patch_embed1.requires_grad = False
+
+ @torch.jit.ignore
+ def no_weight_decay(self):
+ return {'pos_embed1', 'pos_embed2', 'pos_embed3', 'pos_embed4', 'cls_token'} # has pos_embed may be better
+
+ def get_classifier(self):
+ return self.head
+
+ def reset_classifier(self, num_classes, global_pool=''):
+ self.num_classes = num_classes
+ self.head = nn.Linear(self.embed_dim, num_classes) if num_classes > 0 else nn.Identity()
+
+ def forward_features(self, x):
+ B = x.shape[0]
+ outs = []
+
+ # stage 1
+ x, H, W = self.patch_embed1(x)
+ for i, blk in enumerate(self.block1):
+ x = blk(x, H, W)
+ x = self.norm1(x)
+ x = x.reshape(B, H, W, -1).permute(0, 3, 1, 2).contiguous()
+ outs.append(x)
+
+ # stage 2
+ x, H, W = self.patch_embed2(x)
+ for i, blk in enumerate(self.block2):
+ x = blk(x, H, W)
+ x = self.norm2(x)
+ x = x.reshape(B, H, W, -1).permute(0, 3, 1, 2).contiguous()
+ outs.append(x)
+
+ # stage 3
+ x, H, W = self.patch_embed3(x)
+ for i, blk in enumerate(self.block3):
+ x = blk(x, H, W)
+ x = self.norm3(x)
+ x = x.reshape(B, H, W, -1).permute(0, 3, 1, 2).contiguous()
+ outs.append(x)
+
+ # stage 4
+ x, H, W = self.patch_embed4(x)
+ for i, blk in enumerate(self.block4):
+ x = blk(x, H, W)
+ x = self.norm4(x)
+ x = x.reshape(B, H, W, -1).permute(0, 3, 1, 2).contiguous()
+ outs.append(x)
+
+ return outs
+
+ # return x.mean(dim=1)
+
+ def forward(self, x):
+ x = self.forward_features(x)
+ # x = self.head(x)
+
+ return x
+
+
+class DWConv(nn.Module):
+ def __init__(self, dim=768):
+ super(DWConv, self).__init__()
+ self.dwconv = nn.Conv2d(dim, dim, 3, 1, 1, bias=True, groups=dim)
+
+ def forward(self, x, H, W):
+ B, N, C = x.shape
+ x = x.transpose(1, 2).view(B, C, H, W).contiguous()
+ x = self.dwconv(x)
+ x = x.flatten(2).transpose(1, 2)
+
+ return x
+
+
+def _conv_filter(state_dict, patch_size=16):
+ """ convert patch embedding weight from manual patchify + linear proj to conv"""
+ out_dict = {}
+ for k, v in state_dict.items():
+ if 'patch_embed.proj.weight' in k:
+ v = v.reshape((v.shape[0], 3, patch_size, patch_size))
+ out_dict[k] = v
+
+ return out_dict
+
+
+## @register_model
+class pvt_v2_b0(PyramidVisionTransformerImpr):
+ def __init__(self, **kwargs):
+ super(pvt_v2_b0, self).__init__(
+ patch_size=4, embed_dims=[32, 64, 160, 256], num_heads=[1, 2, 5, 8], mlp_ratios=[8, 8, 4, 4],
+ qkv_bias=True, norm_layer=partial(nn.LayerNorm, eps=1e-6), depths=[2, 2, 2, 2], sr_ratios=[8, 4, 2, 1],
+ drop_rate=0.0, drop_path_rate=0.1)
+
+
+
+## @register_model
+class pvt_v2_b1(PyramidVisionTransformerImpr):
+ def __init__(self, **kwargs):
+ super(pvt_v2_b1, self).__init__(
+ patch_size=4, embed_dims=[64, 128, 320, 512], num_heads=[1, 2, 5, 8], mlp_ratios=[8, 8, 4, 4],
+ qkv_bias=True, norm_layer=partial(nn.LayerNorm, eps=1e-6), depths=[2, 2, 2, 2], sr_ratios=[8, 4, 2, 1],
+ drop_rate=0.0, drop_path_rate=0.1)
+
+## @register_model
+class pvt_v2_b2(PyramidVisionTransformerImpr):
+ def __init__(self, in_channels=3, **kwargs):
+ super(pvt_v2_b2, self).__init__(
+ patch_size=4, embed_dims=[64, 128, 320, 512], num_heads=[1, 2, 5, 8], mlp_ratios=[8, 8, 4, 4],
+ qkv_bias=True, norm_layer=partial(nn.LayerNorm, eps=1e-6), depths=[3, 4, 6, 3], sr_ratios=[8, 4, 2, 1],
+ drop_rate=0.0, drop_path_rate=0.1, in_channels=in_channels)
+
+## @register_model
+class pvt_v2_b3(PyramidVisionTransformerImpr):
+ def __init__(self, **kwargs):
+ super(pvt_v2_b3, self).__init__(
+ patch_size=4, embed_dims=[64, 128, 320, 512], num_heads=[1, 2, 5, 8], mlp_ratios=[8, 8, 4, 4],
+ qkv_bias=True, norm_layer=partial(nn.LayerNorm, eps=1e-6), depths=[3, 4, 18, 3], sr_ratios=[8, 4, 2, 1],
+ drop_rate=0.0, drop_path_rate=0.1)
+
+## @register_model
+class pvt_v2_b4(PyramidVisionTransformerImpr):
+ def __init__(self, **kwargs):
+ super(pvt_v2_b4, self).__init__(
+ patch_size=4, embed_dims=[64, 128, 320, 512], num_heads=[1, 2, 5, 8], mlp_ratios=[8, 8, 4, 4],
+ qkv_bias=True, norm_layer=partial(nn.LayerNorm, eps=1e-6), depths=[3, 8, 27, 3], sr_ratios=[8, 4, 2, 1],
+ drop_rate=0.0, drop_path_rate=0.1)
+
+
+## @register_model
+class pvt_v2_b5(PyramidVisionTransformerImpr):
+ def __init__(self, **kwargs):
+ super(pvt_v2_b5, self).__init__(
+ patch_size=4, embed_dims=[64, 128, 320, 512], num_heads=[1, 2, 5, 8], mlp_ratios=[4, 4, 4, 4],
+ qkv_bias=True, norm_layer=partial(nn.LayerNorm, eps=1e-6), depths=[3, 6, 40, 3], sr_ratios=[8, 4, 2, 1],
+ drop_rate=0.0, drop_path_rate=0.1)
diff --git a/py/BiRefNet_legacy/backbones/swin_v1.py b/py/BiRefNet_legacy/backbones/swin_v1.py
new file mode 100644
index 0000000..57501dc
--- /dev/null
+++ b/py/BiRefNet_legacy/backbones/swin_v1.py
@@ -0,0 +1,652 @@
+# --------------------------------------------------------
+# Swin Transformer
+# Copyright (c) 2021 Microsoft
+# Licensed under The MIT License [see LICENSE for details]
+# Written by Ze Liu, Yutong Lin, Yixuan Wei
+# --------------------------------------------------------
+
+import torch
+import torch.nn as nn
+import torch.nn.functional as F
+import torch.utils.checkpoint as checkpoint
+import numpy as np
+from timm.models.layers import DropPath, to_2tuple, trunc_normal_
+
+from ..config import Config
+
+
+config = Config()
+
+class Mlp(nn.Module):
+ """ Multilayer perceptron."""
+
+ def __init__(self, in_features, hidden_features=None, out_features=None, act_layer=nn.GELU, drop=0.):
+ super().__init__()
+ out_features = out_features or in_features
+ hidden_features = hidden_features or in_features
+ self.fc1 = nn.Linear(in_features, hidden_features)
+ self.act = act_layer()
+ self.fc2 = nn.Linear(hidden_features, out_features)
+ self.drop = nn.Dropout(drop)
+
+ def forward(self, x):
+ x = self.fc1(x)
+ x = self.act(x)
+ x = self.drop(x)
+ x = self.fc2(x)
+ x = self.drop(x)
+ return x
+
+
+def window_partition(x, window_size):
+ """
+ Args:
+ x: (B, H, W, C)
+ window_size (int): window size
+
+ Returns:
+ windows: (num_windows*B, window_size, window_size, C)
+ """
+ B, H, W, C = x.shape
+ x = x.view(B, H // window_size, window_size, W // window_size, window_size, C)
+ windows = x.permute(0, 1, 3, 2, 4, 5).contiguous().view(-1, window_size, window_size, C)
+ return windows
+
+
+def window_reverse(windows, window_size, H, W):
+ """
+ Args:
+ windows: (num_windows*B, window_size, window_size, C)
+ window_size (int): Window size
+ H (int): Height of image
+ W (int): Width of image
+
+ Returns:
+ x: (B, H, W, C)
+ """
+ B = int(windows.shape[0] / (H * W / window_size / window_size))
+ x = windows.view(B, H // window_size, W // window_size, window_size, window_size, -1)
+ x = x.permute(0, 1, 3, 2, 4, 5).contiguous().view(B, H, W, -1)
+ return x
+
+
+class WindowAttention(nn.Module):
+ """ Window based multi-head self attention (W-MSA) module with relative position bias.
+ It supports both of shifted and non-shifted window.
+
+ Args:
+ dim (int): Number of input channels.
+ window_size (tuple[int]): The height and width of the window.
+ num_heads (int): Number of attention heads.
+ qkv_bias (bool, optional): If True, add a learnable bias to query, key, value. Default: True
+ qk_scale (float | None, optional): Override default qk scale of head_dim ** -0.5 if set
+ attn_drop (float, optional): Dropout ratio of attention weight. Default: 0.0
+ proj_drop (float, optional): Dropout ratio of output. Default: 0.0
+ """
+
+ def __init__(self, dim, window_size, num_heads, qkv_bias=True, qk_scale=None, attn_drop=0., proj_drop=0.):
+
+ super().__init__()
+ self.dim = dim
+ self.window_size = window_size # Wh, Ww
+ self.num_heads = num_heads
+ head_dim = dim // num_heads
+ self.scale = qk_scale or head_dim ** -0.5
+
+ # define a parameter table of relative position bias
+ self.relative_position_bias_table = nn.Parameter(
+ torch.zeros((2 * window_size[0] - 1) * (2 * window_size[1] - 1), num_heads)) # 2*Wh-1 * 2*Ww-1, nH
+
+ # get pair-wise relative position index for each token inside the window
+ coords_h = torch.arange(self.window_size[0])
+ coords_w = torch.arange(self.window_size[1])
+ coords = torch.stack(torch.meshgrid([coords_h, coords_w], indexing='ij')) # 2, Wh, Ww
+ coords_flatten = torch.flatten(coords, 1) # 2, Wh*Ww
+ relative_coords = coords_flatten[:, :, None] - coords_flatten[:, None, :] # 2, Wh*Ww, Wh*Ww
+ relative_coords = relative_coords.permute(1, 2, 0).contiguous() # Wh*Ww, Wh*Ww, 2
+ relative_coords[:, :, 0] += self.window_size[0] - 1 # shift to start from 0
+ relative_coords[:, :, 1] += self.window_size[1] - 1
+ relative_coords[:, :, 0] *= 2 * self.window_size[1] - 1
+ relative_position_index = relative_coords.sum(-1) # Wh*Ww, Wh*Ww
+ self.register_buffer("relative_position_index", relative_position_index)
+
+ self.qkv = nn.Linear(dim, dim * 3, bias=qkv_bias)
+ self.attn_drop_prob = attn_drop
+ self.attn_drop = nn.Dropout(attn_drop)
+ self.proj = nn.Linear(dim, dim)
+ self.proj_drop = nn.Dropout(proj_drop)
+
+ trunc_normal_(self.relative_position_bias_table, std=.02)
+ self.softmax = nn.Softmax(dim=-1)
+
+ def forward(self, x, mask=None):
+ """ Forward function.
+
+ Args:
+ x: input features with shape of (num_windows*B, N, C)
+ mask: (0/-inf) mask with shape of (num_windows, Wh*Ww, Wh*Ww) or None
+ """
+ B_, N, C = x.shape
+ qkv = self.qkv(x).reshape(B_, N, 3, self.num_heads, C // self.num_heads).permute(2, 0, 3, 1, 4)
+ q, k, v = qkv[0], qkv[1], qkv[2] # make torchscript happy (cannot use tensor as tuple)
+
+ q = q * self.scale
+
+ if config.SDPA_enabled:
+ x = torch.nn.functional.scaled_dot_product_attention(
+ q, k, v,
+ attn_mask=None, dropout_p=self.attn_drop_prob, is_causal=False
+ ).transpose(1, 2).reshape(B_, N, C)
+ else:
+ attn = (q @ k.transpose(-2, -1))
+
+ relative_position_bias = self.relative_position_bias_table[self.relative_position_index.view(-1)].view(
+ self.window_size[0] * self.window_size[1], self.window_size[0] * self.window_size[1], -1) # Wh*Ww,Wh*Ww,nH
+ relative_position_bias = relative_position_bias.permute(2, 0, 1).contiguous() # nH, Wh*Ww, Wh*Ww
+ attn = attn + relative_position_bias.unsqueeze(0)
+
+ if mask is not None:
+ nW = mask.shape[0]
+ attn = attn.view(B_ // nW, nW, self.num_heads, N, N) + mask.unsqueeze(1).unsqueeze(0)
+ attn = attn.view(-1, self.num_heads, N, N)
+ attn = self.softmax(attn)
+ else:
+ attn = self.softmax(attn)
+
+ attn = self.attn_drop(attn)
+
+ x = (attn @ v).transpose(1, 2).reshape(B_, N, C)
+ x = self.proj(x)
+ x = self.proj_drop(x)
+ return x
+
+
+class SwinTransformerBlock(nn.Module):
+ """ Swin Transformer Block.
+
+ Args:
+ dim (int): Number of input channels.
+ num_heads (int): Number of attention heads.
+ window_size (int): Window size.
+ shift_size (int): Shift size for SW-MSA.
+ mlp_ratio (float): Ratio of mlp hidden dim to embedding dim.
+ qkv_bias (bool, optional): If True, add a learnable bias to query, key, value. Default: True
+ qk_scale (float | None, optional): Override default qk scale of head_dim ** -0.5 if set.
+ drop (float, optional): Dropout rate. Default: 0.0
+ attn_drop (float, optional): Attention dropout rate. Default: 0.0
+ drop_path (float, optional): Stochastic depth rate. Default: 0.0
+ act_layer (nn.Module, optional): Activation layer. Default: nn.GELU
+ norm_layer (nn.Module, optional): Normalization layer. Default: nn.LayerNorm
+ """
+
+ def __init__(self, dim, num_heads, window_size=7, shift_size=0,
+ mlp_ratio=4., qkv_bias=True, qk_scale=None, drop=0., attn_drop=0., drop_path=0.,
+ act_layer=nn.GELU, norm_layer=nn.LayerNorm):
+ super().__init__()
+ self.dim = dim
+ self.num_heads = num_heads
+ self.window_size = window_size
+ self.shift_size = shift_size
+ self.mlp_ratio = mlp_ratio
+ assert 0 <= self.shift_size < self.window_size, "shift_size must in 0-window_size"
+
+ self.norm1 = norm_layer(dim)
+ self.attn = WindowAttention(
+ dim, window_size=to_2tuple(self.window_size), num_heads=num_heads,
+ qkv_bias=qkv_bias, qk_scale=qk_scale, attn_drop=attn_drop, proj_drop=drop)
+
+ self.drop_path = DropPath(drop_path) if drop_path > 0. else nn.Identity()
+ self.norm2 = norm_layer(dim)
+ mlp_hidden_dim = int(dim * mlp_ratio)
+ self.mlp = Mlp(in_features=dim, hidden_features=mlp_hidden_dim, act_layer=act_layer, drop=drop)
+
+ self.H = None
+ self.W = None
+
+ def forward(self, x, mask_matrix):
+ """ Forward function.
+
+ Args:
+ x: Input feature, tensor size (B, H*W, C).
+ H, W: Spatial resolution of the input feature.
+ mask_matrix: Attention mask for cyclic shift.
+ """
+ B, L, C = x.shape
+ H, W = self.H, self.W
+ assert L == H * W, "input feature has wrong size"
+
+ shortcut = x
+ x = self.norm1(x)
+ x = x.view(B, H, W, C)
+
+ # pad feature maps to multiples of window size
+ pad_l = pad_t = 0
+ pad_r = (self.window_size - W % self.window_size) % self.window_size
+ pad_b = (self.window_size - H % self.window_size) % self.window_size
+ x = F.pad(x, (0, 0, pad_l, pad_r, pad_t, pad_b))
+ _, Hp, Wp, _ = x.shape
+
+ # cyclic shift
+ if self.shift_size > 0:
+ shifted_x = torch.roll(x, shifts=(-self.shift_size, -self.shift_size), dims=(1, 2))
+ attn_mask = mask_matrix
+ else:
+ shifted_x = x
+ attn_mask = None
+
+ # partition windows
+ x_windows = window_partition(shifted_x, self.window_size) # nW*B, window_size, window_size, C
+ x_windows = x_windows.view(-1, self.window_size * self.window_size, C) # nW*B, window_size*window_size, C
+
+ # W-MSA/SW-MSA
+ attn_windows = self.attn(x_windows, mask=attn_mask) # nW*B, window_size*window_size, C
+
+ # merge windows
+ attn_windows = attn_windows.view(-1, self.window_size, self.window_size, C)
+ shifted_x = window_reverse(attn_windows, self.window_size, Hp, Wp) # B H' W' C
+
+ # reverse cyclic shift
+ if self.shift_size > 0:
+ x = torch.roll(shifted_x, shifts=(self.shift_size, self.shift_size), dims=(1, 2))
+ else:
+ x = shifted_x
+
+ if pad_r > 0 or pad_b > 0:
+ x = x[:, :H, :W, :].contiguous()
+
+ x = x.view(B, H * W, C)
+
+ # FFN
+ x = shortcut + self.drop_path(x)
+ x = x + self.drop_path(self.mlp(self.norm2(x)))
+
+ return x
+
+
+class PatchMerging(nn.Module):
+ """ Patch Merging Layer
+
+ Args:
+ dim (int): Number of input channels.
+ norm_layer (nn.Module, optional): Normalization layer. Default: nn.LayerNorm
+ """
+ def __init__(self, dim, norm_layer=nn.LayerNorm):
+ super().__init__()
+ self.dim = dim
+ self.reduction = nn.Linear(4 * dim, 2 * dim, bias=False)
+ self.norm = norm_layer(4 * dim)
+
+ def forward(self, x, H, W):
+ """ Forward function.
+
+ Args:
+ x: Input feature, tensor size (B, H*W, C).
+ H, W: Spatial resolution of the input feature.
+ """
+ B, L, C = x.shape
+ assert L == H * W, "input feature has wrong size"
+
+ x = x.view(B, H, W, C)
+
+ # padding
+ pad_input = (H % 2 == 1) or (W % 2 == 1)
+ if pad_input:
+ x = F.pad(x, (0, 0, 0, W % 2, 0, H % 2))
+
+ x0 = x[:, 0::2, 0::2, :] # B H/2 W/2 C
+ x1 = x[:, 1::2, 0::2, :] # B H/2 W/2 C
+ x2 = x[:, 0::2, 1::2, :] # B H/2 W/2 C
+ x3 = x[:, 1::2, 1::2, :] # B H/2 W/2 C
+ x = torch.cat([x0, x1, x2, x3], -1) # B H/2 W/2 4*C
+ x = x.view(B, -1, 4 * C) # B H/2*W/2 4*C
+
+ x = self.norm(x)
+ x = self.reduction(x)
+
+ return x
+
+
+class BasicLayer(nn.Module):
+ """ A basic Swin Transformer layer for one stage.
+
+ Args:
+ dim (int): Number of feature channels
+ depth (int): Depths of this stage.
+ num_heads (int): Number of attention head.
+ window_size (int): Local window size. Default: 7.
+ mlp_ratio (float): Ratio of mlp hidden dim to embedding dim. Default: 4.
+ qkv_bias (bool, optional): If True, add a learnable bias to query, key, value. Default: True
+ qk_scale (float | None, optional): Override default qk scale of head_dim ** -0.5 if set.
+ drop (float, optional): Dropout rate. Default: 0.0
+ attn_drop (float, optional): Attention dropout rate. Default: 0.0
+ drop_path (float | tuple[float], optional): Stochastic depth rate. Default: 0.0
+ norm_layer (nn.Module, optional): Normalization layer. Default: nn.LayerNorm
+ downsample (nn.Module | None, optional): Downsample layer at the end of the layer. Default: None
+ use_checkpoint (bool): Whether to use checkpointing to save memory. Default: False.
+ """
+
+ def __init__(self,
+ dim,
+ depth,
+ num_heads,
+ window_size=7,
+ mlp_ratio=4.,
+ qkv_bias=True,
+ qk_scale=None,
+ drop=0.,
+ attn_drop=0.,
+ drop_path=0.,
+ norm_layer=nn.LayerNorm,
+ downsample=None,
+ use_checkpoint=False):
+ super().__init__()
+ self.window_size = window_size
+ self.shift_size = window_size // 2
+ self.depth = depth
+ self.use_checkpoint = use_checkpoint
+
+ # build blocks
+ self.blocks = nn.ModuleList([
+ SwinTransformerBlock(
+ dim=dim,
+ num_heads=num_heads,
+ window_size=window_size,
+ shift_size=0 if (i % 2 == 0) else window_size // 2,
+ mlp_ratio=mlp_ratio,
+ qkv_bias=qkv_bias,
+ qk_scale=qk_scale,
+ drop=drop,
+ attn_drop=attn_drop,
+ drop_path=drop_path[i] if isinstance(drop_path, list) else drop_path,
+ norm_layer=norm_layer)
+ for i in range(depth)])
+
+ # patch merging layer
+ if downsample is not None:
+ self.downsample = downsample(dim=dim, norm_layer=norm_layer)
+ else:
+ self.downsample = None
+
+ def forward(self, x, H, W):
+ """ Forward function.
+
+ Args:
+ x: Input feature, tensor size (B, H*W, C).
+ H, W: Spatial resolution of the input feature.
+ """
+
+ # calculate attention mask for SW-MSA
+ Hp = int(np.ceil(H / self.window_size)) * self.window_size
+ Wp = int(np.ceil(W / self.window_size)) * self.window_size
+ img_mask = torch.zeros((1, Hp, Wp, 1), device=x.device) # 1 Hp Wp 1
+ h_slices = (slice(0, -self.window_size),
+ slice(-self.window_size, -self.shift_size),
+ slice(-self.shift_size, None))
+ w_slices = (slice(0, -self.window_size),
+ slice(-self.window_size, -self.shift_size),
+ slice(-self.shift_size, None))
+ cnt = 0
+ for h in h_slices:
+ for w in w_slices:
+ img_mask[:, h, w, :] = cnt
+ cnt += 1
+
+ mask_windows = window_partition(img_mask, self.window_size) # nW, window_size, window_size, 1
+ mask_windows = mask_windows.view(-1, self.window_size * self.window_size)
+ attn_mask = mask_windows.unsqueeze(1) - mask_windows.unsqueeze(2)
+ attn_mask = attn_mask.masked_fill(attn_mask != 0, float(-100.0)).masked_fill(attn_mask == 0, float(0.0))
+
+ for blk in self.blocks:
+ blk.H, blk.W = H, W
+ if self.use_checkpoint:
+ x = checkpoint.checkpoint(blk, x, attn_mask)
+ else:
+ x = blk(x, attn_mask)
+ if self.downsample is not None:
+ x_down = self.downsample(x, H, W)
+ Wh, Ww = (H + 1) // 2, (W + 1) // 2
+ return x, H, W, x_down, Wh, Ww
+ else:
+ return x, H, W, x, H, W
+
+
+class PatchEmbed(nn.Module):
+ """ Image to Patch Embedding
+
+ Args:
+ patch_size (int): Patch token size. Default: 4.
+ in_channels (int): Number of input image channels. Default: 3.
+ embed_dim (int): Number of linear projection output channels. Default: 96.
+ norm_layer (nn.Module, optional): Normalization layer. Default: None
+ """
+
+ def __init__(self, patch_size=4, in_channels=3, embed_dim=96, norm_layer=None):
+ super().__init__()
+ patch_size = to_2tuple(patch_size)
+ self.patch_size = patch_size
+
+ self.in_channels = in_channels
+ self.embed_dim = embed_dim
+
+ self.proj = nn.Conv2d(in_channels, embed_dim, kernel_size=patch_size, stride=patch_size)
+ if norm_layer is not None:
+ self.norm = norm_layer(embed_dim)
+ else:
+ self.norm = None
+
+ def forward(self, x):
+ """Forward function."""
+ # padding
+ _, _, H, W = x.size()
+ if W % self.patch_size[1] != 0:
+ x = F.pad(x, (0, self.patch_size[1] - W % self.patch_size[1]))
+ if H % self.patch_size[0] != 0:
+ x = F.pad(x, (0, 0, 0, self.patch_size[0] - H % self.patch_size[0]))
+
+ x = self.proj(x) # B C Wh Ww
+ if self.norm is not None:
+ Wh, Ww = x.size(2), x.size(3)
+ x = x.flatten(2).transpose(1, 2)
+ x = self.norm(x)
+ x = x.transpose(1, 2).view(-1, self.embed_dim, Wh, Ww)
+
+ return x
+
+
+class SwinTransformer(nn.Module):
+ """ Swin Transformer backbone.
+ A PyTorch impl of : `Swin Transformer: Hierarchical Vision Transformer using Shifted Windows` -
+ https://arxiv.org/pdf/2103.14030
+
+ Args:
+ pretrain_img_size (int): Input image size for training the pretrained model,
+ used in absolute postion embedding. Default 224.
+ patch_size (int | tuple(int)): Patch size. Default: 4.
+ in_channels (int): Number of input image channels. Default: 3.
+ embed_dim (int): Number of linear projection output channels. Default: 96.
+ depths (tuple[int]): Depths of each Swin Transformer stage.
+ num_heads (tuple[int]): Number of attention head of each stage.
+ window_size (int): Window size. Default: 7.
+ mlp_ratio (float): Ratio of mlp hidden dim to embedding dim. Default: 4.
+ qkv_bias (bool): If True, add a learnable bias to query, key, value. Default: True
+ qk_scale (float): Override default qk scale of head_dim ** -0.5 if set.
+ drop_rate (float): Dropout rate.
+ attn_drop_rate (float): Attention dropout rate. Default: 0.
+ drop_path_rate (float): Stochastic depth rate. Default: 0.2.
+ norm_layer (nn.Module): Normalization layer. Default: nn.LayerNorm.
+ ape (bool): If True, add absolute position embedding to the patch embedding. Default: False.
+ patch_norm (bool): If True, add normalization after patch embedding. Default: True.
+ out_indices (Sequence[int]): Output from which stages.
+ frozen_stages (int): Stages to be frozen (stop grad and set eval mode).
+ -1 means not freezing any parameters.
+ use_checkpoint (bool): Whether to use checkpointing to save memory. Default: False.
+ """
+
+ def __init__(self,
+ pretrain_img_size=224,
+ patch_size=4,
+ in_channels=3,
+ embed_dim=96,
+ depths=[2, 2, 6, 2],
+ num_heads=[3, 6, 12, 24],
+ window_size=7,
+ mlp_ratio=4.,
+ qkv_bias=True,
+ qk_scale=None,
+ drop_rate=0.,
+ attn_drop_rate=0.,
+ drop_path_rate=0.2,
+ norm_layer=nn.LayerNorm,
+ ape=False,
+ patch_norm=True,
+ out_indices=(0, 1, 2, 3),
+ frozen_stages=-1,
+ use_checkpoint=False):
+ super().__init__()
+
+ self.pretrain_img_size = pretrain_img_size
+ self.num_layers = len(depths)
+ self.embed_dim = embed_dim
+ self.ape = ape
+ self.patch_norm = patch_norm
+ self.out_indices = out_indices
+ self.frozen_stages = frozen_stages
+
+ # split image into non-overlapping patches
+ self.patch_embed = PatchEmbed(
+ patch_size=patch_size, in_channels=in_channels, embed_dim=embed_dim,
+ norm_layer=norm_layer if self.patch_norm else None)
+
+ # absolute position embedding
+ if self.ape:
+ pretrain_img_size = to_2tuple(pretrain_img_size)
+ patch_size = to_2tuple(patch_size)
+ patches_resolution = [pretrain_img_size[0] // patch_size[0], pretrain_img_size[1] // patch_size[1]]
+
+ self.absolute_pos_embed = nn.Parameter(torch.zeros(1, embed_dim, patches_resolution[0], patches_resolution[1]))
+ trunc_normal_(self.absolute_pos_embed, std=.02)
+
+ self.pos_drop = nn.Dropout(p=drop_rate)
+
+ # stochastic depth
+ dpr = [x.item() for x in torch.linspace(0, drop_path_rate, sum(depths))] # stochastic depth decay rule
+
+ # build layers
+ self.layers = nn.ModuleList()
+ for i_layer in range(self.num_layers):
+ layer = BasicLayer(
+ dim=int(embed_dim * 2 ** i_layer),
+ depth=depths[i_layer],
+ num_heads=num_heads[i_layer],
+ window_size=window_size,
+ mlp_ratio=mlp_ratio,
+ qkv_bias=qkv_bias,
+ qk_scale=qk_scale,
+ drop=drop_rate,
+ attn_drop=attn_drop_rate,
+ drop_path=dpr[sum(depths[:i_layer]):sum(depths[:i_layer + 1])],
+ norm_layer=norm_layer,
+ downsample=PatchMerging if (i_layer < self.num_layers - 1) else None,
+ use_checkpoint=use_checkpoint)
+ self.layers.append(layer)
+
+ num_features = [int(embed_dim * 2 ** i) for i in range(self.num_layers)]
+ self.num_features = num_features
+
+ # add a norm layer for each output
+ for i_layer in out_indices:
+ layer = norm_layer(num_features[i_layer])
+ layer_name = f'norm{i_layer}'
+ self.add_module(layer_name, layer)
+
+ self._freeze_stages()
+
+ def _freeze_stages(self):
+ if self.frozen_stages >= 0:
+ self.patch_embed.eval()
+ for param in self.patch_embed.parameters():
+ param.requires_grad = False
+
+ if self.frozen_stages >= 1 and self.ape:
+ self.absolute_pos_embed.requires_grad = False
+
+ if self.frozen_stages >= 2:
+ self.pos_drop.eval()
+ for i in range(0, self.frozen_stages - 1):
+ m = self.layers[i]
+ m.eval()
+ for param in m.parameters():
+ param.requires_grad = False
+
+ def init_weights(self, pretrained=None):
+ """Initialize the weights in backbone.
+
+ Args:
+ pretrained (str, optional): Path to pre-trained weights.
+ Defaults to None.
+ """
+
+ def _init_weights(m):
+ if isinstance(m, nn.Linear):
+ trunc_normal_(m.weight, std=.02)
+ if isinstance(m, nn.Linear) and m.bias is not None:
+ nn.init.constant_(m.bias, 0)
+ elif isinstance(m, nn.LayerNorm):
+ nn.init.constant_(m.bias, 0)
+ nn.init.constant_(m.weight, 1.0)
+
+ if isinstance(pretrained, str):
+ self.apply(_init_weights)
+ logger = get_root_logger()
+ load_checkpoint(self, pretrained, strict=False, logger=logger)
+ elif pretrained is None:
+ self.apply(_init_weights)
+ else:
+ raise TypeError('pretrained must be a str or None')
+
+ def forward(self, x):
+ """Forward function."""
+ x = self.patch_embed(x)
+
+ Wh, Ww = x.size(2), x.size(3)
+ if self.ape:
+ # interpolate the position embedding to the corresponding size
+ absolute_pos_embed = F.interpolate(self.absolute_pos_embed, size=(Wh, Ww), mode='bicubic')
+ x = (x + absolute_pos_embed) # B Wh*Ww C
+
+ outs = []#x.contiguous()]
+ x = x.flatten(2).transpose(1, 2)
+ x = self.pos_drop(x)
+ for i in range(self.num_layers):
+ layer = self.layers[i]
+ x_out, H, W, x, Wh, Ww = layer(x, Wh, Ww)
+
+ if i in self.out_indices:
+ norm_layer = getattr(self, f'norm{i}')
+ x_out = norm_layer(x_out)
+
+ out = x_out.view(-1, H, W, self.num_features[i]).permute(0, 3, 1, 2).contiguous()
+ outs.append(out)
+
+ return tuple(outs)
+
+ def train(self, mode=True):
+ """Convert the model into training mode while keep layers freezed."""
+ super(SwinTransformer, self).train(mode)
+ self._freeze_stages()
+
+def swin_v1_t():
+ model = SwinTransformer(embed_dim=96, depths=[2, 2, 6, 2], num_heads=[3, 6, 12, 24], window_size=7)
+ return model
+
+def swin_v1_s():
+ model = SwinTransformer(embed_dim=96, depths=[2, 2, 18, 2], num_heads=[3, 6, 12, 24], window_size=7)
+ return model
+
+def swin_v1_b():
+ model = SwinTransformer(embed_dim=128, depths=[2, 2, 18, 2], num_heads=[4, 8, 16, 32], window_size=12)
+ return model
+
+def swin_v1_l():
+ model = SwinTransformer(embed_dim=192, depths=[2, 2, 18, 2], num_heads=[6, 12, 24, 48], window_size=12)
+ return model
diff --git a/py/BiRefNet_legacy/baseline.py b/py/BiRefNet_legacy/baseline.py
new file mode 100644
index 0000000..e2bb6a9
--- /dev/null
+++ b/py/BiRefNet_legacy/baseline.py
@@ -0,0 +1,292 @@
+import torch
+import torch.nn as nn
+from collections import OrderedDict
+import torch
+import torch.nn as nn
+import torch.nn.functional as F
+from torchvision.models import vgg16, vgg16_bn
+from torchvision.models import resnet50
+from kornia.filters import laplacian
+
+from BiRefNet_legacy.backbones.build_backbone import build_backbone
+from BiRefNet_legacy.modules.decoder_blocks import BasicDecBlk, ResBlk, HierarAttDecBlk
+from BiRefNet_legacy.modules.lateral_blocks import BasicLatBlk
+from BiRefNet_legacy.modules.aspp import ASPP, ASPPDeformable
+from BiRefNet_legacy.modules.ing import *
+from BiRefNet_legacy.refinement.refiner import Refiner, RefinerPVTInChannels4, RefUNet
+from BiRefNet_legacy.refinement.stem_layer import StemLayer
+
+from .config import Config
+from .dataset import class_labels_TR_sorted
+
+
+class BiRefNet(nn.Module):
+ def __init__(self):
+ super(BiRefNet, self).__init__()
+ self.config = Config()
+ self.epoch = 1
+ self.bb = build_backbone(self.config.bb, pretrained=True)
+
+ channels = self.config.lateral_channels_in_collection
+
+ if self.config.auxiliary_classification:
+ self.avgpool = nn.AdaptiveAvgPool2d((1, 1))
+ self.cls_head = nn.Sequential(
+ nn.Linear(channels[0], len(class_labels_TR_sorted))
+ )
+
+ if self.config.squeeze_block:
+ self.squeeze_module = nn.Sequential(*[
+ eval(self.config.squeeze_block.split('_x')[0])(channels[0]+sum(self.config.cxt), channels[0])
+ for _ in range(eval(self.config.squeeze_block.split('_x')[1]))
+ ])
+
+ self.decoder = Decoder(channels)
+
+ if self.config.locate_head:
+ self.locate_header = nn.ModuleList([
+ BasicDecBlk(channels[0], channels[-1]),
+ nn.Sequential(
+ nn.Conv2d(channels[-1], 1, 1, 1, 0),
+ )
+ ])
+
+ if self.config.ender:
+ self.dec_end = nn.Sequential(
+ nn.Conv2d(1, 16, 3, 1, 1),
+ nn.Conv2d(16, 1, 3, 1, 1),
+ nn.ReLU(inplace=True),
+ )
+
+ # refine patch-level segmentation
+ if self.config.refine:
+ if self.config.refine == 'itself':
+ self.stem_layer = StemLayer(in_channels=3+1, inter_channels=48, out_channels=3)
+ else:
+ self.refiner = eval('{}({})'.format(self.config.refine, 'in_channels=3+1'))
+
+ if self.config.freeze_bb:
+ # Freeze the backbone...
+ print(self.named_parameters())
+ for key, value in self.named_parameters():
+ if 'bb.' in key and 'refiner.' not in key:
+ value.requires_grad = False
+
+ def forward_enc(self, x):
+ if self.config.bb in ['vgg16', 'vgg16bn', 'resnet50']:
+ x1 = self.bb.conv1(x); x2 = self.bb.conv2(x1); x3 = self.bb.conv3(x2); x4 = self.bb.conv4(x3)
+ else:
+ x1, x2, x3, x4 = self.bb(x)
+ if self.config.mul_scl_ipt == 'cat':
+ B, C, H, W = x.shape
+ x1_, x2_, x3_, x4_ = self.bb(F.interpolate(x, size=(H//2, W//2), mode='bilinear', align_corners=True))
+ x1 = torch.cat([x1, F.interpolate(x1_, size=x1.shape[2:], mode='bilinear', align_corners=True)], dim=1)
+ x2 = torch.cat([x2, F.interpolate(x2_, size=x2.shape[2:], mode='bilinear', align_corners=True)], dim=1)
+ x3 = torch.cat([x3, F.interpolate(x3_, size=x3.shape[2:], mode='bilinear', align_corners=True)], dim=1)
+ x4 = torch.cat([x4, F.interpolate(x4_, size=x4.shape[2:], mode='bilinear', align_corners=True)], dim=1)
+ elif self.config.mul_scl_ipt == 'add':
+ B, C, H, W = x.shape
+ x1_, x2_, x3_, x4_ = self.bb(F.interpolate(x, size=(H//2, W//2), mode='bilinear', align_corners=True))
+ x1 = x1 + F.interpolate(x1_, size=x1.shape[2:], mode='bilinear', align_corners=True)
+ x2 = x2 + F.interpolate(x2_, size=x2.shape[2:], mode='bilinear', align_corners=True)
+ x3 = x3 + F.interpolate(x3_, size=x3.shape[2:], mode='bilinear', align_corners=True)
+ x4 = x4 + F.interpolate(x4_, size=x4.shape[2:], mode='bilinear', align_corners=True)
+ class_preds = self.cls_head(self.avgpool(x4).view(x4.shape[0], -1)) if self.training and self.config.auxiliary_classification else None
+ if self.config.cxt:
+ x4 = torch.cat(
+ (
+ *[
+ F.interpolate(x1, size=x4.shape[2:], mode='bilinear', align_corners=True),
+ F.interpolate(x2, size=x4.shape[2:], mode='bilinear', align_corners=True),
+ F.interpolate(x3, size=x4.shape[2:], mode='bilinear', align_corners=True),
+ ][-len(self.config.cxt):],
+ x4
+ ),
+ dim=1
+ )
+ return (x1, x2, x3, x4), class_preds
+
+ def forward_ori(self, x):
+ ########## Encoder ##########
+ (x1, x2, x3, x4), class_preds = self.forward_enc(x)
+ if self.config.squeeze_block:
+ x4 = self.squeeze_module(x4)
+ ########## Decoder ##########
+ features = [x, x1, x2, x3, x4]
+ if self.config.out_ref:
+ features.append(laplacian(torch.mean(x, dim=1).unsqueeze(1), kernel_size=5))
+ scaled_preds = self.decoder(features)
+ return scaled_preds, class_preds
+
+ def forward_ref(self, x, pred):
+ # refine patch-level segmentation
+ if pred.shape[2:] != x.shape[2:]:
+ pred = F.interpolate(pred, size=x.shape[2:], mode='bilinear', align_corners=True)
+ # pred = pred.sigmoid()
+ if self.config.refine == 'itself':
+ x = self.stem_layer(torch.cat([x, pred], dim=1))
+ scaled_preds, class_preds = self.forward_ori(x)
+ else:
+ scaled_preds = self.refiner([x, pred])
+ class_preds = None
+ return scaled_preds, class_preds
+
+ def forward_ref_end(self, x):
+ # remove the grids of concatenated preds
+ return self.dec_end(x) if self.config.ender else x
+
+
+ def forward(self, x):
+ scaled_preds, class_preds = self.forward_ori(x)
+ class_preds_lst = [class_preds]
+ return [scaled_preds, class_preds_lst] if self.training else scaled_preds
+
+
+class Decoder(nn.Module):
+ def __init__(self, channels):
+ super(Decoder, self).__init__()
+ self.config = Config()
+ DecoderBlock = eval(self.config.dec_blk)
+ LateralBlock = eval(self.config.lat_blk)
+
+ if self.config.dec_ipt:
+ self.split = self.config.dec_ipt_split
+ N_dec_ipt = 64
+ DBlock = SimpleConvs
+ ic = 64
+ ipt_cha_opt = 1
+ self.ipt_blk4 = DBlock(2**8*3 if self.split else 3, [N_dec_ipt, channels[0]//8][ipt_cha_opt], inter_channels=ic)
+ self.ipt_blk3 = DBlock(2**6*3 if self.split else 3, [N_dec_ipt, channels[1]//8][ipt_cha_opt], inter_channels=ic)
+ self.ipt_blk2 = DBlock(2**4*3 if self.split else 3, [N_dec_ipt, channels[2]//8][ipt_cha_opt], inter_channels=ic)
+ self.ipt_blk1 = DBlock(2**0*3 if self.split else 3, [N_dec_ipt, channels[3]//8][ipt_cha_opt], inter_channels=ic)
+ else:
+ self.split = None
+
+ self.decoder_block4 = DecoderBlock(channels[0], channels[1])
+ self.decoder_block3 = DecoderBlock(channels[1]+([N_dec_ipt, channels[0]//8][ipt_cha_opt] if self.config.dec_ipt else 0), channels[2])
+ self.decoder_block2 = DecoderBlock(channels[2]+([N_dec_ipt, channels[1]//8][ipt_cha_opt] if self.config.dec_ipt else 0), channels[3])
+ self.decoder_block1 = DecoderBlock(channels[3]+([N_dec_ipt, channels[2]//8][ipt_cha_opt] if self.config.dec_ipt else 0), channels[3]//2)
+ self.conv_out1 = nn.Sequential(nn.Conv2d(channels[3]//2+([N_dec_ipt, channels[3]//8][ipt_cha_opt] if self.config.dec_ipt else 0), 1, 1, 1, 0))
+
+ self.lateral_block4 = LateralBlock(channels[1], channels[1])
+ self.lateral_block3 = LateralBlock(channels[2], channels[2])
+ self.lateral_block2 = LateralBlock(channels[3], channels[3])
+
+ if self.config.ms_supervision:
+ self.conv_ms_spvn_4 = nn.Conv2d(channels[1], 1, 1, 1, 0)
+ self.conv_ms_spvn_3 = nn.Conv2d(channels[2], 1, 1, 1, 0)
+ self.conv_ms_spvn_2 = nn.Conv2d(channels[3], 1, 1, 1, 0)
+
+ if self.config.out_ref:
+ _N = 16
+ # self.gdt_convs_4 = nn.Sequential(nn.Conv2d(channels[1], _N, 3, 1, 1), nn.BatchNorm2d(_N), nn.ReLU(inplace=True))
+ self.gdt_convs_3 = nn.Sequential(nn.Conv2d(channels[2], _N, 3, 1, 1), nn.BatchNorm2d(_N), nn.ReLU(inplace=True))
+ self.gdt_convs_2 = nn.Sequential(nn.Conv2d(channels[3], _N, 3, 1, 1), nn.BatchNorm2d(_N), nn.ReLU(inplace=True))
+
+ # self.gdt_convs_pred_4 = nn.Sequential(nn.Conv2d(_N, 1, 1, 1, 0))
+ self.gdt_convs_pred_3 = nn.Sequential(nn.Conv2d(_N, 1, 1, 1, 0))
+ self.gdt_convs_pred_2 = nn.Sequential(nn.Conv2d(_N, 1, 1, 1, 0))
+
+ # self.gdt_convs_attn_4 = nn.Sequential(nn.Conv2d(_N, 1, 1, 1, 0))
+ self.gdt_convs_attn_3 = nn.Sequential(nn.Conv2d(_N, 1, 1, 1, 0))
+ self.gdt_convs_attn_2 = nn.Sequential(nn.Conv2d(_N, 1, 1, 1, 0))
+
+
+ def get_patches_batch(self, x, p):
+ _size_h, _size_w = p.shape[2:]
+ patches_batch = []
+ for idx in range(x.shape[0]):
+ columns_x = torch.split(x[idx], split_size_or_sections=_size_w, dim=-1)
+ patches_x = []
+ for column_x in columns_x:
+ patches_x += [p.unsqueeze(0) for p in torch.split(column_x, split_size_or_sections=_size_h, dim=-2)]
+ patch_sample = torch.cat(patches_x, dim=1)
+ patches_batch.append(patch_sample)
+ return torch.cat(patches_batch, dim=0)
+
+ def forward(self, features):
+ if self.config.out_ref:
+ outs_gdt_pred = []
+ outs_gdt_label = []
+ x, x1, x2, x3, x4, gdt_gt = features
+ else:
+ x, x1, x2, x3, x4 = features
+ outs = []
+ p4 = self.decoder_block4(x4)
+ m4 = self.conv_ms_spvn_4(p4) if self.config.ms_supervision else None
+ _p4 = F.interpolate(p4, size=x3.shape[2:], mode='bilinear', align_corners=True)
+ _p3 = _p4 + self.lateral_block4(x3)
+ if self.config.dec_ipt:
+ patches_batch = self.get_patches_batch(x, _p3) if self.split else x
+ _p3 = torch.cat((_p3, self.ipt_blk4(F.interpolate(patches_batch, size=x3.shape[2:], mode='bilinear', align_corners=True))), 1)
+
+ p3 = self.decoder_block3(_p3)
+ m3 = self.conv_ms_spvn_3(p3) if self.config.ms_supervision else None
+ if self.config.out_ref:
+ # >> GT:
+ # m3 --dilation--> m3_dia
+ # G_3^gt * m3_dia --> G_3^m, which is the label of gradient
+ m3_dia = m3
+ gdt_label_main_3 = gdt_gt * F.interpolate(m3_dia, size=gdt_gt.shape[2:], mode='bilinear', align_corners=True)
+ outs_gdt_label.append(gdt_label_main_3)
+ # >> Pred:
+ # p3 --conv--BN--> F_3^G, where F_3^G predicts the \hat{G_3} with xx
+ # F_3^G --sigmoid--> A_3^G
+ p3_gdt = self.gdt_convs_3(p3)
+ gdt_pred_3 = self.gdt_convs_pred_3(p3_gdt)
+ outs_gdt_pred.append(gdt_pred_3)
+ gdt_attn_3 = self.gdt_convs_attn_3(p3_gdt).sigmoid()
+ # >> Finally:
+ # p3 = p3 * A_3^G
+ p3 = p3 * gdt_attn_3
+ _p3 = F.interpolate(p3, size=x2.shape[2:], mode='bilinear', align_corners=True)
+ _p2 = _p3 + self.lateral_block3(x2)
+ if self.config.dec_ipt:
+ patches_batch = self.get_patches_batch(x, _p2) if self.split else x
+ _p2 = torch.cat((_p2, self.ipt_blk3(F.interpolate(patches_batch, size=x2.shape[2:], mode='bilinear', align_corners=True))), 1)
+
+ p2 = self.decoder_block2(_p2)
+ m2 = self.conv_ms_spvn_2(p2) if self.config.ms_supervision else None
+ if self.config.out_ref:
+ # >> GT:
+ m2_dia = m2
+ gdt_label_main_2 = gdt_gt * F.interpolate(m2_dia, size=gdt_gt.shape[2:], mode='bilinear', align_corners=True)
+ outs_gdt_label.append(gdt_label_main_2)
+ # >> Pred:
+ p2_gdt = self.gdt_convs_2(p2)
+ gdt_pred_2 = self.gdt_convs_pred_2(p2_gdt)
+ outs_gdt_pred.append(gdt_pred_2)
+ gdt_attn_2 = self.gdt_convs_attn_2(p2_gdt).sigmoid()
+ # >> Finally:
+ p2 = p2 * gdt_attn_2
+ _p2 = F.interpolate(p2, size=x1.shape[2:], mode='bilinear', align_corners=True)
+ _p1 = _p2 + self.lateral_block2(x1)
+ if self.config.dec_ipt:
+ patches_batch = self.get_patches_batch(x, _p1) if self.split else x
+ _p1 = torch.cat((_p1, self.ipt_blk2(F.interpolate(patches_batch, size=x1.shape[2:], mode='bilinear', align_corners=True))), 1)
+
+ _p1 = self.decoder_block1(_p1)
+ _p1 = F.interpolate(_p1, size=x.shape[2:], mode='bilinear', align_corners=True)
+ if self.config.dec_ipt:
+ patches_batch = self.get_patches_batch(x, _p1) if self.split else x
+ _p1 = torch.cat((_p1, self.ipt_blk1(F.interpolate(patches_batch, size=x.shape[2:], mode='bilinear', align_corners=True))), 1)
+ p1_out = self.conv_out1(_p1)
+
+ if self.config.ms_supervision:
+ outs.append(m4)
+ outs.append(m3)
+ outs.append(m2)
+ outs.append(p1_out)
+ return outs if not (self.config.out_ref and self.training) else ([outs_gdt_pred, outs_gdt_label], outs)
+
+
+class SimpleConvs(nn.Module):
+ def __init__(
+ self, in_channels: int, out_channels: int, inter_channels=64
+ ) -> None:
+ super().__init__()
+ self.conv1 = nn.Conv2d(in_channels, inter_channels, 3, 1, 1)
+ self.conv_out = nn.Conv2d(inter_channels, out_channels, 3, 1, 1)
+
+ def forward(self, x):
+ return self.conv_out(self.conv1(x))
diff --git a/py/BiRefNet_legacy/config.py b/py/BiRefNet_legacy/config.py
new file mode 100644
index 0000000..e7ee155
--- /dev/null
+++ b/py/BiRefNet_legacy/config.py
@@ -0,0 +1,104 @@
+import os
+import math
+from folder_paths import models_dir
+
+
+class Config():
+ def __init__(self) -> None:
+ self.ms_supervision = True
+ self.out_ref = self.ms_supervision and True
+ self.dec_ipt = True
+ self.dec_ipt_split = True
+ self.locate_head = False
+ self.cxt_num = [0, 3][1] # multi-scale skip connections from encoder
+ self.mul_scl_ipt = ['', 'add', 'cat'][2]
+ self.refine = ['', 'itself', 'RefUNet', 'Refiner', 'RefinerPVTInChannels4'][0]
+ self.progressive_ref = self.refine and True
+ self.ender = self.progressive_ref and False
+ self.scale = self.progressive_ref and 2
+ self.dec_att = ['', 'ASPP', 'ASPPDeformable'][2]
+ self.squeeze_block = ['', 'BasicDecBlk_x1', 'ResBlk_x4', 'ASPP_x3', 'ASPPDeformable_x3'][1]
+ self.dec_blk = ['BasicDecBlk', 'ResBlk', 'HierarAttDecBlk'][0]
+ self.auxiliary_classification = False
+ self.refine_iteration = 1
+ self.freeze_bb = False
+ self.precisionHigh = True
+ self.compile = True
+ self.load_all = True
+ self.verbose_eval = True
+
+ self.size = 1024
+ self.batch_size = 2
+ self.IoU_finetune_last_epochs = [0, -40][1] # choose 0 to skip
+ if self.dec_blk == 'HierarAttDecBlk':
+ self.batch_size = 2 ** [0, 1, 2, 3, 4][2]
+ self.model = [
+ 'BiRefNet',
+ ][0]
+
+ # Components
+ self.lat_blk = ['BasicLatBlk'][0]
+ self.dec_channels_inter = ['fixed', 'adap'][0]
+
+ # Backbone
+ self.bb = [
+ 'vgg16', 'vgg16bn', 'resnet50', # 0, 1, 2
+ 'pvt_v2_b2', 'pvt_v2_b5', # 3-bs10, 4-bs5
+ 'swin_v1_b', 'swin_v1_l' # 5-bs9, 6-bs6
+ ][6]
+ self.lateral_channels_in_collection = {
+ 'vgg16': [512, 256, 128, 64], 'vgg16bn': [512, 256, 128, 64], 'resnet50': [1024, 512, 256, 64],
+ 'pvt_v2_b2': [512, 320, 128, 64], 'pvt_v2_b5': [512, 320, 128, 64],
+ 'swin_v1_b': [1024, 512, 256, 128], 'swin_v1_l': [1536, 768, 384, 192],
+ }[self.bb]
+ if self.mul_scl_ipt == 'cat':
+ self.lateral_channels_in_collection = [channel * 2 for channel in self.lateral_channels_in_collection]
+ self.cxt = self.lateral_channels_in_collection[1:][::-1][-self.cxt_num:] if self.cxt_num else []
+ self.sys_home_dir = models_dir
+ self.weights_root_dir = os.path.join(self.sys_home_dir, "BiRefNet")
+ self.weights = {
+ 'pvt_v2_b2': os.path.join(self.weights_root_dir, 'pvt_v2_b2.pth'),
+ 'pvt_v2_b5': os.path.join(self.weights_root_dir, ['pvt_v2_b5.pth', 'pvt_v2_b5_22k.pth'][0]),
+ 'swin_v1_b': os.path.join(self.weights_root_dir, ['swin_base_patch4_window12_384_22kto1k.pth', 'swin_base_patch4_window12_384_22k.pth'][0]),
+ 'swin_v1_l': os.path.join(self.weights_root_dir, ['swin_large_patch4_window12_384_22kto1k.pth', 'swin_large_patch4_window12_384_22k.pth'][0]),
+ }
+
+ # Training
+ self.num_workers = 5 # will be decrease to min(it, batch_size) at the initialization of the data_loader
+ self.optimizer = ['Adam', 'AdamW'][0]
+ self.lr = 1e-5 * math.sqrt(self.batch_size / 5) # adapt the lr linearly
+ self.lr_decay_epochs = [1e4] # Set to negative N to decay the lr in the last N-th epoch.
+ self.lr_decay_rate = 0.5
+ self.only_S_MAE = False
+ self.SDPA_enabled = False # Bug. Slower and errors occur in multi-GPUs
+
+ # Data
+ self.data_root_dir = os.path.join(self.sys_home_dir, 'datasets/dis')
+ self.dataset = ['DIS5K', 'COD', 'SOD'][0]
+ self.preproc_methods = ['flip', 'enhance', 'rotate', 'pepper', 'crop'][:4]
+
+ # Loss
+ self.lambdas_pix_last = {
+ # not 0 means opening this loss
+ # original rate -- 1 : 30 : 1.5 : 0.2, bce x 30
+ 'bce': 30 * 1, # high performance
+ 'iou': 0.5 * 1, # 0 / 255
+ 'iou_patch': 0.5 * 0, # 0 / 255, win_size = (64, 64)
+ 'mse': 150 * 0, # can smooth the saliency map
+ 'triplet': 3 * 0,
+ 'reg': 100 * 0,
+ 'ssim': 10 * 1, # help contours,
+ 'cnt': 5 * 0, # help contours
+ }
+ self.lambdas_cls = {
+ 'ce': 5.0
+ }
+ # Adv
+ self.lambda_adv_g = 10. * 0 # turn to 0 to avoid adv training
+ self.lambda_adv_d = 3. * (self.lambda_adv_g > 0)
+
+ # others
+ self.device = [0, 'cpu'][0] # .to(0) = .to('cuda:0')
+
+ self.batch_size_valid = 1
+ self.rand_seed = 7
diff --git a/py/BiRefNet_legacy/dataset.py b/py/BiRefNet_legacy/dataset.py
new file mode 100644
index 0000000..cd4280c
--- /dev/null
+++ b/py/BiRefNet_legacy/dataset.py
@@ -0,0 +1,140 @@
+import os
+import cv2
+from tqdm import tqdm
+from PIL import Image
+from torch.utils import data
+from torchvision import transforms
+
+from .preproc import preproc
+from .config import Config
+from glob import glob
+
+
+Image.MAX_IMAGE_PIXELS = None # remove DecompressionBombWarning
+config = Config()
+_class_labels_TR_sorted = 'Airplane, Ant, Antenna, Archery, Axe, BabyCarriage, Bag, BalanceBeam, Balcony, Balloon, Basket, BasketballHoop, Beatle, Bed, Bee, Bench, Bicycle, BicycleFrame, BicycleStand, Boat, Bonsai, BoomLift, Bridge, BunkBed, Butterfly, Button, Cable, CableLift, Cage, Camcorder, Cannon, Canoe, Car, CarParkDropArm, Carriage, Cart, Caterpillar, CeilingLamp, Centipede, Chair, Clip, Clock, Clothes, CoatHanger, Comb, ConcretePumpTruck, Crack, Crane, Cup, DentalChair, Desk, DeskChair, Diagram, DishRack, DoorHandle, Dragonfish, Dragonfly, Drum, Earphone, Easel, ElectricIron, Excavator, Eyeglasses, Fan, Fence, Fencing, FerrisWheel, FireExtinguisher, Fishing, Flag, FloorLamp, Forklift, GasStation, Gate, Gear, Goal, Golf, GymEquipment, Hammock, Handcart, Handcraft, Handrail, HangGlider, Harp, Harvester, Headset, Helicopter, Helmet, Hook, HorizontalBar, Hydrovalve, IroningTable, Jewelry, Key, KidsPlayground, Kitchenware, Kite, Knife, Ladder, LaundryRack, Lightning, Lobster, Locust, Machine, MachineGun, MagazineRack, Mantis, Medal, MemorialArchway, Microphone, Missile, MobileHolder, Monitor, Mosquito, Motorcycle, MovingTrolley, Mower, MusicPlayer, MusicStand, ObservationTower, Octopus, OilWell, OlympicLogo, OperatingTable, OutdoorFitnessEquipment, Parachute, Pavilion, Piano, Pipe, PlowHarrow, PoleVault, Punchbag, Rack, Racket, Rifle, Ring, Robot, RockClimbing, Rope, Sailboat, Satellite, Scaffold, Scale, Scissor, Scooter, Sculpture, Seadragon, Seahorse, Seal, SewingMachine, Ship, Shoe, ShoppingCart, ShoppingTrolley, Shower, Shrimp, Signboard, Skateboarding, Skeleton, Skiing, Spade, SpeedBoat, Spider, Spoon, Stair, Stand, Stationary, SteeringWheel, Stethoscope, Stool, Stove, StreetLamp, SweetStand, Swing, Sword, TV, Table, TableChair, TableLamp, TableTennis, Tank, Tapeline, Teapot, Telescope, Tent, TobaccoPipe, Toy, Tractor, TrafficLight, TrafficSign, Trampoline, TransmissionTower, Tree, Tricycle, TrimmerCover, Tripod, Trombone, Truck, Trumpet, Tuba, UAV, Umbrella, UnevenBars, UtilityPole, VacuumCleaner, Violin, Wakesurfing, Watch, WaterTower, WateringPot, Well, WellLid, Wheel, Wheelchair, WindTurbine, Windmill, WineGlass, WireWhisk, Yacht'
+class_labels_TR_sorted = _class_labels_TR_sorted.split(', ')
+
+
+class MyData(data.Dataset):
+ def __init__(self, data_root, image_size, is_train=True):
+ self.size_train = image_size
+ self.size_test = image_size
+ self.keep_size = not config.size
+ self.data_size = (config.size, config.size)
+ self.is_train = is_train
+ self.load_all = config.load_all
+ self.device = config.device
+ self.dataset = data_root.replace('\\', '/').split('/')[-1]
+ if self.is_train and config.auxiliary_classification:
+ self.cls_name2id = {_name: _id for _id, _name in enumerate(class_labels_TR_sorted)}
+ self.transform_image = transforms.Compose([
+ transforms.Resize(self.data_size),
+ transforms.ToTensor(),
+ transforms.Normalize([0.485, 0.456, 0.406], [0.229, 0.224, 0.225]),
+ ][self.load_all or self.keep_size:])
+ self.transform_label = transforms.Compose([
+ transforms.Resize(self.data_size),
+ transforms.ToTensor(),
+ ][self.load_all or self.keep_size:])
+ ## 'im' and 'gt' need modifying
+ image_root = os.path.join(data_root, 'im')
+ self.image_paths = [os.path.join(image_root, p) for p in os.listdir(image_root)]
+ self.label_paths = [p.replace('/im/', '/gt/').replace('.jpg', '.png') for p in self.image_paths]
+ if self.load_all:
+ self.images_loaded, self.labels_loaded = [], []
+ self.class_labels_loaded = []
+ # for image_path, label_path in zip(self.image_paths, self.label_paths):
+ for image_path, label_path in tqdm(zip(self.image_paths, self.label_paths), total=len(self.image_paths)):
+ _image = cv2.imread(image_path)
+ _label = cv2.imread(label_path, cv2.IMREAD_GRAYSCALE)
+ if not self.keep_size:
+ _image_rs = cv2.resize(_image, (config.size, config.size), interpolation=cv2.INTER_LINEAR)
+ _label_rs = cv2.resize(_label, (config.size, config.size), interpolation=cv2.INTER_LINEAR)
+ self.images_loaded.append(
+ Image.fromarray(cv2.cvtColor(_image_rs, cv2.COLOR_BGR2RGB)).convert('RGB')
+ )
+ self.labels_loaded.append(
+ Image.fromarray(_label_rs).convert('L')
+ )
+ self.class_labels_loaded.append(
+ self.cls_name2id[label_path.split('/')[-1].split('#')[3]] if self.is_train and config.auxiliary_classification else -1
+ )
+
+
+ def __getitem__(self, index):
+
+ if self.load_all:
+ image = self.images_loaded[index]
+ class_label = self.class_labels_loaded[index] if self.is_train and config.auxiliary_classification else -1
+ else:
+ image = Image.open(self.image_paths[index]).convert('RGB')
+
+ # loading image and label
+ if self.is_train:
+ image, label = preproc(image, image, preproc_methods=config.preproc_methods)
+ # else:
+ # if _label.shape[0] > 2048 or _label.shape[1] > 2048:
+ # _image = cv2.resize(_image, (2048, 2048), interpolation=cv2.INTER_LINEAR)
+ # _label = cv2.resize(_label, (2048, 2048), interpolation=cv2.INTER_LINEAR)
+
+ image, label = self.transform_image(image), self.transform_label(label)
+
+ if self.is_train:
+ return image, label, class_label
+ else:
+ return image, label, self.label_paths[index]
+
+ def __len__(self):
+ return len(self.image_paths)
+
+
+class YouData(data.Dataset):
+ def __init__(self, data_root, image_size, is_train=True):
+ self.size_train = image_size
+ self.size_test = image_size
+ self.keep_size = not config.size
+ self.data_size = (config.size, config.size)
+ self.is_train = is_train
+ self.load_all = config.load_all
+ self.device = config.device
+ self.dataset = data_root.replace('\\', '/').split('/')[-1]
+ if self.is_train and config.auxiliary_classification:
+ self.cls_name2id = {_name: _id for _id, _name in enumerate(class_labels_TR_sorted)}
+ self.transform_image = transforms.Compose([
+ transforms.Resize(self.data_size),
+ transforms.ToTensor(),
+ transforms.Normalize([0.485, 0.456, 0.406], [0.229, 0.224, 0.225]),
+ ][self.load_all or self.keep_size:])
+ ## 'im' and 'gt' need modifying
+ self.image_paths = glob(data_root + "/*")
+ self.img_sizes = []
+ if self.load_all:
+ self.images_loaded, self.labels_loaded = [], []
+ for image_path in tqdm(self.image_paths, total=len(self.image_paths)):
+ _image = cv2.imread(image_path)
+ if not self.keep_size:
+ _image_rs = cv2.resize(_image, (config.size, config.size), interpolation=cv2.INTER_LINEAR)
+ self.images_loaded.append(
+ Image.fromarray(cv2.cvtColor(_image_rs, cv2.COLOR_BGR2RGB)).convert('RGB')
+ )
+ self.img_sizes.append(_image.shape[:2])
+
+
+ def __getitem__(self, index):
+
+ if self.load_all:
+ image = self.images_loaded[index]
+ else:
+ image = Image.open(self.image_paths[index]).convert('RGB')
+
+ # loading image and label
+ if self.is_train:
+ image, _ = preproc(image, image, preproc_methods=config.preproc_methods)
+
+ image = self.transform_image(image)
+ size = self.img_sizes[index]
+ return image, size
+
+ def __len__(self):
+ return len(self.image_paths)
diff --git a/py/BiRefNet_legacy/modules/__init__.py b/py/BiRefNet_legacy/modules/__init__.py
new file mode 100644
index 0000000..e69de29
diff --git a/py/BiRefNet_legacy/modules/aspp.py b/py/BiRefNet_legacy/modules/aspp.py
new file mode 100644
index 0000000..a34d69a
--- /dev/null
+++ b/py/BiRefNet_legacy/modules/aspp.py
@@ -0,0 +1,162 @@
+import torch
+import torch.nn as nn
+import torch.nn.functional as F
+from BiRefNet_legacy.modules.deform_conv import DeformableConv2d
+from ..config import Config
+
+
+config = Config()
+
+
+class ASPPComplex(nn.Module):
+ def __init__(self, in_channels=64, out_channels=None, output_stride=16):
+ super(ASPPComplex, self).__init__()
+ self.down_scale = 1
+ if out_channels is None:
+ out_channels = in_channels
+ self.in_channelster = 256 // self.down_scale
+ if output_stride == 16:
+ dilations = [1, 6, 12, 18]
+ elif output_stride == 8:
+ dilations = [1, 12, 24, 36]
+ else:
+ raise NotImplementedError
+
+ self.aspp1 = _ASPPModule(in_channels, self.in_channelster, 1, padding=0, dilation=dilations[0])
+ self.aspp2 = _ASPPModule(in_channels, self.in_channelster, 3, padding=dilations[1], dilation=dilations[1])
+ self.aspp3 = _ASPPModule(in_channels, self.in_channelster, 3, padding=dilations[2], dilation=dilations[2])
+ self.aspp4 = _ASPPModule(in_channels, self.in_channelster, 3, padding=dilations[3], dilation=dilations[3])
+
+ self.global_avg_pool = nn.Sequential(nn.AdaptiveAvgPool2d((1, 1)),
+ nn.Conv2d(in_channels, self.in_channelster, 1, stride=1, bias=False),
+ nn.BatchNorm2d(self.in_channelster) if config.batch_size > 1 else nn.Identity(),
+ nn.ReLU(inplace=True))
+ self.conv1 = nn.Conv2d(self.in_channelster * 5, out_channels, 1, bias=False)
+ self.bn1 = nn.BatchNorm2d(out_channels)
+ self.relu = nn.ReLU(inplace=True)
+ self.dropout = nn.Dropout(0.5)
+
+ def forward(self, x):
+ x1 = self.aspp1(x)
+ x2 = self.aspp2(x)
+ x3 = self.aspp3(x)
+ x4 = self.aspp4(x)
+ x5 = self.global_avg_pool(x)
+ x5 = F.interpolate(x5, size=x1.size()[2:], mode='bilinear', align_corners=True)
+ x = torch.cat((x1, x2, x3, x4, x5), dim=1)
+
+ x = self.conv1(x)
+ x = self.bn1(x)
+ x = self.relu(x)
+
+ return self.dropout(x)
+
+
+class _ASPPModule(nn.Module):
+ def __init__(self, in_channels, planes, kernel_size, padding, dilation):
+ super(_ASPPModule, self).__init__()
+ self.atrous_conv = nn.Conv2d(in_channels, planes, kernel_size=kernel_size,
+ stride=1, padding=padding, dilation=dilation, bias=False)
+ self.bn = nn.BatchNorm2d(planes)
+ self.relu = nn.ReLU(inplace=True)
+
+ def forward(self, x):
+ x = self.atrous_conv(x)
+ x = self.bn(x)
+
+ return self.relu(x)
+
+class ASPP(nn.Module):
+ def __init__(self, in_channels=64, out_channels=None, output_stride=16):
+ super(ASPP, self).__init__()
+ self.down_scale = 1
+ if out_channels is None:
+ out_channels = in_channels
+ self.in_channelster = 256 // self.down_scale
+ if output_stride == 16:
+ dilations = [1, 6, 12, 18]
+ elif output_stride == 8:
+ dilations = [1, 12, 24, 36]
+ else:
+ raise NotImplementedError
+
+ self.aspp1 = _ASPPModule(in_channels, self.in_channelster, 1, padding=0, dilation=dilations[0])
+ self.aspp2 = _ASPPModule(in_channels, self.in_channelster, 3, padding=dilations[1], dilation=dilations[1])
+ self.aspp3 = _ASPPModule(in_channels, self.in_channelster, 3, padding=dilations[2], dilation=dilations[2])
+ self.aspp4 = _ASPPModule(in_channels, self.in_channelster, 3, padding=dilations[3], dilation=dilations[3])
+
+ self.global_avg_pool = nn.Sequential(nn.AdaptiveAvgPool2d((1, 1)),
+ nn.Conv2d(in_channels, self.in_channelster, 1, stride=1, bias=False),
+ nn.BatchNorm2d(self.in_channelster) if config.batch_size > 1 else nn.Identity(),
+ nn.ReLU(inplace=True))
+ self.conv1 = nn.Conv2d(self.in_channelster * 5, out_channels, 1, bias=False)
+ self.bn1 = nn.BatchNorm2d(out_channels)
+ self.relu = nn.ReLU(inplace=True)
+ self.dropout = nn.Dropout(0.5)
+
+ def forward(self, x):
+ x1 = self.aspp1(x)
+ x2 = self.aspp2(x)
+ x3 = self.aspp3(x)
+ x4 = self.aspp4(x)
+ x5 = self.global_avg_pool(x)
+ x5 = F.interpolate(x5, size=x1.size()[2:], mode='bilinear', align_corners=True)
+ x = torch.cat((x1, x2, x3, x4, x5), dim=1)
+
+ x = self.conv1(x)
+ x = self.bn1(x)
+ x = self.relu(x)
+
+ return self.dropout(x)
+
+
+##################### Deformable
+class _ASPPModuleDeformable(nn.Module):
+ def __init__(self, in_channels, planes, kernel_size, padding):
+ super(_ASPPModuleDeformable, self).__init__()
+ self.atrous_conv = DeformableConv2d(in_channels, planes, kernel_size=kernel_size,
+ stride=1, padding=padding, bias=False)
+ self.bn = nn.BatchNorm2d(planes)
+ self.relu = nn.ReLU(inplace=True)
+
+ def forward(self, x):
+ x = self.atrous_conv(x)
+ x = self.bn(x)
+
+ return self.relu(x)
+
+
+class ASPPDeformable(nn.Module):
+ def __init__(self, in_channels, out_channels=None, num_parallel_block=1):
+ super(ASPPDeformable, self).__init__()
+ self.down_scale = 1
+ if out_channels is None:
+ out_channels = in_channels
+ self.in_channelster = 256 // self.down_scale
+
+ self.aspp1 = _ASPPModuleDeformable(in_channels, self.in_channelster, 1, padding=0)
+ self.aspp_deforms = nn.ModuleList([
+ _ASPPModuleDeformable(in_channels, self.in_channelster, 3, padding=1) for _ in range(num_parallel_block)
+ ])
+
+ self.global_avg_pool = nn.Sequential(nn.AdaptiveAvgPool2d((1, 1)),
+ nn.Conv2d(in_channels, self.in_channelster, 1, stride=1, bias=False),
+ nn.BatchNorm2d(self.in_channelster) if config.batch_size > 1 else nn.Identity(),
+ nn.ReLU(inplace=True))
+ self.conv1 = nn.Conv2d(self.in_channelster * (2 + len(self.aspp_deforms)), out_channels, 1, bias=False)
+ self.bn1 = nn.BatchNorm2d(out_channels)
+ self.relu = nn.ReLU(inplace=True)
+ self.dropout = nn.Dropout(0.5)
+
+ def forward(self, x):
+ x1 = self.aspp1(x)
+ x_aspp_deforms = [aspp_deform(x) for aspp_deform in self.aspp_deforms]
+ x5 = self.global_avg_pool(x)
+ x5 = F.interpolate(x5, size=x1.size()[2:], mode='bilinear', align_corners=True)
+ x = torch.cat((x1, *x_aspp_deforms, x5), dim=1)
+
+ x = self.conv1(x)
+ x = self.bn1(x)
+ x = self.relu(x)
+
+ return self.dropout(x)
diff --git a/py/BiRefNet_legacy/modules/attentions.py b/py/BiRefNet_legacy/modules/attentions.py
new file mode 100644
index 0000000..e1032af
--- /dev/null
+++ b/py/BiRefNet_legacy/modules/attentions.py
@@ -0,0 +1,93 @@
+import numpy as np
+import torch
+from torch import nn
+from torch.nn import init
+
+
+class SEWeightModule(nn.Module):
+ def __init__(self, channels, reduction=16):
+ super(SEWeightModule, self).__init__()
+ self.avg_pool = nn.AdaptiveAvgPool2d(1)
+ self.fc1 = nn.Conv2d(channels, channels//reduction, kernel_size=1, padding=0)
+ self.relu = nn.ReLU(inplace=True)
+ self.fc2 = nn.Conv2d(channels//reduction, channels, kernel_size=1, padding=0)
+ self.sigmoid = nn.Sigmoid()
+
+ def forward(self, x):
+ out = self.avg_pool(x)
+ out = self.fc1(out)
+ out = self.relu(out)
+ out = self.fc2(out)
+ weight = self.sigmoid(out)
+ return weight
+
+
+class PSA(nn.Module):
+
+ def __init__(self, in_channels, S=4, reduction=4):
+ super().__init__()
+ self.S = S
+
+ _convs = []
+ for i in range(S):
+ _convs.append(nn.Conv2d(in_channels//S, in_channels//S, kernel_size=2*(i+1)+1, padding=i+1))
+ self.convs = nn.ModuleList(_convs)
+
+ self.se_block = SEWeightModule(in_channels//S, reduction=S*reduction)
+
+ self.softmax = nn.Softmax(dim=1)
+
+ def forward(self, x):
+ b, c, h, w = x.size()
+
+ # Step1: SPC module
+ SPC_out = x.view(b, self.S, c//self.S, h, w) #bs,s,ci,h,w
+ for idx, conv in enumerate(self.convs):
+ SPC_out[:,idx,:,:,:] = conv(SPC_out[:,idx,:,:,:].clone())
+
+ # Step2: SE weight
+ se_out=[]
+ for idx in range(self.S):
+ se_out.append(self.se_block(SPC_out[:, idx, :, :, :]))
+ SE_out = torch.stack(se_out, dim=1)
+ SE_out = SE_out.expand_as(SPC_out)
+
+ # Step3: Softmax
+ softmax_out = self.softmax(SE_out)
+
+ # Step4: SPA
+ PSA_out = SPC_out * softmax_out
+ PSA_out = PSA_out.view(b, -1, h, w)
+
+ return PSA_out
+
+
+class SGE(nn.Module):
+
+ def __init__(self, groups):
+ super().__init__()
+ self.groups=groups
+ self.avg_pool = nn.AdaptiveAvgPool2d(1)
+ self.weight=nn.Parameter(torch.zeros(1,groups,1,1))
+ self.bias=nn.Parameter(torch.zeros(1,groups,1,1))
+ self.sig=nn.Sigmoid()
+
+ def forward(self, x):
+ b, c, h,w=x.shape
+ x=x.view(b*self.groups,-1,h,w) #bs*g,dim//g,h,w
+ xn=x*self.avg_pool(x) #bs*g,dim//g,h,w
+ xn=xn.sum(dim=1,keepdim=True) #bs*g,1,h,w
+ t=xn.view(b*self.groups,-1) #bs*g,h*w
+
+ t=t-t.mean(dim=1,keepdim=True) #bs*g,h*w
+ std=t.std(dim=1,keepdim=True)+1e-5
+ t=t/std #bs*g,h*w
+ t=t.view(b,self.groups,h,w) #bs,g,h*w
+
+ t=t*self.weight+self.bias #bs,g,h*w
+ t=t.view(b*self.groups,1,h,w) #bs*g,1,h*w
+ x=x*self.sig(t)
+ x=x.view(b,c,h,w)
+
+ return x
+
diff --git a/py/BiRefNet_legacy/modules/decoder_blocks.py b/py/BiRefNet_legacy/modules/decoder_blocks.py
new file mode 100644
index 0000000..84c9592
--- /dev/null
+++ b/py/BiRefNet_legacy/modules/decoder_blocks.py
@@ -0,0 +1,101 @@
+import torch
+import torch.nn as nn
+from BiRefNet_legacy.modules.aspp import ASPP, ASPPDeformable
+from BiRefNet_legacy.modules.attentions import PSA, SGE
+from ..config import Config
+
+
+config = Config()
+
+
+class BasicDecBlk(nn.Module):
+ def __init__(self, in_channels=64, out_channels=64, inter_channels=64):
+ super(BasicDecBlk, self).__init__()
+ inter_channels = in_channels // 4 if config.dec_channels_inter == 'adap' else 64
+ self.conv_in = nn.Conv2d(in_channels, inter_channels, 3, 1, padding=1)
+ self.relu_in = nn.ReLU(inplace=True)
+ if config.dec_att == 'ASPP':
+ self.dec_att = ASPP(in_channels=inter_channels)
+ elif config.dec_att == 'ASPPDeformable':
+ self.dec_att = ASPPDeformable(in_channels=inter_channels)
+ self.conv_out = nn.Conv2d(inter_channels, out_channels, 3, 1, padding=1)
+ self.bn_in = nn.BatchNorm2d(inter_channels)
+ self.bn_out = nn.BatchNorm2d(out_channels)
+
+ def forward(self, x):
+ x = self.conv_in(x)
+ x = self.bn_in(x)
+ x = self.relu_in(x)
+ if hasattr(self, 'dec_att'):
+ x = self.dec_att(x)
+ x = self.conv_out(x)
+ x = self.bn_out(x)
+ return x
+
+
+class ResBlk(nn.Module):
+ def __init__(self, in_channels=64, out_channels=None, inter_channels=64):
+ super(ResBlk, self).__init__()
+ if out_channels is None:
+ out_channels = in_channels
+ inter_channels = in_channels // 4 if config.dec_channels_inter == 'adap' else 64
+
+ self.conv_in = nn.Conv2d(in_channels, inter_channels, 3, 1, padding=1)
+ self.bn_in = nn.BatchNorm2d(inter_channels)
+ self.relu_in = nn.ReLU(inplace=True)
+
+ if config.dec_att == 'ASPP':
+ self.dec_att = ASPP(in_channels=inter_channels)
+ elif config.dec_att == 'ASPPDeformable':
+ self.dec_att = ASPPDeformable(in_channels=inter_channels)
+
+ self.conv_out = nn.Conv2d(inter_channels, out_channels, 3, 1, padding=1)
+ self.bn_out = nn.BatchNorm2d(out_channels)
+
+ self.conv_resi = nn.Conv2d(in_channels, out_channels, 1, 1, 0)
+
+ def forward(self, x):
+ _x = self.conv_resi(x)
+ x = self.conv_in(x)
+ x = self.bn_in(x)
+ x = self.relu_in(x)
+ if hasattr(self, 'dec_att'):
+ x = self.dec_att(x)
+ x = self.conv_out(x)
+ x = self.bn_out(x)
+ return x + _x
+
+
+class HierarAttDecBlk(nn.Module):
+ def __init__(self, in_channels=64, out_channels=None, inter_channels=64):
+ super(HierarAttDecBlk, self).__init__()
+ if out_channels is None:
+ out_channels = in_channels
+ inter_channels = in_channels // 4 if config.dec_channels_inter == 'adap' else 64
+ self.split_y = 8 # must be divided by channels of all intermediate features
+ self.split_x = 8
+
+ self.conv_in = nn.Conv2d(in_channels, inter_channels, 3, 1, 1)
+
+ self.psa = PSA(inter_channels*self.split_y*self.split_x, S=config.batch_size)
+ self.sge = SGE(groups=config.batch_size)
+
+ if config.dec_att == 'ASPP':
+ self.dec_att = ASPP(in_channels=inter_channels)
+ elif config.dec_att == 'ASPPDeformable':
+ self.dec_att = ASPPDeformable(in_channels=inter_channels)
+ self.conv_out = nn.Conv2d(inter_channels, out_channels, 3, 1, 1)
+
+ def forward(self, x):
+ x = self.conv_in(x)
+ N, C, H, W = x.shape
+ x_patchs = x.reshape(N, -1, H//self.split_y, W//self.split_x)
+
+ # Hierarchical attention: group attention X patch spatial attention
+ x_patchs = self.psa(x_patchs) # Group Channel Attention -- each group is a single image
+ x_patchs = self.sge(x_patchs) # Patch Spatial Attention
+ x = x.reshape(N, C, H, W)
+ if hasattr(self, 'dec_att'):
+ x = self.dec_att(x)
+ x = self.conv_out(x)
+ return x
diff --git a/py/BiRefNet_legacy/modules/deform_conv.py b/py/BiRefNet_legacy/modules/deform_conv.py
new file mode 100644
index 0000000..43f5e57
--- /dev/null
+++ b/py/BiRefNet_legacy/modules/deform_conv.py
@@ -0,0 +1,66 @@
+import torch
+import torch.nn as nn
+from torchvision.ops import deform_conv2d
+
+
+class DeformableConv2d(nn.Module):
+ def __init__(self,
+ in_channels,
+ out_channels,
+ kernel_size=3,
+ stride=1,
+ padding=1,
+ bias=False):
+
+ super(DeformableConv2d, self).__init__()
+
+ assert type(kernel_size) == tuple or type(kernel_size) == int
+
+ kernel_size = kernel_size if type(kernel_size) == tuple else (kernel_size, kernel_size)
+ self.stride = stride if type(stride) == tuple else (stride, stride)
+ self.padding = padding
+
+ self.offset_conv = nn.Conv2d(in_channels,
+ 2 * kernel_size[0] * kernel_size[1],
+ kernel_size=kernel_size,
+ stride=stride,
+ padding=self.padding,
+ bias=True)
+
+ nn.init.constant_(self.offset_conv.weight, 0.)
+ nn.init.constant_(self.offset_conv.bias, 0.)
+
+ self.modulator_conv = nn.Conv2d(in_channels,
+ 1 * kernel_size[0] * kernel_size[1],
+ kernel_size=kernel_size,
+ stride=stride,
+ padding=self.padding,
+ bias=True)
+
+ nn.init.constant_(self.modulator_conv.weight, 0.)
+ nn.init.constant_(self.modulator_conv.bias, 0.)
+
+ self.regular_conv = nn.Conv2d(in_channels,
+ out_channels=out_channels,
+ kernel_size=kernel_size,
+ stride=stride,
+ padding=self.padding,
+ bias=bias)
+
+ def forward(self, x):
+ #h, w = x.shape[2:]
+ #max_offset = max(h, w)/4.
+
+ offset = self.offset_conv(x)#.clamp(-max_offset, max_offset)
+ modulator = 2. * torch.sigmoid(self.modulator_conv(x))
+
+ x = deform_conv2d(
+ input=x,
+ offset=offset,
+ weight=self.regular_conv.weight,
+ bias=self.regular_conv.bias,
+ padding=self.padding,
+ mask=modulator,
+ stride=self.stride,
+ )
+ return x
diff --git a/py/BiRefNet_legacy/modules/ing.py b/py/BiRefNet_legacy/modules/ing.py
new file mode 100644
index 0000000..3032075
--- /dev/null
+++ b/py/BiRefNet_legacy/modules/ing.py
@@ -0,0 +1,29 @@
+import torch.nn as nn
+from BiRefNet_legacy.modules.mlp import MLPLayer
+
+
+class BlockA(nn.Module):
+ def __init__(self, in_channels=64, out_channels=64, inter_channels=64, mlp_ratio=4.):
+ super(BlockA, self).__init__()
+ inter_channels = in_channels
+ self.conv = nn.Conv2d(in_channels, inter_channels, 3, 1, 1)
+ self.norm1 = nn.LayerNorm(inter_channels)
+ self.ffn = MLPLayer(in_features=inter_channels,
+ hidden_features=int(inter_channels * mlp_ratio),
+ act_layer=nn.GELU,
+ drop=0.)
+ self.norm2 = nn.LayerNorm(inter_channels)
+
+ def forward(self, x):
+ B, C, H, W = x.shape
+ _x = self.conv(x)
+ _x = _x.flatten(2).transpose(1, 2)
+ _x = self.norm1(_x)
+ x = x.reshape(B, H, W, -1).permute(0, 3, 1, 2).contiguous()
+
+ x = x + _x
+ _x1 = self.ffn(x)
+ _x1 = self.norm2(_x1)
+ _x1 = _x1.reshape(B, H, W, -1).permute(0, 3, 1, 2).contiguous()
+ x = x + _x1
+ return x
\ No newline at end of file
diff --git a/py/BiRefNet_legacy/modules/lateral_blocks.py b/py/BiRefNet_legacy/modules/lateral_blocks.py
new file mode 100644
index 0000000..a3022e6
--- /dev/null
+++ b/py/BiRefNet_legacy/modules/lateral_blocks.py
@@ -0,0 +1,21 @@
+import numpy as np
+import torch
+import torch.nn as nn
+import torch.nn.functional as F
+from functools import partial
+
+from ..config import Config
+
+
+config = Config()
+
+
+class BasicLatBlk(nn.Module):
+ def __init__(self, in_channels=64, out_channels=64, inter_channels=64):
+ super(BasicLatBlk, self).__init__()
+ inter_channels = in_channels // 4 if config.dec_channels_inter == 'adap' else 64
+ self.conv = nn.Conv2d(in_channels, out_channels, 1, 1, 0)
+
+ def forward(self, x):
+ x = self.conv(x)
+ return x
diff --git a/py/BiRefNet_legacy/modules/mlp.py b/py/BiRefNet_legacy/modules/mlp.py
new file mode 100644
index 0000000..a383459
--- /dev/null
+++ b/py/BiRefNet_legacy/modules/mlp.py
@@ -0,0 +1,118 @@
+import torch
+import torch.nn as nn
+from functools import partial
+
+from timm.models.layers import DropPath, to_2tuple, trunc_normal_
+from timm.models import register_model
+
+import math
+
+
+class MLPLayer(nn.Module):
+ def __init__(self, in_features, hidden_features=None, out_features=None, act_layer=nn.GELU, drop=0.):
+ super().__init__()
+ out_features = out_features or in_features
+ hidden_features = hidden_features or in_features
+ self.fc1 = nn.Linear(in_features, hidden_features)
+ self.act = act_layer()
+ self.fc2 = nn.Linear(hidden_features, out_features)
+ self.drop = nn.Dropout(drop)
+
+ def forward(self, x):
+ x = self.fc1(x)
+ x = self.act(x)
+ x = self.drop(x)
+ x = self.fc2(x)
+ x = self.drop(x)
+ return x
+
+
+class Attention(nn.Module):
+ def __init__(self, dim, num_heads=8, qkv_bias=False, qk_scale=None, attn_drop=0., proj_drop=0., sr_ratio=1):
+ super().__init__()
+ assert dim % num_heads == 0, f"dim {dim} should be divided by num_heads {num_heads}."
+
+ self.dim = dim
+ self.num_heads = num_heads
+ head_dim = dim // num_heads
+ self.scale = qk_scale or head_dim ** -0.5
+
+ self.q = nn.Linear(dim, dim, bias=qkv_bias)
+ self.kv = nn.Linear(dim, dim * 2, bias=qkv_bias)
+ self.attn_drop = nn.Dropout(attn_drop)
+ self.proj = nn.Linear(dim, dim)
+ self.proj_drop = nn.Dropout(proj_drop)
+
+ self.sr_ratio = sr_ratio
+ if sr_ratio > 1:
+ self.sr = nn.Conv2d(dim, dim, kernel_size=sr_ratio, stride=sr_ratio)
+ self.norm = nn.LayerNorm(dim)
+
+ def forward(self, x, H, W):
+ B, N, C = x.shape
+ q = self.q(x).reshape(B, N, self.num_heads, C // self.num_heads).permute(0, 2, 1, 3)
+
+ if self.sr_ratio > 1:
+ x_ = x.permute(0, 2, 1).reshape(B, C, H, W)
+ x_ = self.sr(x_).reshape(B, C, -1).permute(0, 2, 1)
+ x_ = self.norm(x_)
+ kv = self.kv(x_).reshape(B, -1, 2, self.num_heads, C // self.num_heads).permute(2, 0, 3, 1, 4)
+ else:
+ kv = self.kv(x).reshape(B, -1, 2, self.num_heads, C // self.num_heads).permute(2, 0, 3, 1, 4)
+ k, v = kv[0], kv[1]
+
+ attn = (q @ k.transpose(-2, -1)) * self.scale
+ attn = attn.softmax(dim=-1)
+ attn = self.attn_drop(attn)
+
+ x = (attn @ v).transpose(1, 2).reshape(B, N, C)
+ x = self.proj(x)
+ x = self.proj_drop(x)
+ return x
+
+
+class Block(nn.Module):
+ def __init__(self, dim, num_heads, mlp_ratio=4., qkv_bias=False, qk_scale=None, drop=0., attn_drop=0.,
+ drop_path=0., act_layer=nn.GELU, norm_layer=nn.LayerNorm, sr_ratio=1):
+ super().__init__()
+ self.norm1 = norm_layer(dim)
+ self.attn = Attention(
+ dim,
+ num_heads=num_heads, qkv_bias=qkv_bias, qk_scale=qk_scale,
+ attn_drop=attn_drop, proj_drop=drop, sr_ratio=sr_ratio)
+ # NOTE: drop path for stochastic depth, we shall see if this is better than dropout here
+ self.drop_path = DropPath(drop_path) if drop_path > 0. else nn.Identity()
+ self.norm2 = norm_layer(dim)
+ mlp_hidden_dim = int(dim * mlp_ratio)
+ self.mlp = MLPLayer(in_features=dim, hidden_features=mlp_hidden_dim, act_layer=act_layer, drop=drop)
+
+ def forward(self, x, H, W):
+ x = x + self.drop_path(self.attn(self.norm1(x), H, W))
+ x = x + self.drop_path(self.mlp(self.norm2(x), H, W))
+ return x
+
+
+class OverlapPatchEmbed(nn.Module):
+ """ Image to Patch Embedding
+ """
+
+ def __init__(self, img_size=224, patch_size=7, stride=4, in_channels=3, embed_dim=768):
+ super().__init__()
+ img_size = to_2tuple(img_size)
+ patch_size = to_2tuple(patch_size)
+
+ self.img_size = img_size
+ self.patch_size = patch_size
+ self.H, self.W = img_size[0] // patch_size[0], img_size[1] // patch_size[1]
+ self.num_patches = self.H * self.W
+ self.proj = nn.Conv2d(in_channels, embed_dim, kernel_size=patch_size, stride=stride,
+ padding=(patch_size[0] // 2, patch_size[1] // 2))
+ self.norm = nn.LayerNorm(embed_dim)
+
+ def forward(self, x):
+ x = self.proj(x)
+ _, _, H, W = x.shape
+ x = x.flatten(2).transpose(1, 2)
+ x = self.norm(x)
+ return x, H, W
+
diff --git a/py/BiRefNet_legacy/modules/utils.py b/py/BiRefNet_legacy/modules/utils.py
new file mode 100644
index 0000000..59bd912
--- /dev/null
+++ b/py/BiRefNet_legacy/modules/utils.py
@@ -0,0 +1,54 @@
+import torch.nn as nn
+
+
+def build_act_layer(act_layer):
+ if act_layer == 'ReLU':
+ return nn.ReLU(inplace=True)
+ elif act_layer == 'SiLU':
+ return nn.SiLU(inplace=True)
+ elif act_layer == 'GELU':
+ return nn.GELU()
+
+ raise NotImplementedError(f'build_act_layer does not support {act_layer}')
+
+
+def build_norm_layer(dim,
+ norm_layer,
+ in_format='channels_last',
+ out_format='channels_last',
+ eps=1e-6):
+ layers = []
+ if norm_layer == 'BN':
+ if in_format == 'channels_last':
+ layers.append(to_channels_first())
+ layers.append(nn.BatchNorm2d(dim))
+ if out_format == 'channels_last':
+ layers.append(to_channels_last())
+ elif norm_layer == 'LN':
+ if in_format == 'channels_first':
+ layers.append(to_channels_last())
+ layers.append(nn.LayerNorm(dim, eps=eps))
+ if out_format == 'channels_first':
+ layers.append(to_channels_first())
+ else:
+ raise NotImplementedError(
+ f'build_norm_layer does not support {norm_layer}')
+ return nn.Sequential(*layers)
+
+
+class to_channels_first(nn.Module):
+
+ def __init__(self):
+ super().__init__()
+
+ def forward(self, x):
+ return x.permute(0, 3, 1, 2)
+
+
+class to_channels_last(nn.Module):
+
+ def __init__(self):
+ super().__init__()
+
+ def forward(self, x):
+ return x.permute(0, 2, 3, 1)
diff --git a/py/BiRefNet_legacy/preproc.py b/py/BiRefNet_legacy/preproc.py
new file mode 100644
index 0000000..a059c5d
--- /dev/null
+++ b/py/BiRefNet_legacy/preproc.py
@@ -0,0 +1,85 @@
+from PIL import Image, ImageEnhance
+import random
+import numpy as np
+import random
+
+
+def preproc(image, label, preproc_methods=['flip']):
+ if 'flip' in preproc_methods:
+ image, label = cv_random_flip(image, label)
+ if 'crop' in preproc_methods:
+ image, label = random_crop(image, label)
+ if 'rotate' in preproc_methods:
+ image, label = random_rotate(image, label)
+ if 'enhance' in preproc_methods:
+ image = color_enhance(image)
+ if 'pepper' in preproc_methods:
+ label = random_pepper(label)
+ return image, label
+
+
+def cv_random_flip(img, label):
+ if random.random() > 0.5:
+ img = img.transpose(Image.FLIP_LEFT_RIGHT)
+ label = label.transpose(Image.FLIP_LEFT_RIGHT)
+ return img, label
+
+
+def random_crop(image, label):
+ border = 30
+ image_width = image.size[0]
+ image_height = image.size[1]
+ border = int(min(image_width, image_height) * 0.1)
+ crop_win_width = np.random.randint(image_width - border, image_width)
+ crop_win_height = np.random.randint(image_height - border, image_height)
+ random_region = (
+ (image_width - crop_win_width) >> 1, (image_height - crop_win_height) >> 1, (image_width + crop_win_width) >> 1,
+ (image_height + crop_win_height) >> 1)
+ return image.crop(random_region), label.crop(random_region)
+
+
+def random_rotate(image, label, angle=15):
+ mode = Image.BICUBIC
+ if random.random() > 0.8:
+ random_angle = np.random.randint(-angle, angle)
+ image = image.rotate(random_angle, mode)
+ label = label.rotate(random_angle, mode)
+ return image, label
+
+
+def color_enhance(image):
+ bright_intensity = random.randint(5, 15) / 10.0
+ image = ImageEnhance.Brightness(image).enhance(bright_intensity)
+ contrast_intensity = random.randint(5, 15) / 10.0
+ image = ImageEnhance.Contrast(image).enhance(contrast_intensity)
+ color_intensity = random.randint(0, 20) / 10.0
+ image = ImageEnhance.Color(image).enhance(color_intensity)
+ sharp_intensity = random.randint(0, 30) / 10.0
+ image = ImageEnhance.Sharpness(image).enhance(sharp_intensity)
+ return image
+
+
+def random_gaussian(image, mean=0.1, sigma=0.35):
+ def gaussianNoisy(im, mean=mean, sigma=sigma):
+ for _i in range(len(im)):
+ im[_i] += random.gauss(mean, sigma)
+ return im
+
+ img = np.asarray(image)
+ width, height = img.shape
+ img = gaussianNoisy(img[:].flatten(), mean, sigma)
+ img = img.reshape([width, height])
+ return Image.fromarray(np.uint8(img))
+
+
+def random_pepper(img, N=0.0015):
+ img = np.array(img)
+ noiseNum = int(N * img.shape[0] * img.shape[1])
+ for i in range(noiseNum):
+ randX = random.randint(0, img.shape[0] - 1)
+ randY = random.randint(0, img.shape[1] - 1)
+ if random.randint(0, 1) == 0:
+ img[randX, randY] = 0
+ else:
+ img[randX, randY] = 255
+ return Image.fromarray(img)
diff --git a/py/BiRefNet_legacy/refinement/__init__.py b/py/BiRefNet_legacy/refinement/__init__.py
new file mode 100644
index 0000000..e69de29
diff --git a/py/BiRefNet_legacy/refinement/refiner.py b/py/BiRefNet_legacy/refinement/refiner.py
new file mode 100644
index 0000000..71567a6
--- /dev/null
+++ b/py/BiRefNet_legacy/refinement/refiner.py
@@ -0,0 +1,253 @@
+import torch
+import torch.nn as nn
+from collections import OrderedDict
+import torch
+import torch.nn as nn
+import torch.nn.functional as F
+from torchvision.models import vgg16, vgg16_bn
+from torchvision.models import resnet50
+
+from ..config import Config
+from ..dataset import class_labels_TR_sorted
+from BiRefNet_legacy.backbones.build_backbone import build_backbone
+from BiRefNet_legacy.modules.decoder_blocks import BasicDecBlk
+from BiRefNet_legacy.modules.lateral_blocks import BasicLatBlk
+from BiRefNet_legacy.modules.ing import *
+from BiRefNet_legacy.refinement.stem_layer import StemLayer
+
+
+class RefinerPVTInChannels4(nn.Module):
+ def __init__(self, in_channels=3+1):
+ super(RefinerPVTInChannels4, self).__init__()
+ self.config = Config()
+ self.epoch = 1
+ self.bb = build_backbone(self.config.bb, params_settings='in_channels=4')
+
+ lateral_channels_in_collection = {
+ 'vgg16': [512, 256, 128, 64], 'vgg16bn': [512, 256, 128, 64], 'resnet50': [1024, 512, 256, 64],
+ 'pvt_v2_b2': [512, 320, 128, 64], 'pvt_v2_b5': [512, 320, 128, 64],
+ 'swin_v1_b': [1024, 512, 256, 128], 'swin_v1_l': [1536, 768, 384, 192],
+ }
+ channels = lateral_channels_in_collection[self.config.bb]
+ self.squeeze_module = BasicDecBlk(channels[0], channels[0])
+
+ self.decoder = Decoder(channels)
+
+ if 0:
+ for key, value in self.named_parameters():
+ if 'bb.' in key:
+ value.requires_grad = False
+
+ def forward(self, x):
+ if isinstance(x, list):
+ x = torch.cat(x, dim=1)
+ ########## Encoder ##########
+ if self.config.bb in ['vgg16', 'vgg16bn', 'resnet50']:
+ x1 = self.bb.conv1(x)
+ x2 = self.bb.conv2(x1)
+ x3 = self.bb.conv3(x2)
+ x4 = self.bb.conv4(x3)
+ else:
+ x1, x2, x3, x4 = self.bb(x)
+
+ x4 = self.squeeze_module(x4)
+
+ ########## Decoder ##########
+
+ features = [x, x1, x2, x3, x4]
+ scaled_preds = self.decoder(features)
+
+ return scaled_preds
+
+
+class Refiner(nn.Module):
+ def __init__(self, in_channels=3+1):
+ super(Refiner, self).__init__()
+ self.config = Config()
+ self.epoch = 1
+ self.stem_layer = StemLayer(in_channels=in_channels, inter_channels=48, out_channels=3)
+ self.bb = build_backbone(self.config.bb)
+
+ lateral_channels_in_collection = {
+ 'vgg16': [512, 256, 128, 64], 'vgg16bn': [512, 256, 128, 64], 'resnet50': [1024, 512, 256, 64],
+ 'pvt_v2_b2': [512, 320, 128, 64], 'pvt_v2_b5': [512, 320, 128, 64],
+ 'swin_v1_b': [1024, 512, 256, 128], 'swin_v1_l': [1536, 768, 384, 192],
+ }
+ channels = lateral_channels_in_collection[self.config.bb]
+ self.squeeze_module = BasicDecBlk(channels[0], channels[0])
+
+ self.decoder = Decoder(channels)
+
+ if 0:
+ for key, value in self.named_parameters():
+ if 'bb.' in key:
+ value.requires_grad = False
+
+ def forward(self, x):
+ if isinstance(x, list):
+ x = torch.cat(x, dim=1)
+ x = self.stem_layer(x)
+ ########## Encoder ##########
+ if self.config.bb in ['vgg16', 'vgg16bn', 'resnet50']:
+ x1 = self.bb.conv1(x)
+ x2 = self.bb.conv2(x1)
+ x3 = self.bb.conv3(x2)
+ x4 = self.bb.conv4(x3)
+ else:
+ x1, x2, x3, x4 = self.bb(x)
+
+ x4 = self.squeeze_module(x4)
+
+ ########## Decoder ##########
+
+ features = [x, x1, x2, x3, x4]
+ scaled_preds = self.decoder(features)
+
+ return scaled_preds
+
+
+class Decoder(nn.Module):
+ def __init__(self, channels):
+ super(Decoder, self).__init__()
+ self.config = Config()
+ DecoderBlock = eval('BasicDecBlk')
+ LateralBlock = eval('BasicLatBlk')
+
+ self.decoder_block4 = DecoderBlock(channels[0], channels[1])
+ self.decoder_block3 = DecoderBlock(channels[1], channels[2])
+ self.decoder_block2 = DecoderBlock(channels[2], channels[3])
+ self.decoder_block1 = DecoderBlock(channels[3], channels[3]//2)
+
+ self.lateral_block4 = LateralBlock(channels[1], channels[1])
+ self.lateral_block3 = LateralBlock(channels[2], channels[2])
+ self.lateral_block2 = LateralBlock(channels[3], channels[3])
+
+ if self.config.ms_supervision:
+ self.conv_ms_spvn_4 = nn.Conv2d(channels[1], 1, 1, 1, 0)
+ self.conv_ms_spvn_3 = nn.Conv2d(channels[2], 1, 1, 1, 0)
+ self.conv_ms_spvn_2 = nn.Conv2d(channels[3], 1, 1, 1, 0)
+ self.conv_out1 = nn.Sequential(nn.Conv2d(channels[3]//2, 1, 1, 1, 0))
+
+ def forward(self, features):
+ x, x1, x2, x3, x4 = features
+ outs = []
+ p4 = self.decoder_block4(x4)
+ _p4 = F.interpolate(p4, size=x3.shape[2:], mode='bilinear', align_corners=True)
+ _p3 = _p4 + self.lateral_block4(x3)
+
+ p3 = self.decoder_block3(_p3)
+ _p3 = F.interpolate(p3, size=x2.shape[2:], mode='bilinear', align_corners=True)
+ _p2 = _p3 + self.lateral_block3(x2)
+
+ p2 = self.decoder_block2(_p2)
+ _p2 = F.interpolate(p2, size=x1.shape[2:], mode='bilinear', align_corners=True)
+ _p1 = _p2 + self.lateral_block2(x1)
+
+ _p1 = self.decoder_block1(_p1)
+ _p1 = F.interpolate(_p1, size=x.shape[2:], mode='bilinear', align_corners=True)
+ p1_out = self.conv_out1(_p1)
+
+ if self.config.ms_supervision:
+ outs.append(self.conv_ms_spvn_4(p4))
+ outs.append(self.conv_ms_spvn_3(p3))
+ outs.append(self.conv_ms_spvn_2(p2))
+ outs.append(p1_out)
+ return outs
+
+
+class RefUNet(nn.Module):
+ # Refinement
+ def __init__(self, in_channels=3+1):
+ super(RefUNet, self).__init__()
+ self.encoder_1 = nn.Sequential(
+ nn.Conv2d(in_channels, 64, 3, 1, 1),
+ nn.Conv2d(64, 64, 3, 1, 1),
+ nn.BatchNorm2d(64),
+ nn.ReLU(inplace=True)
+ )
+
+ self.encoder_2 = nn.Sequential(
+ nn.MaxPool2d(2, 2, ceil_mode=True),
+ nn.Conv2d(64, 64, 3, 1, 1),
+ nn.BatchNorm2d(64),
+ nn.ReLU(inplace=True)
+ )
+
+ self.encoder_3 = nn.Sequential(
+ nn.MaxPool2d(2, 2, ceil_mode=True),
+ nn.Conv2d(64, 64, 3, 1, 1),
+ nn.BatchNorm2d(64),
+ nn.ReLU(inplace=True)
+ )
+
+ self.encoder_4 = nn.Sequential(
+ nn.MaxPool2d(2, 2, ceil_mode=True),
+ nn.Conv2d(64, 64, 3, 1, 1),
+ nn.BatchNorm2d(64),
+ nn.ReLU(inplace=True)
+ )
+
+ self.pool4 = nn.MaxPool2d(2, 2, ceil_mode=True)
+ #####
+ self.decoder_5 = nn.Sequential(
+ nn.Conv2d(64, 64, 3, 1, 1),
+ nn.BatchNorm2d(64),
+ nn.ReLU(inplace=True)
+ )
+ #####
+ self.decoder_4 = nn.Sequential(
+ nn.Conv2d(128, 64, 3, 1, 1),
+ nn.BatchNorm2d(64),
+ nn.ReLU(inplace=True)
+ )
+
+ self.decoder_3 = nn.Sequential(
+ nn.Conv2d(128, 64, 3, 1, 1),
+ nn.BatchNorm2d(64),
+ nn.ReLU(inplace=True)
+ )
+
+ self.decoder_2 = nn.Sequential(
+ nn.Conv2d(128, 64, 3, 1, 1),
+ nn.BatchNorm2d(64),
+ nn.ReLU(inplace=True)
+ )
+
+ self.decoder_1 = nn.Sequential(
+ nn.Conv2d(128, 64, 3, 1, 1),
+ nn.BatchNorm2d(64),
+ nn.ReLU(inplace=True)
+ )
+
+ self.conv_d0 = nn.Conv2d(64, 1, 3, 1, 1)
+
+ self.upscore2 = nn.Upsample(scale_factor=2, mode='bilinear', align_corners=True)
+
+ def forward(self, x):
+ outs = []
+ if isinstance(x, list):
+ x = torch.cat(x, dim=1)
+ hx = x
+
+ hx1 = self.encoder_1(hx)
+ hx2 = self.encoder_2(hx1)
+ hx3 = self.encoder_3(hx2)
+ hx4 = self.encoder_4(hx3)
+
+ hx = self.decoder_5(self.pool4(hx4))
+ hx = torch.cat((self.upscore2(hx), hx4), 1)
+
+ d4 = self.decoder_4(hx)
+ hx = torch.cat((self.upscore2(d4), hx3), 1)
+
+ d3 = self.decoder_3(hx)
+ hx = torch.cat((self.upscore2(d3), hx2), 1)
+
+ d2 = self.decoder_2(hx)
+ hx = torch.cat((self.upscore2(d2), hx1), 1)
+
+ d1 = self.decoder_1(hx)
+
+ x = self.conv_d0(d1)
+ outs.append(x)
+ return outs
diff --git a/py/BiRefNet_legacy/refinement/stem_layer.py b/py/BiRefNet_legacy/refinement/stem_layer.py
new file mode 100644
index 0000000..116d546
--- /dev/null
+++ b/py/BiRefNet_legacy/refinement/stem_layer.py
@@ -0,0 +1,45 @@
+import torch.nn as nn
+from BiRefNet_legacy.modules.utils import build_act_layer, build_norm_layer
+
+
+class StemLayer(nn.Module):
+ r""" Stem layer of InternImage
+ Args:
+ in_channels (int): number of input channels
+ out_channels (int): number of output channels
+ act_layer (str): activation layer
+ norm_layer (str): normalization layer
+ """
+
+ def __init__(self,
+ in_channels=3+1,
+ inter_channels=48,
+ out_channels=96,
+ act_layer='GELU',
+ norm_layer='BN'):
+ super().__init__()
+ self.conv1 = nn.Conv2d(in_channels,
+ inter_channels,
+ kernel_size=3,
+ stride=1,
+ padding=1)
+ self.norm1 = build_norm_layer(
+ inter_channels, norm_layer, 'channels_first', 'channels_first'
+ )
+ self.act = build_act_layer(act_layer)
+ self.conv2 = nn.Conv2d(inter_channels,
+ out_channels,
+ kernel_size=3,
+ stride=1,
+ padding=1)
+ self.norm2 = build_norm_layer(
+ out_channels, norm_layer, 'channels_first', 'channels_first'
+ )
+
+ def forward(self, x):
+ x = self.conv1(x)
+ x = self.norm1(x)
+ x = self.act(x)
+ x = self.conv2(x)
+ x = self.norm2(x)
+ return x
diff --git a/py/BiRefNet_v2/LICENSE b/py/BiRefNet_v2/LICENSE
new file mode 100644
index 0000000..485921e
--- /dev/null
+++ b/py/BiRefNet_v2/LICENSE
@@ -0,0 +1,21 @@
+MIT License
+
+Copyright (c) 2024 ZhengPeng
+
+Permission is hereby granted, free of charge, to any person obtaining a copy
+of this software and associated documentation files (the "Software"), to deal
+in the Software without restriction, including without limitation the rights
+to use, copy, modify, merge, publish, distribute, sublicense, and/or sell
+copies of the Software, and to permit persons to whom the Software is
+furnished to do so, subject to the following conditions:
+
+The above copyright notice and this permission notice shall be included in all
+copies or substantial portions of the Software.
+
+THE SOFTWARE IS PROVIDED "AS IS", WITHOUT WARRANTY OF ANY KIND, EXPRESS OR
+IMPLIED, INCLUDING BUT NOT LIMITED TO THE WARRANTIES OF MERCHANTABILITY,
+FITNESS FOR A PARTICULAR PURPOSE AND NONINFRINGEMENT. IN NO EVENT SHALL THE
+AUTHORS OR COPYRIGHT HOLDERS BE LIABLE FOR ANY CLAIM, DAMAGES OR OTHER
+LIABILITY, WHETHER IN AN ACTION OF CONTRACT, TORT OR OTHERWISE, ARISING FROM,
+OUT OF OR IN CONNECTION WITH THE SOFTWARE OR THE USE OR OTHER DEALINGS IN THE
+SOFTWARE.
diff --git a/py/BiRefNet_v2/README.md b/py/BiRefNet_v2/README.md
new file mode 100644
index 0000000..39f71d7
--- /dev/null
+++ b/py/BiRefNet_v2/README.md
@@ -0,0 +1,316 @@
+
Bilateral Reference for High-Resolution Dichotomous Image Segmentation
+
+
+
+
+ 1 Nankai University 2 Northwestern Polytechnical University 3 National University of Defense Technology
+
+ 4 Aalto University 5 Shanghai AI Laboratory 6 University of Trento
+
+
+
+
+
+
+
+| *DIS-Sample_1* | *DIS-Sample_2* |
+| :------------------------------: | :-------------------------------: |
+|
|
|
+
+This repo is the official implementation of "[**Bilateral Reference for High-Resolution Dichotomous Image Segmentation**](https://arxiv.org/pdf/2401.03407)" (___CAAI AIR 2024___).
+
+> [!note]
+> **We need more GPU resources** to push forward the performance of BiRefNet, especially on *matting* tasks, higher-resolution inference (*2K*), and more *efficient* model design. If you are happy to cooperate, please contact me at zhengpeng0108@gmail.com.
+
+## News :newspaper:
+* **`Aug 30, 2024`:** We uploaded notebooks in `tutorials` to run the inference and ONNX conversion locally.
+* **`Aug 23, 2024`:** Our BiRefNet is now officially released [online](https://www.sciopen.com/article/10.26599/AIR.2024.9150038) on CAAI AIR journal. And thanks to the [press release](https://www.eurekalert.org/news-releases/1055380).
+* **`Aug 19, 2024`:** We uploaded the ONNX model files of all weights in the [GitHub release](https://github.com/ZhengPeng7/BiRefNet/releases/tag/v1) and [GDrive folder](https://drive.google.com/drive/u/0/folders/1kZM55bwsRdS__bdnsXpkmH6QPyza-9-N). Check out the **ONNX conversion** part in [model zoo](https://github.com/ZhengPeng7/BiRefNet?tab=readme-ov-file#model-zoo) for more details.
+* **`Jul 30, 2024`:** Thanks to @not-lain for his kind efforts in adding BiRefNet to the official huggingface.js [repo](https://github.com/huggingface/huggingface.js/blob/3a8651fbc6508920475564a692bf0e5b601d9343/packages/tasks/src/model-libraries-snippets.ts#L763).
+* **`Jul 28, 2024`:** We released the [Colab demo for box-guided segmentation](https://colab.research.google.com/drive/1B6aKZ3ekcvKMkSBn0N5mCASLUYMp0whK).
+* **`Jul 15, 2024`:** We deployed our BiRefNet on [Hugging Face Models](https://huggingface.co/ZhengPeng7/BiRefNet) for users to easily load it in one line code.
+* **`Jun 21, 2024`:** We released and uploaded the Chinese version of our original paper to my [GDrive](https://drive.google.com/file/d/1aBnJ_R9lbnC2dm8dqD0-pzP2Cu-U1Xpt/view).
+* **`May 28, 2024`:** We hold a [model zoo](https://github.com/ZhengPeng7/BiRefNet?tab=readme-ov-file#model-zoo) with well-trained weights of our BiRefNet in different sizes and for different tasks, including general use, matting segmentation, DIS, HRSOD, COD, etc.
+* **`May 7, 2024`:** We also released the [Colab demo for multiple images inference](https://colab.research.google.com/drive/14Dqg7oeBkFEtchaHLNpig2BcdkZEogba). Many thanks to @rishabh063 for his support on it.
+* **`Apr 9, 2024`:** Thanks to [Features and Labels Inc.](https://fal.ai/) for deploying a cool online BiRefNet [inference API](https://fal.ai/models/fal-ai/birefnet/playground) and providing me with strong GPU resources for 4 months on more extensive experiments!
+* **`Mar 7, 2024`:** We released BiRefNet codes, the well-trained weights for all tasks in the original papers, and all related stuff in my [GDrive folder](https://drive.google.com/drive/folders/1s2Xe0cjq-2ctnJBR24563yMSCOu4CcxM). Meanwhile, we also deployed our BiRefNet on [Hugging Face Spaces](https://huggingface.co/spaces/ZhengPeng7/BiRefNet_demo) for easier online use and released the [Colab demo for inference and evaluation](https://colab.research.google.com/drive/1MaEiBfJ4xIaZZn0DqKrhydHB8X97hNXl).
+* **`Jan 7, 2024`:** We released our paper on [arXiv](https://arxiv.org/pdf/2401.03407).
+
+
+## :rocket: Load BiRefNet in _ONE LINE_ by HuggingFace, check more: [](https://huggingface.co/ZhengPeng7/birefnet)
+```python
+from transformers import AutoModelForImageSegmentation
+birefnet = AutoModelForImageSegmentation.from_pretrained('zhengpeng7/BiRefNet', trust_remote_code=True)
+```
+## :flight_arrival: Inference Partner:
+We are really happy to collaborate with [FAL](https://fal.ai) to deploy the **inference API** of BiRefNet. You can access this service via the link below:
++ https://fal.ai/models/fal-ai/birefnet
+
+Our BiRefNet has achieved SOTA on many similar HR tasks:
+
+**DIS**: [](https://paperswithcode.com/sota/dichotomous-image-segmentation-on-dis-te1?p=bilateral-reference-for-high-resolution) [](https://paperswithcode.com/sota/dichotomous-image-segmentation-on-dis-te2?p=bilateral-reference-for-high-resolution) [](https://paperswithcode.com/sota/dichotomous-image-segmentation-on-dis-te3?p=bilateral-reference-for-high-resolution) [](https://paperswithcode.com/sota/dichotomous-image-segmentation-on-dis-te4?p=bilateral-reference-for-high-resolution) [](https://paperswithcode.com/sota/dichotomous-image-segmentation-on-dis-vd?p=bilateral-reference-for-high-resolution)
+
+Figure of Comparison on DIS Papers with Codes (by the time of this work):
+
+
+
+
+
+
+
+
+**COD**:[](https://paperswithcode.com/sota/camouflaged-object-segmentation-on-cod?p=bilateral-reference-for-high-resolution) [](https://paperswithcode.com/sota/camouflaged-object-segmentation-on-nc4k?p=bilateral-reference-for-high-resolution) [](https://paperswithcode.com/sota/camouflaged-object-segmentation-on-camo?p=bilateral-reference-for-high-resolution) [](https://paperswithcode.com/sota/camouflaged-object-segmentation-on-chameleon?p=bilateral-reference-for-high-resolution)
+
+Figure of Comparison on COD Papers with Codes (by the time of this work):
+
+
+
+
+
+
+**HRSOD**: [](https://paperswithcode.com/sota/rgb-salient-object-detection-on-davis-s?p=bilateral-reference-for-high-resolution) [](https://paperswithcode.com/sota/rgb-salient-object-detection-on-hrsod?p=bilateral-reference-for-high-resolution) [](https://paperswithcode.com/sota/rgb-salient-object-detection-on-uhrsd?p=bilateral-reference-for-high-resolution) [](https://paperswithcode.com/sota/salient-object-detection-on-duts-te?p=bilateral-reference-for-high-resolution) [](https://paperswithcode.com/sota/salient-object-detection-on-dut-omron?p=bilateral-reference-for-high-resolution)
+
+Figure of Comparison on HRSOD Papers with Codes (by the time of this work):
+
+
+
+
+
+
+
+
+#### Try our online demos for inference:
+
++ **Inference and evaluation** of your given weights: [](https://colab.research.google.com/drive/1MaEiBfJ4xIaZZn0DqKrhydHB8X97hNXl)
++ **Online Inference with GUI** with adjustable resolutions: [](https://huggingface.co/spaces/ZhengPeng7/BiRefNet_demo)
++ Online **Multiple Images Inference** on Colab: [](https://colab.research.google.com/drive/14Dqg7oeBkFEtchaHLNpig2BcdkZEogba)
+
+
+
+
+
+## Model Zoo
+
+> For more general use of our BiRefNet, I extended the original academic one to more general ones for better real-life application.
+>
+> Datasets and datasets are suggested to be downloaded from official pages. But you can also download the packaged ones: [DIS](https://drive.google.com/drive/folders/1hZW6tAGPJwo9mPS7qGGGdpxuvuXiyoMJ), [HRSOD](https://drive.google.com/drive/folders/18_hAE3QM4cwAzEAKXuSNtKjmgFXTQXZN), [COD](https://drive.google.com/drive/folders/1EyHmKWsXfaCR9O0BiZEc3roZbRcs4ECO), [Backbones](https://drive.google.com/drive/folders/1cmce_emsS8A5ha5XT2c_CZiJzlLM81ms).
+>
+> Find performances (almost all metrics) of all models in the `exp-TASK_SETTINGS` folders in [[**stuff**](https://drive.google.com/drive/folders/1s2Xe0cjq-2ctnJBR24563yMSCOu4CcxM)].
+
+
+
+Models in the original paper, for comparison on benchmarks:
+
+| Task | Training Sets | Backbone | Download |
+| :---: | :-------------------------: | :-----------: | :----------------------------------------------------------: |
+| DIS | DIS5K-TR | swin_v1_large | [google-drive](https://drive.google.com/file/d/1J90LucvDQaS3R_-9E7QUh1mgJ8eQvccb/view) |
+| COD | COD10K-TR, CAMO-TR | swin_v1_large | [google-drive](https://drive.google.com/file/d/1tM5M72k7a8aKF-dYy-QXaqvfEhbFaWkC/view) |
+| HRSOD | DUTS-TR | swin_v1_large | [google-drive](https://drive.google.com/file/d/1f7L0Pb1Y3RkOMbqLCW_zO31dik9AiUFa/view) |
+| HRSOD | HRSOD-TR | swin_v1_large | google-drive |
+| HRSOD | UHRSD-TR | swin_v1_large | google-drive |
+| HRSOD | DUTS-TR, HRSOD-TR | swin_v1_large | [google-drive](https://drive.google.com/file/d/1WJooyTkhoDLllaqwbpur_9Hle0XTHEs_/view) |
+| HRSOD | DUTS-TR, UHRSD-TR | swin_v1_large | [google-drive](https://drive.google.com/file/d/1Pu1mv3ORobJatIuUoEuZaWDl2ylP3Gw7/view) |
+| HRSOD | HRSOD-TR, UHRSD-TR | swin_v1_large | [google-drive](https://drive.google.com/file/d/1xEh7fsgWGaS5c3IffMswasv0_u-aVM9E/view) |
+| HRSOD | DUTS-TR, HRSOD-TR, UHRSD-TR | swin_v1_large | [google-drive](https://drive.google.com/file/d/13FaxyyOwyCddfZn2vZo1xG1KNZ3cZ-6B/view) |
+
+
+
+
+
+Models trained with customed data (general, matting), for general use in practical application:
+
+| Task | Training Sets | Backbone | Test Set | Metric (S, wF[, HCE]) | Download |
+| :-----------------------: | :----------------------------------------------------------: | :-----------: | :-------: | :-------------------: | :----------------------------------------------------------: |
+| **general use** | DIS5K-TR,DIS-TEs, DUTS-TR_TE,HRSOD-TR_TE,UHRSD-TR_TE, HRS10K-TR_TE, TR-P3M-10k, TE-P3M-500-NP, TE-P3M-500-P, TR-humans | swin_v1_large | DIS-VD | 0.911, 0.875, 1069 | [google-drive](https://drive.google.com/file/d/1_IfUnu8Fpfn-nerB89FzdNXQ7zk6FKxc/view) |
+| **general use** | DIS5K-TR,DIS-TEs, DUTS-TR_TE,HRSOD-TR_TE,UHRSD-TR_TE, HRS10K-TR_TE, TR-P3M-10k, TE-P3M-500-NP, TE-P3M-500-P, TR-humans | swin_v1_tiny | DIS-VD | 0.882, 0.830, 1175 | [google-drive](https://drive.google.com/file/d/1fzInDWiE2n65tmjaHDSZpqhL0VME6-Yl/view) |
+| **general use** | DIS5K-TR, DIS-TEs | swin_v1_large | DIS-VD | 0.907, 0.865, 1059 | [google-drive](https://drive.google.com/file/d/1P6NJzG3Jf1sl7js2q1CPC3yqvBn_O8UJ/view) |
+| **matting segmentation** | [P3M-10k](https://github.com/JizhiziLi/P3M), [humans](https://huggingface.co/datasets/schirrmacher/humans) | swin_v1_large | P3M-500-P | 0.983, 0.989 | [google-drive](https://drive.google.com/file/d/1uUeXjEUoD2XF_6YjD_fsct-TJp7TFiqh) |
+
+
+
+
+
+Segmentation with box guidance:
+
++ Given box guidance: [](https://colab.research.google.com/drive/1B6aKZ3ekcvKMkSBn0N5mCASLUYMp0whK)
+
+
+
+
+
+Model efficiency:
+
+> Screenshot from the original paper. All tests are conducted on a single A100 GPU.
+
+
+
+
+
+
+
+ONNX conversion:
+
+> We converted from `.pth` weights files to `.onnx` files.
+> We referred a lot to the [Kazuhito00/BiRefNet-ONNX-Sample](https://github.com/Kazuhito00/BiRefNet-ONNX-Sample), many thanks to @Kazuhito00.
+
++ Check our [Colab demo for ONNX conversion](https://colab.research.google.com/drive/1z6OruR52LOvDDpnp516F-N4EyPGrp5om) or the [notebook file for local running](https://drive.google.com/file/d/1cgL2qyvOO5q3ySfhytypX46swdQwZLrJ), where you can do the conversion/inference by yourself and find all relevant info.
++ As tested, BiRefNets with SwinL (default backbone) cost `~90%` more time (the inference costs `~165ms` on an A100 GPU) using ONNX files. Meanwhile, BiRefNets with SwinT (lightweight) cost `~75%` more time (the inference costs `~93.8ms` on an A100 GPU) using ONNX files. Input resolution is `1024x1024` as default.
++ The results of the original pth files and the converted onnx files are slightly different, which is acceptable.
++ Pay attention to the compatibility among `onnxruntime-gpu, CUDA, and CUDNN` (we use `torch==2.0.1, cuda=11.8` here).
+
+
+
+
+## Third-Party Creations
+
+> Concerning edge devices with less computing power, we provide a lightweight version with `swin_v1_tiny` as the backbone, which is x4+ faster and x5+ smaller. The details can be found in [this issue](https://github.com/ZhengPeng7/BiRefNet/issues/11#issuecomment-2041033576) and links there.
+
+We found there've been some 3rd party applications based on our BiRefNet. Many thanks for their contribution to the community!
+Choose the one you like to try with clicks instead of codes:
+1. **Applications**:
+ + Thanks [**lbq779660843/BiRefNet-Tensorrt**](https://github.com/lbq779660843/BiRefNet-Tensorrt) and [**yuanyang1991/birefnet_tensorrt**](https://github.com/yuanyang1991/birefnet_tensorrt): they both provided the project to convert BiRefNet to **TensorRT**, which is faster and better for deployment. Their repos offer solid local establishment (Win and Linux) and [colab demo](https://colab.research.google.com/drive/1r8GkFPyMMO0OkMX6ih5FjZnUCQrl2SHV?usp=sharing), respectively. And @yuanyang1991 kindly offered the comparison among the inference efficiency of naive PyTorch, ONNX, and TensorRT on an RTX 4080S:
+
+| Methods | [Pytorch](https://drive.google.com/file/d/1_IfUnu8Fpfn-nerB89FzdNXQ7zk6FKxc/view) | [ONNX](https://drive.google.com/drive/u/0/folders/1kZM55bwsRdS__bdnsXpkmH6QPyza-9-N) | TensorRT |
+|:------------------------------------------------------------------------------------:|:--------------:|:--------------:|:--------------:|
+| First Inference Time | 0.71s | 5.32s | **0.17s** |
+
+| Methods | [Pytorch](https://drive.google.com/file/d/1_IfUnu8Fpfn-nerB89FzdNXQ7zk6FKxc/view) | [ONNX](https://drive.google.com/drive/u/0/folders/1kZM55bwsRdS__bdnsXpkmH6QPyza-9-N) | TensorRT |
+|:------------------------------------------------------------------------------------:|:--------------:|:--------------:|:--------------:|
+| Avg Inf Time (excluding 1st) | 0.15s | 4.43s | **0.11s** |
+
+ + Thanks [**dimitribarbot/sd-webui-birefnet**](https://github.com/dimitribarbot/sd-webui-birefnet): this project allows to add a BiRefNet section to the original **Stable Diffusion WebUI**'s Extras tab.
+ 
+
+ + Thanks [**fal.ai/birefnet**](https://fal.ai/models/birefnet): this project on `fal.ai` encapsulates BiRefNet **online** with more useful options in **UI** and **API** to call the model.
+ 
+
+ + Thanks [**ZHO-ZHO-ZHO/ComfyUI-BiRefNet-ZHO**](https://github.com/ZHO-ZHO-ZHO/ComfyUI-BiRefNet-ZHO): this project further improves the **UI** for BiRefNet in ComfyUI, especially for **video data**.
+ 
+
+
+
+ + Thanks [**viperyl/ComfyUI-BiRefNet**](https://github.com/viperyl/ComfyUI-BiRefNet): this project packs BiRefNet as **ComfyUI nodes**, and makes this SOTA model easier use for everyone.
+ 
+
+ + Thanks [**Rishabh**](https://github.com/rishabh063) for offering a demo for the [easier multiple images inference on colab](https://colab.research.google.com/drive/14Dqg7oeBkFEtchaHLNpig2BcdkZEogba).
+
+2. **More Visual Comparisons**
+ + Thanks [**twitter.com/ZHOZHO672070**](https://twitter.com/ZHOZHO672070) for the comparison with more background-removal methods in images:
+
+
+
+ + Thanks [**twitter.com/toyxyz3**](https://twitter.com/toyxyz3) for the comparison with more background-removal methods in videos:
+
+
+
+
+
+
+## Usage
+
+#### Environment Setup
+
+```shell
+# PyTorch==2.0.1 is used for faster training with compilation.
+conda create -n birefnet python=3.9 -y && conda activate birefnet
+pip install -r requirements.txt
+```
+
+#### Dataset Preparation
+
+Download combined training / test sets I have organized well from: [DIS](https://drive.google.com/drive/folders/1hZW6tAGPJwo9mPS7qGGGdpxuvuXiyoMJ)--[COD](https://drive.google.com/drive/folders/1EyHmKWsXfaCR9O0BiZEc3roZbRcs4ECO)--[HRSOD](https://drive.google.com/drive/folders/18_hAE3QM4cwAzEAKXuSNtKjmgFXTQXZN) or the single official ones in the `single_ones` folder, or their official pages. You can also find the same ones on my **BaiduDisk**: [DIS](https://pan.baidu.com/s/1O_pQIGAE4DKqL93xOxHpxw?pwd=PSWD)--[COD](https://pan.baidu.com/s/1RnxAzaHSTGBC1N6r_RfeqQ?pwd=PSWD)--[HRSOD](https://pan.baidu.com/s/1_Del53_0lBuG0DKJJAk4UA?pwd=PSWD).
+
+#### Weights Preparation
+
+Download backbone weights from [my google-drive folder](https://drive.google.com/drive/folders/1s2Xe0cjq-2ctnJBR24563yMSCOu4CcxM) or their official pages.
+
+## Run
+
+```shell
+# Train & Test & Evaluation
+./train_test.sh RUN_NAME GPU_NUMBERS_FOR_TRAINING GPU_NUMBERS_FOR_TEST
+# Example: ./train_test.sh tmp-proj 0,1,2,3,4,5,6,7 0
+
+# See train.sh / test.sh for only training / test-evaluation.
+# After the evaluation, run `gen_best_ep.py` to select the best ckpt from a specific metric (you choose it from Sm, wFm, HCE (DIS only)).
+```
+
+#### Well-trained weights:
+
+Download the `BiRefNet-{TASK}-{EPOCH}.pth` from [[**stuff**](https://drive.google.com/drive/folders/1s2Xe0cjq-2ctnJBR24563yMSCOu4CcxM)]. Info of the corresponding (predicted\_maps/performance/training\_log) weights can be also found in folders like `exp-BiRefNet-{TASK_SETTINGS}` in the same directory.
+
+You can also download the weights from the release of this repo.
+
+The results might be a bit different from those in the original paper, you can see them in the `eval_results-BiRefNet-{TASK_SETTINGS}` folder in each `exp-xx`, we will update them in the following days. Due to the very high cost I used (A100-80G x 8) which many people cannot afford to (including myself....), I re-trained BiRefNet on a single A100-40G only and achieve the performance on the same level (even better). It means you can directly train the model on a single GPU with 36.5G+ memory. BTW, 5.5G GPU memory is needed for inference in 1024x1024. (I personally paid a lot for renting an A100-40G to re-train BiRefNet on the three tasks... T_T. Hope it can help you.)
+
+But if you have more and more powerful GPUs, you can set GPU IDs and increase the batch size in `config.py` to accelerate the training. We have made all this kind of things adaptive in scripts to seamlessly switch between single-card training and multi-card training. Enjoy it :)
+
+#### Some of my messages:
+
+This project was originally built for DIS only. But after the updates one by one, I made it larger and larger with many functions embedded together. Finally, you can **use it for any binary image segmentation tasks**, such as DIS/COD/SOD, medical image segmentation, anomaly segmentation, etc. You can eaily open/close below things (usually in `config.py`):
++ Multi-GPU training: open/close with one variable.
++ Backbone choices: Swin_v1, PVT_v2, ConvNets, ...
++ Weighted losses: BCE, IoU, SSIM, MAE, Reg, ...
++ Adversarial loss for binary segmentation (proposed in my previous work [MCCL](https://arxiv.org/pdf/2302.14485)).
++ Training tricks: multi-scale supervision, freezing backbone, multi-scale input...
++ Data collator: loading all in memory, smooth combination of different datasets for combined training and test.
++ ...
+I really hope you enjoy this project and use it in more works to achieve new SOTAs.
+
+
+### Quantitative Results
+
+
+
+
+
+
+
+### Qualitative Results
+
+
+
+
+
+
+
+### Citation
+
+```
+@article{zheng2024birefnet,
+ title={Bilateral Reference for High-Resolution Dichotomous Image Segmentation},
+ author={Zheng, Peng and Gao, Dehong and Fan, Deng-Ping and Liu, Li and Laaksonen, Jorma and Ouyang, Wanli and Sebe, Nicu},
+ journal={CAAI Artificial Intelligence Research},
+ volume = {3},
+ pages = {9150038},
+ year={2024}
+}
+```
+
+
+
+## Contact
+
+Any questions, discussions, or even complaints, feel free to leave issues here or send me e-mails (zhengpeng0108@gmail.com). You can also join the Discord Group (https://discord.gg/d9NN5sgFrq) or QQ Group (https://qm.qq.com/q/y6WPy7WOIK) if you want to talk a lot publicly.
+
diff --git a/py/BiRefNet_v2/__init__.py b/py/BiRefNet_v2/__init__.py
new file mode 100644
index 0000000..e69de29
diff --git a/py/BiRefNet_v2/config.py b/py/BiRefNet_v2/config.py
new file mode 100644
index 0000000..c333160
--- /dev/null
+++ b/py/BiRefNet_v2/config.py
@@ -0,0 +1,174 @@
+import os
+import math
+
+
+class Config():
+ def __init__(self) -> None:
+ # PATH settings
+ # Make up your file system as: SYS_HOME_DIR/codes/dis/BiRefNet, SYS_HOME_DIR/datasets/dis/xx, SYS_HOME_DIR/weights/xx
+ if os.name == 'nt':
+ self.sys_home_dir = os.environ['USERPROFILE'] # For windows system
+ else:
+ self.sys_home_dir = os.environ['HOME'] # For Linux system
+
+ # TASK settings
+ self.task = ['DIS5K', 'COD', 'HRSOD', 'General', 'Matting'][0]
+ self.training_set = {
+ 'DIS5K': ['DIS-TR', 'DIS-TR+DIS-TE1+DIS-TE2+DIS-TE3+DIS-TE4'][0],
+ 'COD': 'TR-COD10K+TR-CAMO',
+ 'HRSOD': ['TR-DUTS', 'TR-HRSOD', 'TR-UHRSD', 'TR-DUTS+TR-HRSOD', 'TR-DUTS+TR-UHRSD', 'TR-HRSOD+TR-UHRSD', 'TR-DUTS+TR-HRSOD+TR-UHRSD'][5],
+ 'General': 'DIS-TE1+DIS-TE2+DIS-TE3+DIS-TE4+DIS-TR+TR-HRSOD+TE-HRSOD+TR-HRS10K+TE-HRS10K+TR-UHRSD+TE-UHRSD+TR-P3M-10k+TE-P3M-500-NP+TE-P3M-500-P+TR-humans', # leave DIS-VD for evaluation.
+ 'Matting': 'TR-P3M-10k+TE-P3M-500-NP+TR-humans+TR-Distrinctions-646',
+ }[self.task]
+ self.prompt4loc = ['dense', 'sparse'][0]
+
+ # Faster-Training settings
+ self.load_all = False # Turn it on/off by your case. It may consume a lot of CPU memory. And for multi-GPU (N), it would cost N times the CPU memory to load the data.
+ self.use_fp16 = False # It may cause nan in training.
+ self.compile = True and (not self.use_fp16) # 1. Trigger CPU memory leak in some extend, which is an inherent problem of PyTorch.
+ # Machines with > 70GB CPU memory can run the whole training on DIS5K with default setting.
+ # 2. Higher PyTorch version may fix it: https://github.com/pytorch/pytorch/issues/119607.
+ # 3. But compile in Pytorch > 2.0.1 seems to bring no acceleration for training.
+ self.precisionHigh = True
+
+ # MODEL settings
+ self.ms_supervision = True
+ self.out_ref = self.ms_supervision and True
+ self.dec_ipt = True
+ self.dec_ipt_split = True
+ self.cxt_num = [0, 3][1] # multi-scale skip connections from encoder
+ self.mul_scl_ipt = ['', 'add', 'cat'][2]
+ self.dec_att = ['', 'ASPP', 'ASPPDeformable'][2]
+ self.squeeze_block = ['', 'BasicDecBlk_x1', 'ResBlk_x4', 'ASPP_x3', 'ASPPDeformable_x3'][1]
+ self.dec_blk = ['BasicDecBlk', 'ResBlk'][0]
+
+ # TRAINING settings
+ self.batch_size = 4
+ self.finetune_last_epochs = [
+ ('IoU', 0),
+ {
+ 'DIS5K': ('IoU', -30),
+ 'COD': ('IoU', -20),
+ 'HRSOD': ('IoU', -20),
+ 'General': ('MAE', -10),
+ 'Matting': ('MAE', -10),
+ }[self.task]
+ ][1] # choose 0 to skip
+ self.lr = (1e-4 if 'DIS5K' in self.task else 1e-5) * math.sqrt(self.batch_size / 4) # DIS needs high lr to converge faster. Adapt the lr linearly
+ self.size = 1024
+ self.num_workers = max(4, self.batch_size) # will be decrease to min(it, batch_size) at the initialization of the data_loader
+
+ # Backbone settings
+ self.bb = [
+ 'vgg16', 'vgg16bn', 'resnet50', # 0, 1, 2
+ 'swin_v1_t', 'swin_v1_s', # 3, 4
+ 'swin_v1_b', 'swin_v1_l', # 5-bs9, 6-bs4
+ 'pvt_v2_b0', 'pvt_v2_b1', # 7, 8
+ 'pvt_v2_b2', 'pvt_v2_b5', # 9-bs10, 10-bs5
+ ][6]
+ self.lateral_channels_in_collection = {
+ 'vgg16': [512, 256, 128, 64], 'vgg16bn': [512, 256, 128, 64], 'resnet50': [1024, 512, 256, 64],
+ 'pvt_v2_b2': [512, 320, 128, 64], 'pvt_v2_b5': [512, 320, 128, 64],
+ 'swin_v1_b': [1024, 512, 256, 128], 'swin_v1_l': [1536, 768, 384, 192],
+ 'swin_v1_t': [768, 384, 192, 96], 'swin_v1_s': [768, 384, 192, 96],
+ 'pvt_v2_b0': [256, 160, 64, 32], 'pvt_v2_b1': [512, 320, 128, 64],
+ }[self.bb]
+ if self.mul_scl_ipt == 'cat':
+ self.lateral_channels_in_collection = [channel * 2 for channel in self.lateral_channels_in_collection]
+ self.cxt = self.lateral_channels_in_collection[1:][::-1][-self.cxt_num:] if self.cxt_num else []
+
+ # MODEL settings - inactive
+ self.lat_blk = ['BasicLatBlk'][0]
+ self.dec_channels_inter = ['fixed', 'adap'][0]
+ self.refine = ['', 'itself', 'RefUNet', 'Refiner', 'RefinerPVTInChannels4'][0]
+ self.progressive_ref = self.refine and True
+ self.ender = self.progressive_ref and False
+ self.scale = self.progressive_ref and 2
+ self.auxiliary_classification = False # Only for DIS5K, where class labels are saved in `dataset.py`.
+ self.refine_iteration = 1
+ self.freeze_bb = False
+ self.model = [
+ 'BiRefNet',
+ ][0]
+
+ # TRAINING settings - inactive
+ self.preproc_methods = ['flip', 'enhance', 'rotate', 'pepper', 'crop'][:4]
+ self.optimizer = ['Adam', 'AdamW'][1]
+ self.lr_decay_epochs = [1e5] # Set to negative N to decay the lr in the last N-th epoch.
+ self.lr_decay_rate = 0.5
+ # Loss
+ if self.task not in ['Matting']:
+ self.lambdas_pix_last = {
+ # not 0 means opening this loss
+ # original rate -- 1 : 30 : 1.5 : 0.2, bce x 30
+ 'bce': 30 * 1, # high performance
+ 'iou': 0.5 * 1, # 0 / 255
+ 'iou_patch': 0.5 * 0, # 0 / 255, win_size = (64, 64)
+ 'mae': 30 * 0,
+ 'mse': 30 * 0, # can smooth the saliency map
+ 'triplet': 3 * 0,
+ 'reg': 100 * 0,
+ 'ssim': 10 * 1, # help contours,
+ 'cnt': 5 * 0, # help contours
+ 'structure': 5 * 0, # structure loss from codes of MVANet. A little improvement on DIS-TE[1,2,3], a bit more decrease on DIS-TE4.
+ }
+ else:
+ self.lambdas_pix_last = {
+ # not 0 means opening this loss
+ # original rate -- 1 : 30 : 1.5 : 0.2, bce x 30
+ 'bce': 30 * 0, # high performance
+ 'iou': 0.5 * 0, # 0 / 255
+ 'iou_patch': 0.5 * 0, # 0 / 255, win_size = (64, 64)
+ 'mae': 100 * 1,
+ 'mse': 30 * 0, # can smooth the saliency map
+ 'triplet': 3 * 0,
+ 'reg': 100 * 0,
+ 'ssim': 10 * 1, # help contours,
+ 'cnt': 5 * 0, # help contours
+ 'structure': 5 * 0, # structure loss from codes of MVANet. A little improvement on DIS-TE[1,2,3], a bit more decrease on DIS-TE4.
+ }
+ self.lambdas_cls = {
+ 'ce': 5.0
+ }
+ # Adv
+ self.lambda_adv_g = 10. * 0 # turn to 0 to avoid adv training
+ self.lambda_adv_d = 3. * (self.lambda_adv_g > 0)
+
+ # PATH settings - inactive
+ self.data_root_dir = os.path.join(self.sys_home_dir, 'datasets/dis')
+ self.weights_root_dir = os.path.join(self.sys_home_dir, 'weights')
+ self.weights = {
+ 'pvt_v2_b2': os.path.join(self.weights_root_dir, 'pvt_v2_b2.pth'),
+ 'pvt_v2_b5': os.path.join(self.weights_root_dir, ['pvt_v2_b5.pth', 'pvt_v2_b5_22k.pth'][0]),
+ 'swin_v1_b': os.path.join(self.weights_root_dir, ['swin_base_patch4_window12_384_22kto1k.pth', 'swin_base_patch4_window12_384_22k.pth'][0]),
+ 'swin_v1_l': os.path.join(self.weights_root_dir, ['swin_large_patch4_window12_384_22kto1k.pth', 'swin_large_patch4_window12_384_22k.pth'][0]),
+ 'swin_v1_t': os.path.join(self.weights_root_dir, ['swin_tiny_patch4_window7_224_22kto1k_finetune.pth'][0]),
+ 'swin_v1_s': os.path.join(self.weights_root_dir, ['swin_small_patch4_window7_224_22kto1k_finetune.pth'][0]),
+ 'pvt_v2_b0': os.path.join(self.weights_root_dir, ['pvt_v2_b0.pth'][0]),
+ 'pvt_v2_b1': os.path.join(self.weights_root_dir, ['pvt_v2_b1.pth'][0]),
+ }
+
+ # Callbacks - inactive
+ self.verbose_eval = True
+ self.only_S_MAE = False
+ self.SDPA_enabled = False # Bugs. Slower and errors occur in multi-GPUs
+
+ # others
+ self.device = [0, 'cpu'][0] # .to(0) == .to('cuda:0')
+
+ self.batch_size_valid = 1
+ self.rand_seed = 7
+ run_sh_file = [f for f in os.listdir('.') if 'train.sh' == f] + [os.path.join('..', f) for f in os.listdir('..') if 'train.sh' == f]
+ if run_sh_file:
+ with open(run_sh_file[0], 'r') as f:
+ lines = f.readlines()
+ self.save_last = int([l.strip() for l in lines if '"{}")'.format(self.task) in l and 'val_last=' in l][0].split('val_last=')[-1].split()[0])
+
+ def print_task(self) -> None:
+ # Return task for choosing settings in shell scripts.
+ print(self.task)
+
+if __name__ == '__main__':
+ config = Config()
+ config.print_task()
+
diff --git a/py/BiRefNet_v2/dataset.py b/py/BiRefNet_v2/dataset.py
new file mode 100644
index 0000000..a7d9e13
--- /dev/null
+++ b/py/BiRefNet_v2/dataset.py
@@ -0,0 +1,118 @@
+import os
+import cv2
+from tqdm import tqdm
+from PIL import Image
+from torch.utils import data
+from torchvision import transforms
+
+from .image_proc import preproc
+from .config import Config
+from .utils import path_to_image
+
+
+Image.MAX_IMAGE_PIXELS = None # remove DecompressionBombWarning
+config = Config()
+_class_labels_TR_sorted = (
+ 'Airplane, Ant, Antenna, Archery, Axe, BabyCarriage, Bag, BalanceBeam, Balcony, Balloon, Basket, BasketballHoop, Beatle, Bed, Bee, Bench, Bicycle, '
+ 'BicycleFrame, BicycleStand, Boat, Bonsai, BoomLift, Bridge, BunkBed, Butterfly, Button, Cable, CableLift, Cage, Camcorder, Cannon, Canoe, Car, '
+ 'CarParkDropArm, Carriage, Cart, Caterpillar, CeilingLamp, Centipede, Chair, Clip, Clock, Clothes, CoatHanger, Comb, ConcretePumpTruck, Crack, Crane, '
+ 'Cup, DentalChair, Desk, DeskChair, Diagram, DishRack, DoorHandle, Dragonfish, Dragonfly, Drum, Earphone, Easel, ElectricIron, Excavator, Eyeglasses, '
+ 'Fan, Fence, Fencing, FerrisWheel, FireExtinguisher, Fishing, Flag, FloorLamp, Forklift, GasStation, Gate, Gear, Goal, Golf, GymEquipment, Hammock, '
+ 'Handcart, Handcraft, Handrail, HangGlider, Harp, Harvester, Headset, Helicopter, Helmet, Hook, HorizontalBar, Hydrovalve, IroningTable, Jewelry, Key, '
+ 'KidsPlayground, Kitchenware, Kite, Knife, Ladder, LaundryRack, Lightning, Lobster, Locust, Machine, MachineGun, MagazineRack, Mantis, Medal, MemorialArchway, '
+ 'Microphone, Missile, MobileHolder, Monitor, Mosquito, Motorcycle, MovingTrolley, Mower, MusicPlayer, MusicStand, ObservationTower, Octopus, OilWell, '
+ 'OlympicLogo, OperatingTable, OutdoorFitnessEquipment, Parachute, Pavilion, Piano, Pipe, PlowHarrow, PoleVault, Punchbag, Rack, Racket, Rifle, Ring, Robot, '
+ 'RockClimbing, Rope, Sailboat, Satellite, Scaffold, Scale, Scissor, Scooter, Sculpture, Seadragon, Seahorse, Seal, SewingMachine, Ship, Shoe, ShoppingCart, '
+ 'ShoppingTrolley, Shower, Shrimp, Signboard, Skateboarding, Skeleton, Skiing, Spade, SpeedBoat, Spider, Spoon, Stair, Stand, Stationary, SteeringWheel, '
+ 'Stethoscope, Stool, Stove, StreetLamp, SweetStand, Swing, Sword, TV, Table, TableChair, TableLamp, TableTennis, Tank, Tapeline, Teapot, Telescope, Tent, '
+ 'TobaccoPipe, Toy, Tractor, TrafficLight, TrafficSign, Trampoline, TransmissionTower, Tree, Tricycle, TrimmerCover, Tripod, Trombone, Truck, Trumpet, Tuba, '
+ 'UAV, Umbrella, UnevenBars, UtilityPole, VacuumCleaner, Violin, Wakesurfing, Watch, WaterTower, WateringPot, Well, WellLid, Wheel, Wheelchair, WindTurbine, Windmill, WineGlass, WireWhisk, Yacht'
+)
+class_labels_TR_sorted = _class_labels_TR_sorted.split(', ')
+
+
+class MyData(data.Dataset):
+ def __init__(self, datasets, image_size, is_train=True):
+ self.size_train = image_size
+ self.size_test = image_size
+ self.keep_size = not config.size
+ self.data_size = (config.size, config.size)
+ self.is_train = is_train
+ self.load_all = config.load_all
+ self.device = config.device
+ valid_extensions = ['.png', '.jpg', '.PNG', '.JPG', '.JPEG']
+
+ if self.is_train and config.auxiliary_classification:
+ self.cls_name2id = {_name: _id for _id, _name in enumerate(class_labels_TR_sorted)}
+ self.transform_image = transforms.Compose([
+ transforms.Resize(self.data_size),
+ transforms.ToTensor(),
+ transforms.Normalize([0.485, 0.456, 0.406], [0.229, 0.224, 0.225]),
+ ][self.load_all or self.keep_size:])
+ self.transform_label = transforms.Compose([
+ transforms.Resize(self.data_size),
+ transforms.ToTensor(),
+ ][self.load_all or self.keep_size:])
+ dataset_root = os.path.join(config.data_root_dir, config.task)
+ # datasets can be a list of different datasets for training on combined sets.
+ self.image_paths = []
+ for dataset in datasets.split('+'):
+ image_root = os.path.join(dataset_root, dataset, 'im')
+ self.image_paths += [os.path.join(image_root, p) for p in os.listdir(image_root) if any(p.endswith(ext) for ext in valid_extensions)]
+ self.label_paths = []
+ for p in self.image_paths:
+ for ext in valid_extensions:
+ ## 'im' and 'gt' may need modifying
+ p_gt = p.replace('/im/', '/gt/')[:-(len(p.split('.')[-1])+1)] + ext
+ file_exists = False
+ if os.path.exists(p_gt):
+ self.label_paths.append(p_gt)
+ file_exists = True
+ break
+ if not file_exists:
+ print('Not exists:', p_gt)
+
+ if len(self.label_paths) != len(self.image_paths):
+ raise ValueError(f"There are different numbers of images ({len(self.label_paths)}) and labels ({len(self.image_paths)})")
+
+ if self.load_all:
+ self.images_loaded, self.labels_loaded = [], []
+ self.class_labels_loaded = []
+ # for image_path, label_path in zip(self.image_paths, self.label_paths):
+ for image_path, label_path in tqdm(zip(self.image_paths, self.label_paths), total=len(self.image_paths)):
+ _image = path_to_image(image_path, size=(config.size, config.size), color_type='rgb')
+ _label = path_to_image(label_path, size=(config.size, config.size), color_type='gray')
+ self.images_loaded.append(_image)
+ self.labels_loaded.append(_label)
+ self.class_labels_loaded.append(
+ self.cls_name2id[label_path.split('/')[-1].split('#')[3]] if self.is_train and config.auxiliary_classification else -1
+ )
+
+ def __getitem__(self, index):
+
+ if self.load_all:
+ image = self.images_loaded[index]
+ label = self.labels_loaded[index]
+ class_label = self.class_labels_loaded[index] if self.is_train and config.auxiliary_classification else -1
+ else:
+ image = path_to_image(self.image_paths[index], size=(config.size, config.size), color_type='rgb')
+ label = path_to_image(self.label_paths[index], size=(config.size, config.size), color_type='gray')
+ class_label = self.cls_name2id[self.label_paths[index].split('/')[-1].split('#')[3]] if self.is_train and config.auxiliary_classification else -1
+
+ # loading image and label
+ if self.is_train:
+ image, label = preproc(image, label, preproc_methods=config.preproc_methods)
+ # else:
+ # if _label.shape[0] > 2048 or _label.shape[1] > 2048:
+ # _image = cv2.resize(_image, (2048, 2048), interpolation=cv2.INTER_LINEAR)
+ # _label = cv2.resize(_label, (2048, 2048), interpolation=cv2.INTER_LINEAR)
+
+ image, label = self.transform_image(image), self.transform_label(label)
+
+ if self.is_train:
+ return image, label, class_label
+ else:
+ return image, label, self.label_paths[index]
+
+ def __len__(self):
+ return len(self.image_paths)
diff --git a/py/BiRefNet_v2/eval_existingOnes.py b/py/BiRefNet_v2/eval_existingOnes.py
new file mode 100644
index 0000000..9a66c93
--- /dev/null
+++ b/py/BiRefNet_v2/eval_existingOnes.py
@@ -0,0 +1,146 @@
+import os
+import argparse
+from glob import glob
+import prettytable as pt
+
+from .evaluation.evaluate import evaluator
+from .config import Config
+
+
+config = Config()
+
+
+def do_eval(args):
+ # evaluation for whole dataset
+ # dataset first in evaluation
+ for _data_name in args.data_lst.split('+'):
+ pred_data_dir = sorted(glob(os.path.join(args.pred_root, args.model_lst[0], _data_name)))
+ if not pred_data_dir:
+ print('Skip dataset {}.'.format(_data_name))
+ continue
+ gt_src = os.path.join(args.gt_root, _data_name)
+ gt_paths = sorted(glob(os.path.join(gt_src, 'gt', '*')))
+ print('#' * 20, _data_name, '#' * 20)
+ filename = os.path.join(args.save_dir, '{}_eval.txt'.format(_data_name))
+ tb = pt.PrettyTable()
+ tb.vertical_char = '&'
+ if config.task == 'DIS5K':
+ tb.field_names = ["Dataset", "Method", "maxFm", "wFmeasure", 'MAE', "Smeasure", "meanEm", "HCE", "maxEm", "meanFm", "adpEm", "adpFm", 'mBA', 'maxBIoU', 'meanBIoU']
+ elif config.task == 'COD':
+ tb.field_names = ["Dataset", "Method", "Smeasure", "wFmeasure", "meanFm", "meanEm", "maxEm", 'MAE', "maxFm", "adpEm", "adpFm", "HCE", 'mBA', 'maxBIoU', 'meanBIoU']
+ elif config.task == 'HRSOD':
+ tb.field_names = ["Dataset", "Method", "Smeasure", "maxFm", "meanEm", 'MAE', "maxEm", "meanFm", "wFmeasure", "adpEm", "adpFm", "HCE", 'mBA', 'maxBIoU', 'meanBIoU']
+ elif config.task == 'General':
+ tb.field_names = ["Dataset", "Method", "maxFm", "wFmeasure", 'MAE', "Smeasure", "meanEm", "HCE", "maxEm", "meanFm", "adpEm", "adpFm", 'mBA', 'maxBIoU', 'meanBIoU']
+ elif config.task == 'Matting':
+ tb.field_names = ["Dataset", "Method", "Smeasure", "maxFm", "meanEm", 'MSE', "maxEm", "meanFm", "wFmeasure", "adpEm", "adpFm", "HCE", 'mBA', 'maxBIoU', 'meanBIoU']
+ else:
+ tb.field_names = ["Dataset", "Method", "Smeasure", 'MAE', "maxEm", "meanEm", "maxFm", "meanFm", "wFmeasure", "adpEm", "adpFm", "HCE", 'mBA', 'maxBIoU', 'meanBIoU']
+ for _model_name in args.model_lst[:]:
+ print('\t', 'Evaluating model: {}...'.format(_model_name))
+ pred_paths = [p.replace(args.gt_root, os.path.join(args.pred_root, _model_name)).replace('/gt/', '/') for p in gt_paths]
+ # print(pred_paths[:1], gt_paths[:1])
+ em, sm, fm, mae, wfm, hce, mba, biou = evaluator(
+ gt_paths=gt_paths,
+ pred_paths=pred_paths,
+ metrics=args.metrics.split('+'),
+ verbose=config.verbose_eval
+ )
+ if config.task == 'DIS5K':
+ scores = [
+ fm['curve'].max().round(3), wfm.round(3), mae.round(3), sm.round(3), em['curve'].mean().round(3), int(hce.round()),
+ em['curve'].max().round(3), fm['curve'].mean().round(3), em['adp'].round(3), fm['adp'].round(3),
+ mba.round(3), biou['curve'].max().round(3), biou['curve'].mean().round(3),
+ ]
+ elif config.task == 'COD':
+ scores = [
+ sm.round(3), wfm.round(3), fm['curve'].mean().round(3), em['curve'].mean().round(3), em['curve'].max().round(3), mae.round(3),
+ fm['curve'].max().round(3), em['adp'].round(3), fm['adp'].round(3), int(hce.round()),
+ mba.round(3), biou['curve'].max().round(3), biou['curve'].mean().round(3),
+ ]
+ elif config.task == 'HRSOD':
+ scores = [
+ sm.round(3), fm['curve'].max().round(3), em['curve'].mean().round(3), mae.round(3),
+ em['curve'].max().round(3), fm['curve'].mean().round(3), wfm.round(3), em['adp'].round(3), fm['adp'].round(3), int(hce.round()),
+ mba.round(3), biou['curve'].max().round(3), biou['curve'].mean().round(3),
+ ]
+ elif config.task == 'General':
+ scores = [
+ fm['curve'].max().round(3), wfm.round(3), mae.round(3), sm.round(3), em['curve'].mean().round(3), int(hce.round()),
+ em['curve'].max().round(3), fm['curve'].mean().round(3), em['adp'].round(3), fm['adp'].round(3),
+ mba.round(3), biou['curve'].max().round(3), biou['curve'].mean().round(3),
+ ]
+ elif config.task == 'Matting':
+ scores = [
+ sm.round(3), fm['curve'].max().round(3), em['curve'].mean().round(3), mse.round(3),
+ em['curve'].max().round(3), fm['curve'].mean().round(3), wfm.round(3), em['adp'].round(3), fm['adp'].round(3), int(hce.round()),
+ mba.round(3), biou['curve'].max().round(3), biou['curve'].mean().round(3),
+ ]
+ else:
+ scores = [
+ sm.round(3), mae.round(3), em['curve'].max().round(3), em['curve'].mean().round(3),
+ fm['curve'].max().round(3), fm['curve'].mean().round(3), wfm.round(3),
+ em['adp'].round(3), fm['adp'].round(3), int(hce.round()),
+ mba.round(3), biou['curve'].max().round(3), biou['curve'].mean().round(3),
+ ]
+
+ for idx_score, score in enumerate(scores):
+ scores[idx_score] = '.' + format(score, '.3f').split('.')[-1] if score <= 1 else format(score, '<4')
+ records = [_data_name, _model_name] + scores
+ tb.add_row(records)
+ # Write results after every check.
+ with open(filename, 'w+') as file_to_write:
+ file_to_write.write(str(tb)+'\n')
+ print(tb)
+
+
+if __name__ == '__main__':
+ # set parameters
+ parser = argparse.ArgumentParser()
+ parser.add_argument(
+ '--gt_root', type=str, help='ground-truth root',
+ default=os.path.join(config.data_root_dir, config.task))
+ parser.add_argument(
+ '--pred_root', type=str, help='prediction root',
+ default='./e_preds')
+ parser.add_argument(
+ '--data_lst', type=str, help='test dataset',
+ default={
+ 'DIS5K': '+'.join(['DIS-VD', 'DIS-TE1', 'DIS-TE2', 'DIS-TE3', 'DIS-TE4'][:]),
+ 'COD': '+'.join(['TE-COD10K', 'NC4K', 'TE-CAMO', 'CHAMELEON'][:]),
+ 'HRSOD': '+'.join(['DAVIS-S', 'TE-HRSOD', 'TE-UHRSD', 'TE-DUTS', 'DUT-OMRON'][:]),
+ 'General': '+'.join(['DIS-VD'][:]),
+ 'Matting': '+'.join(['TE-P3M-500-P'][:]),
+ }[config.task])
+ parser.add_argument(
+ '--save_dir', type=str, help='candidate competitors',
+ default='e_results')
+ parser.add_argument(
+ '--check_integrity', type=bool, help='whether to check the file integrity',
+ default=False)
+ parser.add_argument(
+ '--metrics', type=str, help='candidate competitors',
+ default='+'.join(['S', 'MAE', 'E', 'F', 'WF', 'MBA', 'BIoU', 'HCE'][:100 if 'DIS5K' in config.task else -1]))
+ args = parser.parse_args()
+ args.metrics = '+'.join(['S', 'MAE', 'E', 'F', 'WF', 'MBA', 'BIoU', 'HCE'][:100 if sum(['DIS-' in _data for _data in args.data_lst.split('+')]) else -1])
+
+ os.makedirs(args.save_dir, exist_ok=True)
+ try:
+ args.model_lst = [m for m in sorted(os.listdir(args.pred_root), key=lambda x: int(x.split('epoch_')[-1]), reverse=True) if int(m.split('epoch_')[-1]) % 1 == 0]
+ except:
+ args.model_lst = [m for m in sorted(os.listdir(args.pred_root))]
+
+ # check the integrity of each candidates
+ if args.check_integrity:
+ for _data_name in args.data_lst.split('+'):
+ for _model_name in args.model_lst:
+ gt_pth = os.path.join(args.gt_root, _data_name)
+ pred_pth = os.path.join(args.pred_root, _model_name, _data_name)
+ if not sorted(os.listdir(gt_pth)) == sorted(os.listdir(pred_pth)):
+ print(len(sorted(os.listdir(gt_pth))), len(sorted(os.listdir(pred_pth))))
+ print('The {} Dataset of {} Model is not matching to the ground-truth'.format(_data_name, _model_name))
+ else:
+ print('>>> skip check the integrity of each candidates')
+
+ # start engine
+ do_eval(args)
diff --git a/py/BiRefNet_v2/evaluation/metrics.py b/py/BiRefNet_v2/evaluation/metrics.py
new file mode 100644
index 0000000..76ebc45
--- /dev/null
+++ b/py/BiRefNet_v2/evaluation/metrics.py
@@ -0,0 +1,763 @@
+import os
+from tqdm import tqdm
+import cv2
+import numpy as np
+from scipy.ndimage import convolve, distance_transform_edt as bwdist
+from skimage.morphology import skeletonize
+from skimage.morphology import disk
+from skimage.measure import label
+
+
+_EPS = np.spacing(1)
+_TYPE = np.float64
+
+
+def evaluator(gt_paths, pred_paths, metrics=['S', 'MAE', 'E', 'F', 'WF', 'MBA', 'BIoU', 'HCE'], verbose=False):
+ # define measures
+ if 'E' in metrics:
+ EM = EMeasure()
+ if 'S' in metrics:
+ SM = SMeasure()
+ if 'F' in metrics:
+ FM = FMeasure()
+ if 'MAE' in metrics:
+ MAE = MAEMeasure()
+ if 'WF' in metrics:
+ WFM = WeightedFMeasure()
+ if 'HCE' in metrics:
+ HCE = HCEMeasure()
+ if 'MBA' in metrics:
+ MBA = MBAMeasure()
+ if 'BIoU' in metrics:
+ BIoU = BIoUMeasure()
+
+ if isinstance(gt_paths, list) and isinstance(pred_paths, list):
+ # print(len(gt_paths), len(pred_paths))
+ assert len(gt_paths) == len(pred_paths)
+
+ for idx_sample in tqdm(range(len(gt_paths)), total=len(gt_paths)) if verbose else range(len(gt_paths)):
+ gt = gt_paths[idx_sample]
+ pred = pred_paths[idx_sample]
+
+ pred = pred[:-4] + '.png'
+ valid_extensions = ['.png', '.jpg', '.PNG', '.JPG', '.JPEG']
+ file_exists = False
+ for ext in valid_extensions:
+ if os.path.exists(pred[:-4] + ext):
+ pred = pred[:-4] + ext
+ file_exists = True
+ break
+ if file_exists:
+ pred_ary = cv2.imread(pred, cv2.IMREAD_GRAYSCALE)
+ else:
+ print('Not exists:', pred)
+
+ gt_ary = cv2.imread(gt, cv2.IMREAD_GRAYSCALE)
+ pred_ary = cv2.resize(pred_ary, (gt_ary.shape[1], gt_ary.shape[0]))
+
+ if 'E' in metrics:
+ EM.step(pred=pred_ary, gt=gt_ary)
+ if 'S' in metrics:
+ SM.step(pred=pred_ary, gt=gt_ary)
+ if 'F' in metrics:
+ FM.step(pred=pred_ary, gt=gt_ary)
+ if 'MAE' in metrics:
+ MAE.step(pred=pred_ary, gt=gt_ary)
+ if 'WF' in metrics:
+ WFM.step(pred=pred_ary, gt=gt_ary)
+ if 'HCE' in metrics:
+ ske_path = gt.replace('/gt/', '/ske/')
+ if os.path.exists(ske_path):
+ ske_ary = cv2.imread(ske_path, cv2.IMREAD_GRAYSCALE)
+ ske_ary = ske_ary > 128
+ else:
+ ske_ary = skeletonize(gt_ary > 128)
+ ske_save_dir = os.path.join(*ske_path.split(os.sep)[:-1])
+ if ske_path[0] == os.sep:
+ ske_save_dir = os.sep + ske_save_dir
+ os.makedirs(ske_save_dir, exist_ok=True)
+ cv2.imwrite(ske_path, ske_ary.astype(np.uint8) * 255)
+ HCE.step(pred=pred_ary, gt=gt_ary, gt_ske=ske_ary)
+ if 'MBA' in metrics:
+ MBA.step(pred=pred_ary, gt=gt_ary)
+ if 'BIoU' in metrics:
+ BIoU.step(pred=pred_ary, gt=gt_ary)
+
+ if 'E' in metrics:
+ em = EM.get_results()['em']
+ else:
+ em = {'curve': np.array([np.float64(-1)]), 'adp': np.float64(-1)}
+ if 'S' in metrics:
+ sm = SM.get_results()['sm']
+ else:
+ sm = np.float64(-1)
+ if 'F' in metrics:
+ fm = FM.get_results()['fm']
+ else:
+ fm = {'curve': np.array([np.float64(-1)]), 'adp': np.float64(-1)}
+ if 'MAE' in metrics:
+ mae = MAE.get_results()['mae']
+ else:
+ mae = np.float64(-1)
+ if 'WF' in metrics:
+ wfm = WFM.get_results()['wfm']
+ else:
+ wfm = np.float64(-1)
+ if 'HCE' in metrics:
+ hce = HCE.get_results()['hce']
+ else:
+ hce = np.float64(-1)
+ if 'MBA' in metrics:
+ mba = MBA.get_results()['mba']
+ else:
+ mba = np.float64(-1)
+ if 'BIoU' in metrics:
+ biou = BIoU.get_results()['biou']
+ else:
+ biou = {'curve': np.array([np.float64(-1)])}
+
+ return em, sm, fm, mae, wfm, hce, mba, biou
+
+
+def _prepare_data(pred: np.ndarray, gt: np.ndarray) -> tuple:
+ gt = gt > 128
+ pred = pred / 255
+ if pred.max() != pred.min():
+ pred = (pred - pred.min()) / (pred.max() - pred.min())
+ return pred, gt
+
+
+def _get_adaptive_threshold(matrix: np.ndarray, max_value: float = 1) -> float:
+ return min(2 * matrix.mean(), max_value)
+
+
+class FMeasure(object):
+ def __init__(self, beta: float = 0.3):
+ self.beta = beta
+ self.precisions = []
+ self.recalls = []
+ self.adaptive_fms = []
+ self.changeable_fms = []
+
+ def step(self, pred: np.ndarray, gt: np.ndarray):
+ pred, gt = _prepare_data(pred, gt)
+
+ adaptive_fm = self.cal_adaptive_fm(pred=pred, gt=gt)
+ self.adaptive_fms.append(adaptive_fm)
+
+ precisions, recalls, changeable_fms = self.cal_pr(pred=pred, gt=gt)
+ self.precisions.append(precisions)
+ self.recalls.append(recalls)
+ self.changeable_fms.append(changeable_fms)
+
+ def cal_adaptive_fm(self, pred: np.ndarray, gt: np.ndarray) -> float:
+ adaptive_threshold = _get_adaptive_threshold(pred, max_value=1)
+ binary_predcition = pred >= adaptive_threshold
+ area_intersection = binary_predcition[gt].sum()
+ if area_intersection == 0:
+ adaptive_fm = 0
+ else:
+ pre = area_intersection / np.count_nonzero(binary_predcition)
+ rec = area_intersection / np.count_nonzero(gt)
+ adaptive_fm = (1 + self.beta) * pre * rec / (self.beta * pre + rec)
+ return adaptive_fm
+
+ def cal_pr(self, pred: np.ndarray, gt: np.ndarray) -> tuple:
+ pred = (pred * 255).astype(np.uint8)
+ bins = np.linspace(0, 256, 257)
+ fg_hist, _ = np.histogram(pred[gt], bins=bins)
+ bg_hist, _ = np.histogram(pred[~gt], bins=bins)
+ fg_w_thrs = np.cumsum(np.flip(fg_hist), axis=0)
+ bg_w_thrs = np.cumsum(np.flip(bg_hist), axis=0)
+ TPs = fg_w_thrs
+ Ps = fg_w_thrs + bg_w_thrs
+ Ps[Ps == 0] = 1
+ T = max(np.count_nonzero(gt), 1)
+ precisions = TPs / Ps
+ recalls = TPs / T
+ numerator = (1 + self.beta) * precisions * recalls
+ denominator = np.where(numerator == 0, 1, self.beta * precisions + recalls)
+ changeable_fms = numerator / denominator
+ return precisions, recalls, changeable_fms
+
+ def get_results(self) -> dict:
+ adaptive_fm = np.mean(np.array(self.adaptive_fms, _TYPE))
+ changeable_fm = np.mean(np.array(self.changeable_fms, dtype=_TYPE), axis=0)
+ precision = np.mean(np.array(self.precisions, dtype=_TYPE), axis=0) # N, 256
+ recall = np.mean(np.array(self.recalls, dtype=_TYPE), axis=0) # N, 256
+ return dict(fm=dict(adp=adaptive_fm, curve=changeable_fm),
+ pr=dict(p=precision, r=recall))
+
+
+class MAEMeasure(object):
+ def __init__(self):
+ self.maes = []
+
+ def step(self, pred: np.ndarray, gt: np.ndarray):
+ pred, gt = _prepare_data(pred, gt)
+
+ mae = self.cal_mae(pred, gt)
+ self.maes.append(mae)
+
+ def cal_mae(self, pred: np.ndarray, gt: np.ndarray) -> float:
+ mae = np.mean(np.abs(pred - gt))
+ return mae
+
+ def get_results(self) -> dict:
+ mae = np.mean(np.array(self.maes, _TYPE))
+ return dict(mae=mae)
+
+
+class SMeasure(object):
+ def __init__(self, alpha: float = 0.5):
+ self.sms = []
+ self.alpha = alpha
+
+ def step(self, pred: np.ndarray, gt: np.ndarray):
+ pred, gt = _prepare_data(pred=pred, gt=gt)
+
+ sm = self.cal_sm(pred, gt)
+ self.sms.append(sm)
+
+ def cal_sm(self, pred: np.ndarray, gt: np.ndarray) -> float:
+ y = np.mean(gt)
+ if y == 0:
+ sm = 1 - np.mean(pred)
+ elif y == 1:
+ sm = np.mean(pred)
+ else:
+ sm = self.alpha * self.object(pred, gt) + (1 - self.alpha) * self.region(pred, gt)
+ sm = max(0, sm)
+ return sm
+
+ def object(self, pred: np.ndarray, gt: np.ndarray) -> float:
+ fg = pred * gt
+ bg = (1 - pred) * (1 - gt)
+ u = np.mean(gt)
+ object_score = u * self.s_object(fg, gt) + (1 - u) * self.s_object(bg, 1 - gt)
+ return object_score
+
+ def s_object(self, pred: np.ndarray, gt: np.ndarray) -> float:
+ x = np.mean(pred[gt == 1])
+ sigma_x = np.std(pred[gt == 1], ddof=1)
+ score = 2 * x / (np.power(x, 2) + 1 + sigma_x + _EPS)
+ return score
+
+ def region(self, pred: np.ndarray, gt: np.ndarray) -> float:
+ x, y = self.centroid(gt)
+ part_info = self.divide_with_xy(pred, gt, x, y)
+ w1, w2, w3, w4 = part_info['weight']
+ pred1, pred2, pred3, pred4 = part_info['pred']
+ gt1, gt2, gt3, gt4 = part_info['gt']
+ score1 = self.ssim(pred1, gt1)
+ score2 = self.ssim(pred2, gt2)
+ score3 = self.ssim(pred3, gt3)
+ score4 = self.ssim(pred4, gt4)
+
+ return w1 * score1 + w2 * score2 + w3 * score3 + w4 * score4
+
+ def centroid(self, matrix: np.ndarray) -> tuple:
+ h, w = matrix.shape
+ area_object = np.count_nonzero(matrix)
+ if area_object == 0:
+ x = np.round(w / 2)
+ y = np.round(h / 2)
+ else:
+ # More details can be found at: https://www.yuque.com/lart/blog/gpbigm
+ y, x = np.argwhere(matrix).mean(axis=0).round()
+ return int(x) + 1, int(y) + 1
+
+ def divide_with_xy(self, pred: np.ndarray, gt: np.ndarray, x, y) -> dict:
+ h, w = gt.shape
+ area = h * w
+
+ gt_LT = gt[0:y, 0:x]
+ gt_RT = gt[0:y, x:w]
+ gt_LB = gt[y:h, 0:x]
+ gt_RB = gt[y:h, x:w]
+
+ pred_LT = pred[0:y, 0:x]
+ pred_RT = pred[0:y, x:w]
+ pred_LB = pred[y:h, 0:x]
+ pred_RB = pred[y:h, x:w]
+
+ w1 = x * y / area
+ w2 = y * (w - x) / area
+ w3 = (h - y) * x / area
+ w4 = 1 - w1 - w2 - w3
+
+ return dict(gt=(gt_LT, gt_RT, gt_LB, gt_RB),
+ pred=(pred_LT, pred_RT, pred_LB, pred_RB),
+ weight=(w1, w2, w3, w4))
+
+ def ssim(self, pred: np.ndarray, gt: np.ndarray) -> float:
+ h, w = pred.shape
+ N = h * w
+
+ x = np.mean(pred)
+ y = np.mean(gt)
+
+ sigma_x = np.sum((pred - x) ** 2) / (N - 1)
+ sigma_y = np.sum((gt - y) ** 2) / (N - 1)
+ sigma_xy = np.sum((pred - x) * (gt - y)) / (N - 1)
+
+ alpha = 4 * x * y * sigma_xy
+ beta = (x ** 2 + y ** 2) * (sigma_x + sigma_y)
+
+ if alpha != 0:
+ score = alpha / (beta + _EPS)
+ elif alpha == 0 and beta == 0:
+ score = 1
+ else:
+ score = 0
+ return score
+
+ def get_results(self) -> dict:
+ sm = np.mean(np.array(self.sms, dtype=_TYPE))
+ return dict(sm=sm)
+
+
+class EMeasure(object):
+ def __init__(self):
+ self.adaptive_ems = []
+ self.changeable_ems = []
+
+ def step(self, pred: np.ndarray, gt: np.ndarray):
+ pred, gt = _prepare_data(pred=pred, gt=gt)
+ self.gt_fg_numel = np.count_nonzero(gt)
+ self.gt_size = gt.shape[0] * gt.shape[1]
+
+ changeable_ems = self.cal_changeable_em(pred, gt)
+ self.changeable_ems.append(changeable_ems)
+ adaptive_em = self.cal_adaptive_em(pred, gt)
+ self.adaptive_ems.append(adaptive_em)
+
+ def cal_adaptive_em(self, pred: np.ndarray, gt: np.ndarray) -> float:
+ adaptive_threshold = _get_adaptive_threshold(pred, max_value=1)
+ adaptive_em = self.cal_em_with_threshold(pred, gt, threshold=adaptive_threshold)
+ return adaptive_em
+
+ def cal_changeable_em(self, pred: np.ndarray, gt: np.ndarray) -> np.ndarray:
+ changeable_ems = self.cal_em_with_cumsumhistogram(pred, gt)
+ return changeable_ems
+
+ def cal_em_with_threshold(self, pred: np.ndarray, gt: np.ndarray, threshold: float) -> float:
+ binarized_pred = pred >= threshold
+ fg_fg_numel = np.count_nonzero(binarized_pred & gt)
+ fg_bg_numel = np.count_nonzero(binarized_pred & ~gt)
+
+ fg___numel = fg_fg_numel + fg_bg_numel
+ bg___numel = self.gt_size - fg___numel
+
+ if self.gt_fg_numel == 0:
+ enhanced_matrix_sum = bg___numel
+ elif self.gt_fg_numel == self.gt_size:
+ enhanced_matrix_sum = fg___numel
+ else:
+ parts_numel, combinations = self.generate_parts_numel_combinations(
+ fg_fg_numel=fg_fg_numel, fg_bg_numel=fg_bg_numel,
+ pred_fg_numel=fg___numel, pred_bg_numel=bg___numel,
+ )
+
+ results_parts = []
+ for i, (part_numel, combination) in enumerate(zip(parts_numel, combinations)):
+ align_matrix_value = 2 * (combination[0] * combination[1]) / \
+ (combination[0] ** 2 + combination[1] ** 2 + _EPS)
+ enhanced_matrix_value = (align_matrix_value + 1) ** 2 / 4
+ results_parts.append(enhanced_matrix_value * part_numel)
+ enhanced_matrix_sum = sum(results_parts)
+
+ em = enhanced_matrix_sum / (self.gt_size - 1 + _EPS)
+ return em
+
+ def cal_em_with_cumsumhistogram(self, pred: np.ndarray, gt: np.ndarray) -> np.ndarray:
+ pred = (pred * 255).astype(np.uint8)
+ bins = np.linspace(0, 256, 257)
+ fg_fg_hist, _ = np.histogram(pred[gt], bins=bins)
+ fg_bg_hist, _ = np.histogram(pred[~gt], bins=bins)
+ fg_fg_numel_w_thrs = np.cumsum(np.flip(fg_fg_hist), axis=0)
+ fg_bg_numel_w_thrs = np.cumsum(np.flip(fg_bg_hist), axis=0)
+
+ fg___numel_w_thrs = fg_fg_numel_w_thrs + fg_bg_numel_w_thrs
+ bg___numel_w_thrs = self.gt_size - fg___numel_w_thrs
+
+ if self.gt_fg_numel == 0:
+ enhanced_matrix_sum = bg___numel_w_thrs
+ elif self.gt_fg_numel == self.gt_size:
+ enhanced_matrix_sum = fg___numel_w_thrs
+ else:
+ parts_numel_w_thrs, combinations = self.generate_parts_numel_combinations(
+ fg_fg_numel=fg_fg_numel_w_thrs, fg_bg_numel=fg_bg_numel_w_thrs,
+ pred_fg_numel=fg___numel_w_thrs, pred_bg_numel=bg___numel_w_thrs,
+ )
+
+ results_parts = np.empty(shape=(4, 256), dtype=np.float64)
+ for i, (part_numel, combination) in enumerate(zip(parts_numel_w_thrs, combinations)):
+ align_matrix_value = 2 * (combination[0] * combination[1]) / \
+ (combination[0] ** 2 + combination[1] ** 2 + _EPS)
+ enhanced_matrix_value = (align_matrix_value + 1) ** 2 / 4
+ results_parts[i] = enhanced_matrix_value * part_numel
+ enhanced_matrix_sum = results_parts.sum(axis=0)
+
+ em = enhanced_matrix_sum / (self.gt_size - 1 + _EPS)
+ return em
+
+ def generate_parts_numel_combinations(self, fg_fg_numel, fg_bg_numel, pred_fg_numel, pred_bg_numel):
+ bg_fg_numel = self.gt_fg_numel - fg_fg_numel
+ bg_bg_numel = pred_bg_numel - bg_fg_numel
+
+ parts_numel = [fg_fg_numel, fg_bg_numel, bg_fg_numel, bg_bg_numel]
+
+ mean_pred_value = pred_fg_numel / self.gt_size
+ mean_gt_value = self.gt_fg_numel / self.gt_size
+
+ demeaned_pred_fg_value = 1 - mean_pred_value
+ demeaned_pred_bg_value = 0 - mean_pred_value
+ demeaned_gt_fg_value = 1 - mean_gt_value
+ demeaned_gt_bg_value = 0 - mean_gt_value
+
+ combinations = [
+ (demeaned_pred_fg_value, demeaned_gt_fg_value),
+ (demeaned_pred_fg_value, demeaned_gt_bg_value),
+ (demeaned_pred_bg_value, demeaned_gt_fg_value),
+ (demeaned_pred_bg_value, demeaned_gt_bg_value)
+ ]
+ return parts_numel, combinations
+
+ def get_results(self) -> dict:
+ adaptive_em = np.mean(np.array(self.adaptive_ems, dtype=_TYPE))
+ changeable_em = np.mean(np.array(self.changeable_ems, dtype=_TYPE), axis=0)
+ return dict(em=dict(adp=adaptive_em, curve=changeable_em))
+
+
+class WeightedFMeasure(object):
+ def __init__(self, beta: float = 1):
+ self.beta = beta
+ self.weighted_fms = []
+
+ def step(self, pred: np.ndarray, gt: np.ndarray):
+ pred, gt = _prepare_data(pred=pred, gt=gt)
+
+ if np.all(~gt):
+ wfm = 0
+ else:
+ wfm = self.cal_wfm(pred, gt)
+ self.weighted_fms.append(wfm)
+
+ def cal_wfm(self, pred: np.ndarray, gt: np.ndarray) -> float:
+ # [Dst,IDXT] = bwdist(dGT);
+ Dst, Idxt = bwdist(gt == 0, return_indices=True)
+
+ # %Pixel dependency
+ # E = abs(FG-dGT);
+ E = np.abs(pred - gt)
+ Et = np.copy(E)
+ Et[gt == 0] = Et[Idxt[0][gt == 0], Idxt[1][gt == 0]]
+
+ # K = fspecial('gaussian',7,5);
+ # EA = imfilter(Et,K);
+ K = self.matlab_style_gauss2D((7, 7), sigma=5)
+ EA = convolve(Et, weights=K, mode="constant", cval=0)
+ # MIN_E_EA = E;
+ # MIN_E_EA(GT & EA np.ndarray:
+ """
+ 2D gaussian mask - should give the same result as MATLAB's
+ fspecial('gaussian',[shape],[sigma])
+ """
+ m, n = [(ss - 1) / 2 for ss in shape]
+ y, x = np.ogrid[-m: m + 1, -n: n + 1]
+ h = np.exp(-(x * x + y * y) / (2 * sigma * sigma))
+ h[h < np.finfo(h.dtype).eps * h.max()] = 0
+ sumh = h.sum()
+ if sumh != 0:
+ h /= sumh
+ return h
+
+ def get_results(self) -> dict:
+ weighted_fm = np.mean(np.array(self.weighted_fms, dtype=_TYPE))
+ return dict(wfm=weighted_fm)
+
+
+class HCEMeasure(object):
+ def __init__(self):
+ self.hces = []
+
+ def step(self, pred: np.ndarray, gt: np.ndarray, gt_ske):
+ # pred, gt = _prepare_data(pred, gt)
+
+ hce = self.cal_hce(pred, gt, gt_ske)
+ self.hces.append(hce)
+
+ def get_results(self) -> dict:
+ hce = np.mean(np.array(self.hces, _TYPE))
+ return dict(hce=hce)
+
+
+ def cal_hce(self, pred: np.ndarray, gt: np.ndarray, gt_ske: np.ndarray, relax=5, epsilon=2.0) -> float:
+ # Binarize gt
+ if(len(gt.shape)>2):
+ gt = gt[:, :, 0]
+
+ epsilon_gt = 128#(np.amin(gt)+np.amax(gt))/2.0
+ gt = (gt>epsilon_gt).astype(np.uint8)
+
+ # Binarize pred
+ if(len(pred.shape)>2):
+ pred = pred[:, :, 0]
+ epsilon_pred = 128#(np.amin(pred)+np.amax(pred))/2.0
+ pred = (pred>epsilon_pred).astype(np.uint8)
+
+ Union = np.logical_or(gt, pred)
+ TP = np.logical_and(gt, pred)
+ FP = pred - TP
+ FN = gt - TP
+
+ # relax the Union of gt and pred
+ Union_erode = Union.copy()
+ Union_erode = cv2.erode(Union_erode.astype(np.uint8), disk(1), iterations=relax)
+
+ # --- get the relaxed False Positive regions for computing the human efforts in correcting them ---
+ FP_ = np.logical_and(FP, Union_erode) # get the relaxed FP
+ for i in range(0, relax):
+ FP_ = cv2.dilate(FP_.astype(np.uint8), disk(1))
+ FP_ = np.logical_and(FP_, 1-np.logical_or(TP, FN))
+ FP_ = np.logical_and(FP, FP_)
+
+ # --- get the relaxed False Negative regions for computing the human efforts in correcting them ---
+ FN_ = np.logical_and(FN, Union_erode) # preserve the structural components of FN
+ ## recover the FN, where pixels are not close to the TP borders
+ for i in range(0, relax):
+ FN_ = cv2.dilate(FN_.astype(np.uint8), disk(1))
+ FN_ = np.logical_and(FN_, 1-np.logical_or(TP, FP))
+ FN_ = np.logical_and(FN, FN_)
+ FN_ = np.logical_or(FN_, np.logical_xor(gt_ske, np.logical_and(TP, gt_ske))) # preserve the structural components of FN
+
+ ## 2. =============Find exact polygon control points and independent regions==============
+ ## find contours from FP_
+ ctrs_FP, hier_FP = cv2.findContours(FP_.astype(np.uint8), cv2.RETR_TREE, cv2.CHAIN_APPROX_NONE)
+ ## find control points and independent regions for human correction
+ bdies_FP, indep_cnt_FP = self.filter_bdy_cond(ctrs_FP, FP_, np.logical_or(TP,FN_))
+ ## find contours from FN_
+ ctrs_FN, hier_FN = cv2.findContours(FN_.astype(np.uint8), cv2.RETR_TREE, cv2.CHAIN_APPROX_NONE)
+ ## find control points and independent regions for human correction
+ bdies_FN, indep_cnt_FN = self.filter_bdy_cond(ctrs_FN, FN_, 1-np.logical_or(np.logical_or(TP, FP_), FN_))
+
+ poly_FP, poly_FP_len, poly_FP_point_cnt = self.approximate_RDP(bdies_FP, epsilon=epsilon)
+ poly_FN, poly_FN_len, poly_FN_point_cnt = self.approximate_RDP(bdies_FN, epsilon=epsilon)
+
+ # FP_points+FP_indep+FN_points+FN_indep
+ return poly_FP_point_cnt+indep_cnt_FP+poly_FN_point_cnt+indep_cnt_FN
+
+ def filter_bdy_cond(self, bdy_, mask, cond):
+
+ cond = cv2.dilate(cond.astype(np.uint8), disk(1))
+ labels = label(mask) # find the connected regions
+ lbls = np.unique(labels) # the indices of the connected regions
+ indep = np.ones(lbls.shape[0]) # the label of each connected regions
+ indep[0] = 0 # 0 indicate the background region
+
+ boundaries = []
+ h,w = cond.shape[0:2]
+ ind_map = np.zeros((h, w))
+ indep_cnt = 0
+
+ for i in range(0, len(bdy_)):
+ tmp_bdies = []
+ tmp_bdy = []
+ for j in range(0, bdy_[i].shape[0]):
+ r, c = bdy_[i][j,0,1],bdy_[i][j,0,0]
+
+ if(np.sum(cond[r, c])==0 or ind_map[r, c]!=0):
+ if(len(tmp_bdy)>0):
+ tmp_bdies.append(tmp_bdy)
+ tmp_bdy = []
+ continue
+ tmp_bdy.append([c, r])
+ ind_map[r, c] = ind_map[r, c] + 1
+ indep[labels[r, c]] = 0 # indicates part of the boundary of this region needs human correction
+ if(len(tmp_bdy)>0):
+ tmp_bdies.append(tmp_bdy)
+
+ # check if the first and the last boundaries are connected
+ # if yes, invert the first boundary and attach it after the last boundary
+ if(len(tmp_bdies)>1):
+ first_x, first_y = tmp_bdies[0][0]
+ last_x, last_y = tmp_bdies[-1][-1]
+ if((abs(first_x-last_x)==1 and first_y==last_y) or
+ (first_x==last_x and abs(first_y-last_y)==1) or
+ (abs(first_x-last_x)==1 and abs(first_y-last_y)==1)
+ ):
+ tmp_bdies[-1].extend(tmp_bdies[0][::-1])
+ del tmp_bdies[0]
+
+ for k in range(0, len(tmp_bdies)):
+ tmp_bdies[k] = np.array(tmp_bdies[k])[:, np.newaxis, :]
+ if(len(tmp_bdies)>0):
+ boundaries.extend(tmp_bdies)
+
+ return boundaries, np.sum(indep)
+
+ # this function approximate each boundary by DP algorithm
+ # https://en.wikipedia.org/wiki/Ramer%E2%80%93Douglas%E2%80%93Peucker_algorithm
+ def approximate_RDP(self, boundaries, epsilon=1.0):
+
+ boundaries_ = []
+ boundaries_len_ = []
+ pixel_cnt_ = 0
+
+ # polygon approximate of each boundary
+ for i in range(0, len(boundaries)):
+ boundaries_.append(cv2.approxPolyDP(boundaries[i], epsilon, False))
+
+ # count the control points number of each boundary and the total control points number of all the boundaries
+ for i in range(0, len(boundaries_)):
+ boundaries_len_.append(len(boundaries_[i]))
+ pixel_cnt_ = pixel_cnt_ + len(boundaries_[i])
+
+ return boundaries_, boundaries_len_, pixel_cnt_
+
+
+class MBAMeasure(object):
+ def __init__(self):
+ self.bas = []
+ self.all_h = 0
+ self.all_w = 0
+ self.all_max = 0
+
+ def step(self, pred: np.ndarray, gt: np.ndarray):
+ # pred, gt = _prepare_data(pred, gt)
+
+ refined = gt.copy()
+
+ rmin = cmin = 0
+ rmax, cmax = gt.shape
+
+ self.all_h += rmax
+ self.all_w += cmax
+ self.all_max += max(rmax, cmax)
+
+ refined_h, refined_w = refined.shape
+ if refined_h != cmax:
+ refined = np.array(Image.fromarray(pred).resize((cmax, rmax), Image.BILINEAR))
+
+ if not(gt.sum() < 32*32):
+ if not((cmax==cmin) or (rmax==rmin)):
+ class_refined_prob = np.array(Image.fromarray(pred).resize((cmax-cmin, rmax-rmin), Image.BILINEAR))
+ refined[rmin:rmax, cmin:cmax] = class_refined_prob
+
+ pred = pred > 128
+ gt = gt > 128
+
+ ba = self.cal_ba(pred, gt)
+ self.bas.append(ba)
+
+ def get_disk_kernel(self, radius):
+ return cv2.getStructuringElement(cv2.MORPH_ELLIPSE, (radius*2+1, radius*2+1))
+
+ def cal_ba(self, pred: np.ndarray, gt: np.ndarray) -> np.ndarray:
+ """
+ Calculate the mean absolute error.
+
+ :return: ba
+ """
+
+ gt = gt.astype(np.uint8)
+ pred = pred.astype(np.uint8)
+
+ h, w = gt.shape
+
+ min_radius = 1
+ max_radius = (w+h)/300
+ num_steps = 5
+
+ pred_acc = [None] * num_steps
+
+ for i in range(num_steps):
+ curr_radius = min_radius + int((max_radius-min_radius)/num_steps*i)
+
+ kernel = self.get_disk_kernel(curr_radius)
+ boundary_region = cv2.morphologyEx(gt, cv2.MORPH_GRADIENT, kernel) > 0
+
+ gt_in_bound = gt[boundary_region]
+ pred_in_bound = pred[boundary_region]
+
+ num_edge_pixels = (boundary_region).sum()
+ num_pred_gd_pix = ((gt_in_bound) * (pred_in_bound) + (1-gt_in_bound) * (1-pred_in_bound)).sum()
+
+ pred_acc[i] = num_pred_gd_pix / num_edge_pixels
+
+ ba = sum(pred_acc)/num_steps
+ return ba
+
+ def get_results(self) -> dict:
+ mba = np.mean(np.array(self.bas, _TYPE))
+ return dict(mba=mba)
+
+
+class BIoUMeasure(object):
+ def __init__(self, dilation_ratio=0.02):
+ self.bious = []
+ self.dilation_ratio = dilation_ratio
+
+ def mask_to_boundary(self, mask):
+ h, w = mask.shape
+ img_diag = np.sqrt(h ** 2 + w ** 2)
+ dilation = int(round(self.dilation_ratio * img_diag))
+ if dilation < 1:
+ dilation = 1
+ # Pad image so mask truncated by the image border is also considered as boundary.
+ new_mask = cv2.copyMakeBorder(mask, 1, 1, 1, 1, cv2.BORDER_CONSTANT, value=0)
+ kernel = np.ones((3, 3), dtype=np.uint8)
+ new_mask_erode = cv2.erode(new_mask, kernel, iterations=dilation)
+ mask_erode = new_mask_erode[1 : h + 1, 1 : w + 1]
+ # G_d intersects G in the paper.
+ return mask - mask_erode
+
+ def step(self, pred: np.ndarray, gt: np.ndarray):
+ pred, gt = _prepare_data(pred, gt)
+
+ bious = self.cal_biou(pred=pred, gt=gt)
+ self.bious.append(bious)
+
+ def cal_biou(self, pred, gt):
+ pred = (pred * 255).astype(np.uint8)
+ pred = self.mask_to_boundary(pred)
+ gt = (gt * 255).astype(np.uint8)
+ gt = self.mask_to_boundary(gt)
+ gt = gt > 128
+
+ bins = np.linspace(0, 256, 257)
+ fg_hist, _ = np.histogram(pred[gt], bins=bins) # ture positive
+ bg_hist, _ = np.histogram(pred[~gt], bins=bins) # false positive
+ fg_w_thrs = np.cumsum(np.flip(fg_hist), axis=0)
+ bg_w_thrs = np.cumsum(np.flip(bg_hist), axis=0)
+ TPs = fg_w_thrs
+ Ps = fg_w_thrs + bg_w_thrs # positives
+ Ps[Ps == 0] = 1
+ T = max(np.count_nonzero(gt), 1)
+
+ ious = TPs / (T + bg_w_thrs)
+ return ious
+
+ def get_results(self) -> dict:
+ biou = np.mean(np.array(self.bious, dtype=_TYPE), axis=0)
+ return dict(biou=dict(curve=biou))
diff --git a/py/BiRefNet_v2/gen_best_ep.py b/py/BiRefNet_v2/gen_best_ep.py
new file mode 100644
index 0000000..8e59868
--- /dev/null
+++ b/py/BiRefNet_v2/gen_best_ep.py
@@ -0,0 +1,86 @@
+import os
+from glob import glob
+import numpy as np
+
+from .config import Config
+
+
+config = Config()
+
+eval_txts = sorted(glob('e_results/*_eval.txt'))
+print('eval_txts:', [_.split(os.sep)[-1] for _ in eval_txts])
+score_panel = {}
+sep = '&'
+metrics = ['sm', 'wfm', 'hce'] # we used HCE for DIS and wFm for others.
+if 'DIS5K' not in config.task:
+ metrics.remove('hce')
+
+for metric in metrics:
+ print('Metric:', metric)
+ current_line_nums = []
+ for idx_et, eval_txt in enumerate(eval_txts):
+ with open(eval_txt, 'r') as f:
+ lines = [l for l in f.readlines()[3:] if '.' in l]
+ current_line_nums.append(len(lines))
+ for idx_et, eval_txt in enumerate(eval_txts):
+ with open(eval_txt, 'r') as f:
+ lines = [l for l in f.readlines()[3:] if '.' in l]
+ for idx_line, line in enumerate(lines[:min(current_line_nums)]): # Consist line numbers by the minimal result file.
+ properties = line.strip().strip(sep).split(sep)
+ dataset = properties[0].strip()
+ ckpt = properties[1].strip()
+ if int(ckpt.split('--epoch_')[-1].strip()) < 0:
+ continue
+ targe_idx = {
+ 'sm': [5, 2, 2, 5, 2],
+ 'wfm': [3, 3, 8, 3, 8],
+ 'hce': [7, -1, -1, 7, -1]
+ }[metric][['DIS5K', 'COD', 'HRSOD', 'General', 'Matting'].index(config.task)]
+ if metric != 'hce':
+ score_sm = float(properties[targe_idx].strip())
+ else:
+ score_sm = int(properties[targe_idx].strip().strip('.'))
+ if idx_et == 0:
+ score_panel[ckpt] = []
+ score_panel[ckpt].append(score_sm)
+
+ metrics_min = ['hce', 'mae']
+ max_or_min = min if metric in metrics_min else max
+ score_max = max_or_min(score_panel.values(), key=lambda x: np.sum(x))
+
+ good_models = []
+ for k, v in score_panel.items():
+ if (np.sum(v) <= np.sum(score_max)) if metric in metrics_min else (np.sum(v) >= np.sum(score_max)):
+ print(k, v)
+ good_models.append(k)
+
+ # Write
+ with open(eval_txt, 'r') as f:
+ lines = f.readlines()
+ info4good_models = lines[:3]
+ metric_names = [m.strip() for m in lines[1].strip().strip('&').split('&')[2:]]
+ testset_mean_values = {metric_name: [] for metric_name in metric_names}
+ for good_model in good_models:
+ for idx_et, eval_txt in enumerate(eval_txts):
+ with open(eval_txt, 'r') as f:
+ lines = f.readlines()
+ for line in lines:
+ if set([good_model]) & set([_.strip() for _ in line.split(sep)]):
+ info4good_models.append(line)
+ metric_scores = [float(m.strip()) for m in line.strip().strip('&').split('&')[2:]]
+ for idx_score, metric_score in enumerate(metric_scores):
+ testset_mean_values[metric_names[idx_score]].append(metric_score)
+
+ if 'DIS5K' in config.task:
+ testset_mean_values_lst = ['{:<4}'.format(int(np.mean(v_lst[:-1]).round())) if name == 'HCE' else '{:.3f}'.format(np.mean(v_lst[:-1])).lstrip('0') for name, v_lst in testset_mean_values.items()] # [:-1] to remove DIS-VD
+ sample_line_for_placing_mean_values = info4good_models[-2]
+ numbers_placed_well = sample_line_for_placing_mean_values.replace(sample_line_for_placing_mean_values.split('&')[1].strip(), 'DIS-TEs').strip().split('&')[3:]
+ for idx_number, (number_placed_well, testset_mean_value) in enumerate(zip(numbers_placed_well, testset_mean_values_lst)):
+ numbers_placed_well[idx_number] = number_placed_well.replace(number_placed_well.strip(), testset_mean_value)
+ testset_mean_line = '&'.join(sample_line_for_placing_mean_values.replace(sample_line_for_placing_mean_values.split('&')[1].strip(), 'DIS-TEs').split('&')[:3] + numbers_placed_well) + '\n'
+ info4good_models.append(testset_mean_line)
+ info4good_models.append(lines[-1])
+ info = ''.join(info4good_models)
+ print(info)
+ with open(os.path.join('e_results', 'eval-{}_best_on_{}.txt'.format(config.task, metric)), 'w') as f:
+ f.write(info + '\n')
diff --git a/py/BiRefNet_v2/image_proc.py b/py/BiRefNet_v2/image_proc.py
new file mode 100644
index 0000000..2ebfbfa
--- /dev/null
+++ b/py/BiRefNet_v2/image_proc.py
@@ -0,0 +1,119 @@
+import random
+from PIL import Image, ImageEnhance
+import numpy as np
+import cv2
+
+
+def refine_foreground(image, mask, r=90):
+ if mask.size != image.size:
+ mask = mask.resize(image.size)
+ image = np.array(image) / 255.0
+ mask = np.array(mask) / 255.0
+ estimated_foreground = FB_blur_fusion_foreground_estimator_2(image, mask, r=r)
+ image_masked = Image.fromarray((estimated_foreground * 255.0).astype(np.uint8))
+ return image_masked
+
+
+def FB_blur_fusion_foreground_estimator_2(image, alpha, r=90):
+ # Thanks to the source: https://github.com/Photoroom/fast-foreground-estimation
+ alpha = alpha[:, :, None]
+ F, blur_B = FB_blur_fusion_foreground_estimator(
+ image, image, image, alpha, r)
+ return FB_blur_fusion_foreground_estimator(image, F, blur_B, alpha, r=6)[0]
+
+
+def FB_blur_fusion_foreground_estimator(image, F, B, alpha, r=90):
+ if isinstance(image, Image.Image):
+ image = np.array(image) / 255.0
+ blurred_alpha = cv2.blur(alpha, (r, r))[:, :, None]
+
+ blurred_FA = cv2.blur(F * alpha, (r, r))
+ blurred_F = blurred_FA / (blurred_alpha + 1e-5)
+
+ blurred_B1A = cv2.blur(B * (1 - alpha), (r, r))
+ blurred_B = blurred_B1A / ((1 - blurred_alpha) + 1e-5)
+ F = blurred_F + alpha * \
+ (image - alpha * blurred_F - (1 - alpha) * blurred_B)
+ F = np.clip(F, 0, 1)
+ return F, blurred_B
+
+
+def preproc(image, label, preproc_methods=['flip']):
+ if 'flip' in preproc_methods:
+ image, label = cv_random_flip(image, label)
+ if 'crop' in preproc_methods:
+ image, label = random_crop(image, label)
+ if 'rotate' in preproc_methods:
+ image, label = random_rotate(image, label)
+ if 'enhance' in preproc_methods:
+ image = color_enhance(image)
+ if 'pepper' in preproc_methods:
+ label = random_pepper(label)
+ return image, label
+
+
+def cv_random_flip(img, label):
+ if random.random() > 0.5:
+ img = img.transpose(Image.FLIP_LEFT_RIGHT)
+ label = label.transpose(Image.FLIP_LEFT_RIGHT)
+ return img, label
+
+
+def random_crop(image, label):
+ border = 30
+ image_width = image.size[0]
+ image_height = image.size[1]
+ border = int(min(image_width, image_height) * 0.1)
+ crop_win_width = np.random.randint(image_width - border, image_width)
+ crop_win_height = np.random.randint(image_height - border, image_height)
+ random_region = (
+ (image_width - crop_win_width) >> 1, (image_height - crop_win_height) >> 1, (image_width + crop_win_width) >> 1,
+ (image_height + crop_win_height) >> 1)
+ return image.crop(random_region), label.crop(random_region)
+
+
+def random_rotate(image, label, angle=15):
+ mode = Image.BICUBIC
+ if random.random() > 0.8:
+ random_angle = np.random.randint(-angle, angle)
+ image = image.rotate(random_angle, mode)
+ label = label.rotate(random_angle, mode)
+ return image, label
+
+
+def color_enhance(image):
+ bright_intensity = random.randint(5, 15) / 10.0
+ image = ImageEnhance.Brightness(image).enhance(bright_intensity)
+ contrast_intensity = random.randint(5, 15) / 10.0
+ image = ImageEnhance.Contrast(image).enhance(contrast_intensity)
+ color_intensity = random.randint(0, 20) / 10.0
+ image = ImageEnhance.Color(image).enhance(color_intensity)
+ sharp_intensity = random.randint(0, 30) / 10.0
+ image = ImageEnhance.Sharpness(image).enhance(sharp_intensity)
+ return image
+
+
+def random_gaussian(image, mean=0.1, sigma=0.35):
+ def gaussianNoisy(im, mean=mean, sigma=sigma):
+ for _i in range(len(im)):
+ im[_i] += random.gauss(mean, sigma)
+ return im
+
+ img = np.asarray(image)
+ width, height = img.shape
+ img = gaussianNoisy(img[:].flatten(), mean, sigma)
+ img = img.reshape([width, height])
+ return Image.fromarray(np.uint8(img))
+
+
+def random_pepper(img, N=0.0015):
+ img = np.array(img)
+ noiseNum = int(N * img.shape[0] * img.shape[1])
+ for i in range(noiseNum):
+ randX = random.randint(0, img.shape[0] - 1)
+ randY = random.randint(0, img.shape[1] - 1)
+ if random.randint(0, 1) == 0:
+ img[randX, randY] = 0
+ else:
+ img[randX, randY] = 255
+ return Image.fromarray(img)
diff --git a/py/BiRefNet_v2/inference.py b/py/BiRefNet_v2/inference.py
new file mode 100644
index 0000000..21ed88f
--- /dev/null
+++ b/py/BiRefNet_v2/inference.py
@@ -0,0 +1,105 @@
+import os
+import argparse
+from glob import glob
+from tqdm import tqdm
+import cv2
+import torch
+
+from .dataset import MyData
+from .models.birefnet import BiRefNet
+from .utils import save_tensor_img, check_state_dict
+from .config import Config
+
+
+config = Config()
+
+
+def inference(model, data_loader_test, pred_root, method, testset, device=0):
+ model_training = model.training
+ if model_training:
+ model.eval()
+ for batch in tqdm(data_loader_test, total=len(data_loader_test)) if 1 or config.verbose_eval else data_loader_test:
+ inputs = batch[0].to(device)
+ # gts = batch[1].to(device)
+ label_paths = batch[-1]
+ with torch.no_grad():
+ scaled_preds = model(inputs)[-1].sigmoid()
+
+ os.makedirs(os.path.join(pred_root, method, testset), exist_ok=True)
+
+ for idx_sample in range(scaled_preds.shape[0]):
+ res = torch.nn.functional.interpolate(
+ scaled_preds[idx_sample].unsqueeze(0),
+ size=cv2.imread(label_paths[idx_sample], cv2.IMREAD_GRAYSCALE).shape[:2],
+ mode='bilinear',
+ align_corners=True
+ )
+ save_tensor_img(res, os.path.join(os.path.join(pred_root, method, testset), label_paths[idx_sample].replace('\\', '/').split('/')[-1])) # test set dir + file name
+ if model_training:
+ model.train()
+ return None
+
+
+def main(args):
+ # Init model
+
+ device = config.device
+ if args.ckpt_folder:
+ print('Testing with models in {}'.format(args.ckpt_folder))
+ else:
+ print('Testing with model {}'.format(args.ckpt))
+
+ if config.model == 'BiRefNet':
+ model = BiRefNet(bb_pretrained=False)
+ weights_lst = sorted(
+ glob(os.path.join(args.ckpt_folder, '*.pth')) if args.ckpt_folder else [args.ckpt],
+ key=lambda x: int(x.split('epoch_')[-1].split('.pth')[0]),
+ reverse=True
+ )
+ for testset in args.testsets.split('+'):
+ print('>>>> Testset: {}...'.format(testset))
+ data_loader_test = torch.utils.data.DataLoader(
+ dataset=MyData(testset, image_size=config.size, is_train=False),
+ batch_size=config.batch_size_valid, shuffle=False, num_workers=config.num_workers, pin_memory=True
+ )
+ for weights in weights_lst:
+ if int(weights.strip('.pth').split('epoch_')[-1]) % 1 != 0:
+ continue
+ print('\tInferencing {}...'.format(weights))
+ # model.load_state_dict(torch.load(weights, map_location='cpu'))
+ state_dict = torch.load(weights, map_location='cpu')
+ state_dict = check_state_dict(state_dict)
+ model.load_state_dict(state_dict)
+ model = model.to(device)
+ inference(
+ model, data_loader_test=data_loader_test, pred_root=args.pred_root,
+ method='--'.join([w.rstrip('.pth') for w in weights.split(os.sep)[-2:]]),
+ testset=testset, device=config.device
+ )
+
+
+if __name__ == '__main__':
+ # Parameter from command line
+ parser = argparse.ArgumentParser(description='')
+ parser.add_argument('--ckpt', type=str, help='model folder')
+ parser.add_argument('--ckpt_folder', default=sorted(glob(os.path.join('ckpt', '*')))[-1], type=str, help='model folder')
+ parser.add_argument('--pred_root', default='e_preds', type=str, help='Output folder')
+ parser.add_argument('--testsets',
+ default={
+ 'DIS5K': 'DIS-VD+DIS-TE1+DIS-TE2+DIS-TE3+DIS-TE4',
+ 'COD': 'TE-COD10K+NC4K+TE-CAMO+CHAMELEON',
+ 'HRSOD': 'DAVIS-S+TE-HRSOD+TE-UHRSD+TE-DUTS+DUT-OMRON',
+ 'General': 'DIS-VD',
+ 'Matting': 'TE-P3M-500-P',
+ 'DIS5K-': 'DIS-VD',
+ 'COD-': 'TE-COD10K',
+ 'SOD-': 'DAVIS-S+TE-HRSOD+TE-UHRSD',
+ }[config.task + ''],
+ type=str,
+ help="Test all sets: , 'DIS-VD+DIS-TE1+DIS-TE2+DIS-TE3+DIS-TE4'")
+
+ args = parser.parse_args()
+
+ if config.precisionHigh:
+ torch.set_float32_matmul_precision('high')
+ main(args)
diff --git a/py/BiRefNet_v2/loss.py b/py/BiRefNet_v2/loss.py
new file mode 100644
index 0000000..ee0c5a2
--- /dev/null
+++ b/py/BiRefNet_v2/loss.py
@@ -0,0 +1,277 @@
+import torch
+from torch import nn
+import torch.nn.functional as F
+from torch.autograd import Variable
+from math import exp
+
+from .config import Config
+
+
+class Discriminator(nn.Module):
+ def __init__(self, channels=1, img_size=256):
+ super(Discriminator, self).__init__()
+
+ def discriminator_block(in_filters, out_filters, bn=Config().batch_size > 1):
+ block = [nn.Conv2d(in_filters, out_filters, 3, 2, 1), nn.LeakyReLU(0.2, inplace=True), nn.Dropout2d(0.25)]
+ if bn:
+ block.append(nn.BatchNorm2d(out_filters, 0.8))
+ return block
+
+ self.model = nn.Sequential(
+ *discriminator_block(channels, 16, bn=False),
+ *discriminator_block(16, 32),
+ *discriminator_block(32, 64),
+ *discriminator_block(64, 128),
+ )
+
+ # The height and width of downsampled image
+ ds_size = img_size // 2 ** 4
+ self.adv_layer = nn.Sequential(nn.Linear(128 * ds_size ** 2, 1), nn.Sigmoid())
+
+ def forward(self, img):
+ out = self.model(img)
+ out = out.view(out.shape[0], -1)
+ validity = self.adv_layer(out)
+
+ return validity
+
+
+class ContourLoss(torch.nn.Module):
+ def __init__(self):
+ super(ContourLoss, self).__init__()
+
+ def forward(self, pred, target, weight=10):
+ '''
+ target, pred: tensor of shape (B, C, H, W), where target[:,:,region_in_contour] == 1,
+ target[:,:,region_out_contour] == 0.
+ weight: scalar, length term weight.
+ '''
+ # length term
+ delta_r = pred[:,:,1:,:] - pred[:,:,:-1,:] # horizontal gradient (B, C, H-1, W)
+ delta_c = pred[:,:,:,1:] - pred[:,:,:,:-1] # vertical gradient (B, C, H, W-1)
+
+ delta_r = delta_r[:,:,1:,:-2]**2 # (B, C, H-2, W-2)
+ delta_c = delta_c[:,:,:-2,1:]**2 # (B, C, H-2, W-2)
+ delta_pred = torch.abs(delta_r + delta_c)
+
+ epsilon = 1e-8 # where is a parameter to avoid square root is zero in practice.
+ length = torch.mean(torch.sqrt(delta_pred + epsilon)) # eq.(11) in the paper, mean is used instead of sum.
+
+ c_in = torch.ones_like(pred)
+ c_out = torch.zeros_like(pred)
+
+ region_in = torch.mean( pred * (target - c_in )**2 ) # equ.(12) in the paper, mean is used instead of sum.
+ region_out = torch.mean( (1-pred) * (target - c_out)**2 )
+ region = region_in + region_out
+
+ loss = weight * length + region
+
+ return loss
+
+
+class IoULoss(torch.nn.Module):
+ def __init__(self):
+ super(IoULoss, self).__init__()
+
+ def forward(self, pred, target):
+ b = pred.shape[0]
+ IoU = 0.0
+ for i in range(0, b):
+ # compute the IoU of the foreground
+ Iand1 = torch.sum(target[i, :, :, :] * pred[i, :, :, :])
+ Ior1 = torch.sum(target[i, :, :, :]) + torch.sum(pred[i, :, :, :]) - Iand1
+ IoU1 = Iand1 / Ior1
+ # IoU loss is (1-IoU1)
+ IoU = IoU + (1-IoU1)
+ # return IoU/b
+ return IoU
+
+
+class StructureLoss(torch.nn.Module):
+ def __init__(self):
+ super(StructureLoss, self).__init__()
+
+ def forward(self, pred, target):
+ weit = 1+5*torch.abs(F.avg_pool2d(target, kernel_size=31, stride=1, padding=15)-target)
+ wbce = F.binary_cross_entropy_with_logits(pred, target, reduction='none')
+ wbce = (weit*wbce).sum(dim=(2,3))/weit.sum(dim=(2,3))
+
+ pred = torch.sigmoid(pred)
+ inter = ((pred * target) * weit).sum(dim=(2, 3))
+ union = ((pred + target) * weit).sum(dim=(2, 3))
+ wiou = 1-(inter+1)/(union-inter+1)
+
+ return (wbce+wiou).mean()
+
+
+class PatchIoULoss(torch.nn.Module):
+ def __init__(self):
+ super(PatchIoULoss, self).__init__()
+ self.iou_loss = IoULoss()
+
+ def forward(self, pred, target):
+ win_y, win_x = 64, 64
+ iou_loss = 0.
+ for anchor_y in range(0, target.shape[0], win_y):
+ for anchor_x in range(0, target.shape[1], win_y):
+ patch_pred = pred[:, :, anchor_y:anchor_y+win_y, anchor_x:anchor_x+win_x]
+ patch_target = target[:, :, anchor_y:anchor_y+win_y, anchor_x:anchor_x+win_x]
+ patch_iou_loss = self.iou_loss(patch_pred, patch_target)
+ iou_loss += patch_iou_loss
+ return iou_loss
+
+
+class ThrReg_loss(torch.nn.Module):
+ def __init__(self):
+ super(ThrReg_loss, self).__init__()
+
+ def forward(self, pred, gt=None):
+ return torch.mean(1 - ((pred - 0) ** 2 + (pred - 1) ** 2))
+
+
+class ClsLoss(nn.Module):
+ """
+ Auxiliary classification loss for each refined class output.
+ """
+ def __init__(self):
+ super(ClsLoss, self).__init__()
+ self.config = Config()
+ self.lambdas_cls = self.config.lambdas_cls
+
+ self.criterions_last = {
+ 'ce': nn.CrossEntropyLoss()
+ }
+
+ def forward(self, preds, gt):
+ loss = 0.
+ for _, pred_lvl in enumerate(preds):
+ if pred_lvl is None:
+ continue
+ for criterion_name, criterion in self.criterions_last.items():
+ loss += criterion(pred_lvl, gt) * self.lambdas_cls[criterion_name]
+ return loss
+
+
+class PixLoss(nn.Module):
+ """
+ Pixel loss for each refined map output.
+ """
+ def __init__(self):
+ super(PixLoss, self).__init__()
+ self.config = Config()
+ self.lambdas_pix_last = self.config.lambdas_pix_last
+
+ self.criterions_last = {}
+ if 'bce' in self.lambdas_pix_last and self.lambdas_pix_last['bce']:
+ self.criterions_last['bce'] = nn.BCELoss() if not self.config.use_fp16 else nn.BCEWithLogitsLoss()
+ if 'iou' in self.lambdas_pix_last and self.lambdas_pix_last['iou']:
+ self.criterions_last['iou'] = IoULoss()
+ if 'iou_patch' in self.lambdas_pix_last and self.lambdas_pix_last['iou_patch']:
+ self.criterions_last['iou_patch'] = PatchIoULoss()
+ if 'ssim' in self.lambdas_pix_last and self.lambdas_pix_last['ssim']:
+ self.criterions_last['ssim'] = SSIMLoss()
+ if 'mae' in self.lambdas_pix_last and self.lambdas_pix_last['mae']:
+ self.criterions_last['mae'] = nn.L1Loss()
+ if 'mse' in self.lambdas_pix_last and self.lambdas_pix_last['mse']:
+ self.criterions_last['mse'] = nn.MSELoss()
+ if 'reg' in self.lambdas_pix_last and self.lambdas_pix_last['reg']:
+ self.criterions_last['reg'] = ThrReg_loss()
+ if 'cnt' in self.lambdas_pix_last and self.lambdas_pix_last['cnt']:
+ self.criterions_last['cnt'] = ContourLoss()
+ if 'structure' in self.lambdas_pix_last and self.lambdas_pix_last['structure']:
+ self.criterions_last['structure'] = StructureLoss()
+
+ def forward(self, scaled_preds, gt):
+ loss = 0.
+ criterions_embedded_with_sigmoid = ['structure', ] + ['bce'] if self.config.use_fp16 else []
+ for _, pred_lvl in enumerate(scaled_preds):
+ if pred_lvl.shape != gt.shape:
+ pred_lvl = nn.functional.interpolate(pred_lvl, size=gt.shape[2:], mode='bilinear', align_corners=True)
+ for criterion_name, criterion in self.criterions_last.items():
+ _loss = criterion(pred_lvl.sigmoid() if criterion_name not in criterions_embedded_with_sigmoid else pred_lvl, gt) * self.lambdas_pix_last[criterion_name]
+ loss += _loss
+ # print(criterion_name, _loss.item())
+ return loss
+
+
+class SSIMLoss(torch.nn.Module):
+ def __init__(self, window_size=11, size_average=True):
+ super(SSIMLoss, self).__init__()
+ self.window_size = window_size
+ self.size_average = size_average
+ self.channel = 1
+ self.window = create_window(window_size, self.channel)
+
+ def forward(self, img1, img2):
+ (_, channel, _, _) = img1.size()
+ if channel == self.channel and self.window.data.type() == img1.data.type():
+ window = self.window
+ else:
+ window = create_window(self.window_size, channel)
+ if img1.is_cuda:
+ window = window.cuda(img1.get_device())
+ window = window.type_as(img1)
+ self.window = window
+ self.channel = channel
+ return 1 - _ssim(img1, img2, window, self.window_size, channel, self.size_average)
+
+
+def gaussian(window_size, sigma):
+ gauss = torch.Tensor([exp(-(x - window_size//2)**2/float(2*sigma**2)) for x in range(window_size)])
+ return gauss/gauss.sum()
+
+
+def create_window(window_size, channel):
+ _1D_window = gaussian(window_size, 1.5).unsqueeze(1)
+ _2D_window = _1D_window.mm(_1D_window.t()).float().unsqueeze(0).unsqueeze(0)
+ window = Variable(_2D_window.expand(channel, 1, window_size, window_size).contiguous())
+ return window
+
+
+def _ssim(img1, img2, window, window_size, channel, size_average=True):
+ mu1 = F.conv2d(img1, window, padding = window_size//2, groups=channel)
+ mu2 = F.conv2d(img2, window, padding = window_size//2, groups=channel)
+
+ mu1_sq = mu1.pow(2)
+ mu2_sq = mu2.pow(2)
+ mu1_mu2 = mu1*mu2
+
+ sigma1_sq = F.conv2d(img1*img1, window, padding=window_size//2, groups=channel) - mu1_sq
+ sigma2_sq = F.conv2d(img2*img2, window, padding=window_size//2, groups=channel) - mu2_sq
+ sigma12 = F.conv2d(img1*img2, window, padding=window_size//2, groups=channel) - mu1_mu2
+
+ C1 = 0.01**2
+ C2 = 0.03**2
+
+ ssim_map = ((2*mu1_mu2 + C1)*(2*sigma12 + C2))/((mu1_sq + mu2_sq + C1)*(sigma1_sq + sigma2_sq + C2))
+
+ if size_average:
+ return ssim_map.mean()
+ else:
+ return ssim_map.mean(1).mean(1).mean(1)
+
+
+def SSIM(x, y):
+ C1 = 0.01 ** 2
+ C2 = 0.03 ** 2
+
+ mu_x = nn.AvgPool2d(3, 1, 1)(x)
+ mu_y = nn.AvgPool2d(3, 1, 1)(y)
+ mu_x_mu_y = mu_x * mu_y
+ mu_x_sq = mu_x.pow(2)
+ mu_y_sq = mu_y.pow(2)
+
+ sigma_x = nn.AvgPool2d(3, 1, 1)(x * x) - mu_x_sq
+ sigma_y = nn.AvgPool2d(3, 1, 1)(y * y) - mu_y_sq
+ sigma_xy = nn.AvgPool2d(3, 1, 1)(x * y) - mu_x_mu_y
+
+ SSIM_n = (2 * mu_x_mu_y + C1) * (2 * sigma_xy + C2)
+ SSIM_d = (mu_x_sq + mu_y_sq + C1) * (sigma_x + sigma_y + C2)
+ SSIM = SSIM_n / SSIM_d
+
+ return torch.clamp((1 - SSIM) / 2, 0, 1)
+
+
+def saliency_structure_consistency(x, y):
+ ssim = torch.mean(SSIM(x,y))
+ return ssim
diff --git a/py/BiRefNet_v2/make_a_copy.sh b/py/BiRefNet_v2/make_a_copy.sh
new file mode 100644
index 0000000..97a35fb
--- /dev/null
+++ b/py/BiRefNet_v2/make_a_copy.sh
@@ -0,0 +1,18 @@
+#!/bin/bash
+# Set dst repo here.
+repo=$1
+mkdir ../${repo}
+mkdir ../${repo}/evaluation
+mkdir ../${repo}/models
+mkdir ../${repo}/models/backbones
+mkdir ../${repo}/models/modules
+mkdir ../${repo}/models/refinement
+
+cp ./*.sh ../${repo}
+cp ./*.py ../${repo}
+cp ./evaluation/*.py ../${repo}/evaluation
+cp ./models/*.py ../${repo}/models
+cp ./models/backbones/*.py ../${repo}/models/backbones
+cp ./models/modules/*.py ../${repo}/models/modules
+cp ./models/refinement/*.py ../${repo}/models/refinement
+cp -r ./.git* ../${repo}
diff --git a/py/BiRefNet_v2/models/backbones/build_backbone.py b/py/BiRefNet_v2/models/backbones/build_backbone.py
new file mode 100644
index 0000000..65761a1
--- /dev/null
+++ b/py/BiRefNet_v2/models/backbones/build_backbone.py
@@ -0,0 +1,44 @@
+import torch
+import torch.nn as nn
+from collections import OrderedDict
+from torchvision.models import vgg16, vgg16_bn, VGG16_Weights, VGG16_BN_Weights, resnet50, ResNet50_Weights
+from ...models.backbones.pvt_v2 import pvt_v2_b0, pvt_v2_b1, pvt_v2_b2, pvt_v2_b5
+from ...models.backbones.swin_v1 import swin_v1_t, swin_v1_s, swin_v1_b, swin_v1_l
+from ...config import Config
+
+
+config = Config()
+
+def build_backbone(bb_name, pretrained=True, params_settings=''):
+ if bb_name == 'vgg16':
+ bb_net = list(vgg16(pretrained=VGG16_Weights.DEFAULT if pretrained else None).children())[0]
+ bb = nn.Sequential(OrderedDict({'conv1': bb_net[:4], 'conv2': bb_net[4:9], 'conv3': bb_net[9:16], 'conv4': bb_net[16:23]}))
+ elif bb_name == 'vgg16bn':
+ bb_net = list(vgg16_bn(pretrained=VGG16_BN_Weights.DEFAULT if pretrained else None).children())[0]
+ bb = nn.Sequential(OrderedDict({'conv1': bb_net[:6], 'conv2': bb_net[6:13], 'conv3': bb_net[13:23], 'conv4': bb_net[23:33]}))
+ elif bb_name == 'resnet50':
+ bb_net = list(resnet50(pretrained=ResNet50_Weights.DEFAULT if pretrained else None).children())
+ bb = nn.Sequential(OrderedDict({'conv1': nn.Sequential(*bb_net[0:3]), 'conv2': bb_net[4], 'conv3': bb_net[5], 'conv4': bb_net[6]}))
+ else:
+ bb = eval('{}({})'.format(bb_name, params_settings))
+ if pretrained:
+ bb = load_weights(bb, bb_name)
+ return bb
+
+def load_weights(model, model_name):
+ save_model = torch.load(config.weights[model_name], map_location='cpu')
+ model_dict = model.state_dict()
+ state_dict = {k: v if v.size() == model_dict[k].size() else model_dict[k] for k, v in save_model.items() if k in model_dict.keys()}
+ # to ignore the weights with mismatched size when I modify the backbone itself.
+ if not state_dict:
+ save_model_keys = list(save_model.keys())
+ sub_item = save_model_keys[0] if len(save_model_keys) == 1 else None
+ state_dict = {k: v if v.size() == model_dict[k].size() else model_dict[k] for k, v in save_model[sub_item].items() if k in model_dict.keys()}
+ if not state_dict or not sub_item:
+ print('Weights are not successully loaded. Check the state dict of weights file.')
+ return None
+ else:
+ print('Found correct weights in the "{}" item of loaded state_dict.'.format(sub_item))
+ model_dict.update(state_dict)
+ model.load_state_dict(model_dict)
+ return model
diff --git a/py/BiRefNet_v2/models/backbones/pvt_v2.py b/py/BiRefNet_v2/models/backbones/pvt_v2.py
new file mode 100644
index 0000000..4b902dd
--- /dev/null
+++ b/py/BiRefNet_v2/models/backbones/pvt_v2.py
@@ -0,0 +1,435 @@
+import torch
+import torch.nn as nn
+from functools import partial
+
+from timm.models.layers import DropPath, to_2tuple, trunc_normal_
+from timm.models import register_model
+
+import math
+
+from ...config import Config
+
+config = Config()
+
+class Mlp(nn.Module):
+ def __init__(self, in_features, hidden_features=None, out_features=None, act_layer=nn.GELU, drop=0.):
+ super().__init__()
+ out_features = out_features or in_features
+ hidden_features = hidden_features or in_features
+ self.fc1 = nn.Linear(in_features, hidden_features)
+ self.dwconv = DWConv(hidden_features)
+ self.act = act_layer()
+ self.fc2 = nn.Linear(hidden_features, out_features)
+ self.drop = nn.Dropout(drop)
+
+ self.apply(self._init_weights)
+
+ def _init_weights(self, m):
+ if isinstance(m, nn.Linear):
+ trunc_normal_(m.weight, std=.02)
+ if isinstance(m, nn.Linear) and m.bias is not None:
+ nn.init.constant_(m.bias, 0)
+ elif isinstance(m, nn.LayerNorm):
+ nn.init.constant_(m.bias, 0)
+ nn.init.constant_(m.weight, 1.0)
+ elif isinstance(m, nn.Conv2d):
+ fan_out = m.kernel_size[0] * m.kernel_size[1] * m.out_channels
+ fan_out //= m.groups
+ m.weight.data.normal_(0, math.sqrt(2.0 / fan_out))
+ if m.bias is not None:
+ m.bias.data.zero_()
+
+ def forward(self, x, H, W):
+ x = self.fc1(x)
+ x = self.dwconv(x, H, W)
+ x = self.act(x)
+ x = self.drop(x)
+ x = self.fc2(x)
+ x = self.drop(x)
+ return x
+
+
+class Attention(nn.Module):
+ def __init__(self, dim, num_heads=8, qkv_bias=False, qk_scale=None, attn_drop=0., proj_drop=0., sr_ratio=1):
+ super().__init__()
+ assert dim % num_heads == 0, f"dim {dim} should be divided by num_heads {num_heads}."
+
+ self.dim = dim
+ self.num_heads = num_heads
+ head_dim = dim // num_heads
+ self.scale = qk_scale or head_dim ** -0.5
+
+ self.q = nn.Linear(dim, dim, bias=qkv_bias)
+ self.kv = nn.Linear(dim, dim * 2, bias=qkv_bias)
+ self.attn_drop_prob = attn_drop
+ self.attn_drop = nn.Dropout(attn_drop)
+ self.proj = nn.Linear(dim, dim)
+ self.proj_drop = nn.Dropout(proj_drop)
+
+ self.sr_ratio = sr_ratio
+ if sr_ratio > 1:
+ self.sr = nn.Conv2d(dim, dim, kernel_size=sr_ratio, stride=sr_ratio)
+ self.norm = nn.LayerNorm(dim)
+
+ self.apply(self._init_weights)
+
+ def _init_weights(self, m):
+ if isinstance(m, nn.Linear):
+ trunc_normal_(m.weight, std=.02)
+ if isinstance(m, nn.Linear) and m.bias is not None:
+ nn.init.constant_(m.bias, 0)
+ elif isinstance(m, nn.LayerNorm):
+ nn.init.constant_(m.bias, 0)
+ nn.init.constant_(m.weight, 1.0)
+ elif isinstance(m, nn.Conv2d):
+ fan_out = m.kernel_size[0] * m.kernel_size[1] * m.out_channels
+ fan_out //= m.groups
+ m.weight.data.normal_(0, math.sqrt(2.0 / fan_out))
+ if m.bias is not None:
+ m.bias.data.zero_()
+
+ def forward(self, x, H, W):
+ B, N, C = x.shape
+ q = self.q(x).reshape(B, N, self.num_heads, C // self.num_heads).permute(0, 2, 1, 3)
+
+ if self.sr_ratio > 1:
+ x_ = x.permute(0, 2, 1).reshape(B, C, H, W)
+ x_ = self.sr(x_).reshape(B, C, -1).permute(0, 2, 1)
+ x_ = self.norm(x_)
+ kv = self.kv(x_).reshape(B, -1, 2, self.num_heads, C // self.num_heads).permute(2, 0, 3, 1, 4)
+ else:
+ kv = self.kv(x).reshape(B, -1, 2, self.num_heads, C // self.num_heads).permute(2, 0, 3, 1, 4)
+ k, v = kv[0], kv[1]
+
+ if config.SDPA_enabled:
+ x = torch.nn.functional.scaled_dot_product_attention(
+ q, k, v,
+ attn_mask=None, dropout_p=self.attn_drop_prob, is_causal=False
+ ).transpose(1, 2).reshape(B, N, C)
+ else:
+ attn = (q @ k.transpose(-2, -1)) * self.scale
+ attn = attn.softmax(dim=-1)
+ attn = self.attn_drop(attn)
+
+ x = (attn @ v).transpose(1, 2).reshape(B, N, C)
+ x = self.proj(x)
+ x = self.proj_drop(x)
+
+ return x
+
+
+class Block(nn.Module):
+
+ def __init__(self, dim, num_heads, mlp_ratio=4., qkv_bias=False, qk_scale=None, drop=0., attn_drop=0.,
+ drop_path=0., act_layer=nn.GELU, norm_layer=nn.LayerNorm, sr_ratio=1):
+ super().__init__()
+ self.norm1 = norm_layer(dim)
+ self.attn = Attention(
+ dim,
+ num_heads=num_heads, qkv_bias=qkv_bias, qk_scale=qk_scale,
+ attn_drop=attn_drop, proj_drop=drop, sr_ratio=sr_ratio)
+ # NOTE: drop path for stochastic depth, we shall see if this is better than dropout here
+ self.drop_path = DropPath(drop_path) if drop_path > 0. else nn.Identity()
+ self.norm2 = norm_layer(dim)
+ mlp_hidden_dim = int(dim * mlp_ratio)
+ self.mlp = Mlp(in_features=dim, hidden_features=mlp_hidden_dim, act_layer=act_layer, drop=drop)
+
+ self.apply(self._init_weights)
+
+ def _init_weights(self, m):
+ if isinstance(m, nn.Linear):
+ trunc_normal_(m.weight, std=.02)
+ if isinstance(m, nn.Linear) and m.bias is not None:
+ nn.init.constant_(m.bias, 0)
+ elif isinstance(m, nn.LayerNorm):
+ nn.init.constant_(m.bias, 0)
+ nn.init.constant_(m.weight, 1.0)
+ elif isinstance(m, nn.Conv2d):
+ fan_out = m.kernel_size[0] * m.kernel_size[1] * m.out_channels
+ fan_out //= m.groups
+ m.weight.data.normal_(0, math.sqrt(2.0 / fan_out))
+ if m.bias is not None:
+ m.bias.data.zero_()
+
+ def forward(self, x, H, W):
+ x = x + self.drop_path(self.attn(self.norm1(x), H, W))
+ x = x + self.drop_path(self.mlp(self.norm2(x), H, W))
+
+ return x
+
+
+class OverlapPatchEmbed(nn.Module):
+ """ Image to Patch Embedding
+ """
+
+ def __init__(self, img_size=224, patch_size=7, stride=4, in_channels=3, embed_dim=768):
+ super().__init__()
+ img_size = to_2tuple(img_size)
+ patch_size = to_2tuple(patch_size)
+
+ self.img_size = img_size
+ self.patch_size = patch_size
+ self.H, self.W = img_size[0] // patch_size[0], img_size[1] // patch_size[1]
+ self.num_patches = self.H * self.W
+ self.proj = nn.Conv2d(in_channels, embed_dim, kernel_size=patch_size, stride=stride,
+ padding=(patch_size[0] // 2, patch_size[1] // 2))
+ self.norm = nn.LayerNorm(embed_dim)
+
+ self.apply(self._init_weights)
+
+ def _init_weights(self, m):
+ if isinstance(m, nn.Linear):
+ trunc_normal_(m.weight, std=.02)
+ if isinstance(m, nn.Linear) and m.bias is not None:
+ nn.init.constant_(m.bias, 0)
+ elif isinstance(m, nn.LayerNorm):
+ nn.init.constant_(m.bias, 0)
+ nn.init.constant_(m.weight, 1.0)
+ elif isinstance(m, nn.Conv2d):
+ fan_out = m.kernel_size[0] * m.kernel_size[1] * m.out_channels
+ fan_out //= m.groups
+ m.weight.data.normal_(0, math.sqrt(2.0 / fan_out))
+ if m.bias is not None:
+ m.bias.data.zero_()
+
+ def forward(self, x):
+ x = self.proj(x)
+ _, _, H, W = x.shape
+ x = x.flatten(2).transpose(1, 2)
+ x = self.norm(x)
+
+ return x, H, W
+
+
+class PyramidVisionTransformerImpr(nn.Module):
+ def __init__(self, img_size=224, patch_size=16, in_channels=3, num_classes=1000, embed_dims=[64, 128, 256, 512],
+ num_heads=[1, 2, 4, 8], mlp_ratios=[4, 4, 4, 4], qkv_bias=False, qk_scale=None, drop_rate=0.,
+ attn_drop_rate=0., drop_path_rate=0., norm_layer=nn.LayerNorm,
+ depths=[3, 4, 6, 3], sr_ratios=[8, 4, 2, 1]):
+ super().__init__()
+ self.num_classes = num_classes
+ self.depths = depths
+
+ # patch_embed
+ self.patch_embed1 = OverlapPatchEmbed(img_size=img_size, patch_size=7, stride=4, in_channels=in_channels,
+ embed_dim=embed_dims[0])
+ self.patch_embed2 = OverlapPatchEmbed(img_size=img_size // 4, patch_size=3, stride=2, in_channels=embed_dims[0],
+ embed_dim=embed_dims[1])
+ self.patch_embed3 = OverlapPatchEmbed(img_size=img_size // 8, patch_size=3, stride=2, in_channels=embed_dims[1],
+ embed_dim=embed_dims[2])
+ self.patch_embed4 = OverlapPatchEmbed(img_size=img_size // 16, patch_size=3, stride=2, in_channels=embed_dims[2],
+ embed_dim=embed_dims[3])
+
+ # transformer encoder
+ dpr = [x.item() for x in torch.linspace(0, drop_path_rate, sum(depths))] # stochastic depth decay rule
+ cur = 0
+ self.block1 = nn.ModuleList([Block(
+ dim=embed_dims[0], num_heads=num_heads[0], mlp_ratio=mlp_ratios[0], qkv_bias=qkv_bias, qk_scale=qk_scale,
+ drop=drop_rate, attn_drop=attn_drop_rate, drop_path=dpr[cur + i], norm_layer=norm_layer,
+ sr_ratio=sr_ratios[0])
+ for i in range(depths[0])])
+ self.norm1 = norm_layer(embed_dims[0])
+
+ cur += depths[0]
+ self.block2 = nn.ModuleList([Block(
+ dim=embed_dims[1], num_heads=num_heads[1], mlp_ratio=mlp_ratios[1], qkv_bias=qkv_bias, qk_scale=qk_scale,
+ drop=drop_rate, attn_drop=attn_drop_rate, drop_path=dpr[cur + i], norm_layer=norm_layer,
+ sr_ratio=sr_ratios[1])
+ for i in range(depths[1])])
+ self.norm2 = norm_layer(embed_dims[1])
+
+ cur += depths[1]
+ self.block3 = nn.ModuleList([Block(
+ dim=embed_dims[2], num_heads=num_heads[2], mlp_ratio=mlp_ratios[2], qkv_bias=qkv_bias, qk_scale=qk_scale,
+ drop=drop_rate, attn_drop=attn_drop_rate, drop_path=dpr[cur + i], norm_layer=norm_layer,
+ sr_ratio=sr_ratios[2])
+ for i in range(depths[2])])
+ self.norm3 = norm_layer(embed_dims[2])
+
+ cur += depths[2]
+ self.block4 = nn.ModuleList([Block(
+ dim=embed_dims[3], num_heads=num_heads[3], mlp_ratio=mlp_ratios[3], qkv_bias=qkv_bias, qk_scale=qk_scale,
+ drop=drop_rate, attn_drop=attn_drop_rate, drop_path=dpr[cur + i], norm_layer=norm_layer,
+ sr_ratio=sr_ratios[3])
+ for i in range(depths[3])])
+ self.norm4 = norm_layer(embed_dims[3])
+
+ # classification head
+ # self.head = nn.Linear(embed_dims[3], num_classes) if num_classes > 0 else nn.Identity()
+
+ self.apply(self._init_weights)
+
+ def _init_weights(self, m):
+ if isinstance(m, nn.Linear):
+ trunc_normal_(m.weight, std=.02)
+ if isinstance(m, nn.Linear) and m.bias is not None:
+ nn.init.constant_(m.bias, 0)
+ elif isinstance(m, nn.LayerNorm):
+ nn.init.constant_(m.bias, 0)
+ nn.init.constant_(m.weight, 1.0)
+ elif isinstance(m, nn.Conv2d):
+ fan_out = m.kernel_size[0] * m.kernel_size[1] * m.out_channels
+ fan_out //= m.groups
+ m.weight.data.normal_(0, math.sqrt(2.0 / fan_out))
+ if m.bias is not None:
+ m.bias.data.zero_()
+
+ def init_weights(self, pretrained=None):
+ if isinstance(pretrained, str):
+ logger = 1
+ #load_checkpoint(self, pretrained, map_location='cpu', strict=False, logger=logger)
+
+ def reset_drop_path(self, drop_path_rate):
+ dpr = [x.item() for x in torch.linspace(0, drop_path_rate, sum(self.depths))]
+ cur = 0
+ for i in range(self.depths[0]):
+ self.block1[i].drop_path.drop_prob = dpr[cur + i]
+
+ cur += self.depths[0]
+ for i in range(self.depths[1]):
+ self.block2[i].drop_path.drop_prob = dpr[cur + i]
+
+ cur += self.depths[1]
+ for i in range(self.depths[2]):
+ self.block3[i].drop_path.drop_prob = dpr[cur + i]
+
+ cur += self.depths[2]
+ for i in range(self.depths[3]):
+ self.block4[i].drop_path.drop_prob = dpr[cur + i]
+
+ def freeze_patch_emb(self):
+ self.patch_embed1.requires_grad = False
+
+ @torch.jit.ignore
+ def no_weight_decay(self):
+ return {'pos_embed1', 'pos_embed2', 'pos_embed3', 'pos_embed4', 'cls_token'} # has pos_embed may be better
+
+ def get_classifier(self):
+ return self.head
+
+ def reset_classifier(self, num_classes, global_pool=''):
+ self.num_classes = num_classes
+ self.head = nn.Linear(self.embed_dim, num_classes) if num_classes > 0 else nn.Identity()
+
+ def forward_features(self, x):
+ B = x.shape[0]
+ outs = []
+
+ # stage 1
+ x, H, W = self.patch_embed1(x)
+ for i, blk in enumerate(self.block1):
+ x = blk(x, H, W)
+ x = self.norm1(x)
+ x = x.reshape(B, H, W, -1).permute(0, 3, 1, 2).contiguous()
+ outs.append(x)
+
+ # stage 2
+ x, H, W = self.patch_embed2(x)
+ for i, blk in enumerate(self.block2):
+ x = blk(x, H, W)
+ x = self.norm2(x)
+ x = x.reshape(B, H, W, -1).permute(0, 3, 1, 2).contiguous()
+ outs.append(x)
+
+ # stage 3
+ x, H, W = self.patch_embed3(x)
+ for i, blk in enumerate(self.block3):
+ x = blk(x, H, W)
+ x = self.norm3(x)
+ x = x.reshape(B, H, W, -1).permute(0, 3, 1, 2).contiguous()
+ outs.append(x)
+
+ # stage 4
+ x, H, W = self.patch_embed4(x)
+ for i, blk in enumerate(self.block4):
+ x = blk(x, H, W)
+ x = self.norm4(x)
+ x = x.reshape(B, H, W, -1).permute(0, 3, 1, 2).contiguous()
+ outs.append(x)
+
+ return outs
+
+ # return x.mean(dim=1)
+
+ def forward(self, x):
+ x = self.forward_features(x)
+ # x = self.head(x)
+
+ return x
+
+
+class DWConv(nn.Module):
+ def __init__(self, dim=768):
+ super(DWConv, self).__init__()
+ self.dwconv = nn.Conv2d(dim, dim, 3, 1, 1, bias=True, groups=dim)
+
+ def forward(self, x, H, W):
+ B, N, C = x.shape
+ x = x.transpose(1, 2).view(B, C, H, W).contiguous()
+ x = self.dwconv(x)
+ x = x.flatten(2).transpose(1, 2)
+
+ return x
+
+
+def _conv_filter(state_dict, patch_size=16):
+ """ convert patch embedding weight from manual patchify + linear proj to conv"""
+ out_dict = {}
+ for k, v in state_dict.items():
+ if 'patch_embed.proj.weight' in k:
+ v = v.reshape((v.shape[0], 3, patch_size, patch_size))
+ out_dict[k] = v
+
+ return out_dict
+
+
+## @register_model
+class pvt_v2_b0(PyramidVisionTransformerImpr):
+ def __init__(self, **kwargs):
+ super(pvt_v2_b0, self).__init__(
+ patch_size=4, embed_dims=[32, 64, 160, 256], num_heads=[1, 2, 5, 8], mlp_ratios=[8, 8, 4, 4],
+ qkv_bias=True, norm_layer=partial(nn.LayerNorm, eps=1e-6), depths=[2, 2, 2, 2], sr_ratios=[8, 4, 2, 1],
+ drop_rate=0.0, drop_path_rate=0.1)
+
+
+
+## @register_model
+class pvt_v2_b1(PyramidVisionTransformerImpr):
+ def __init__(self, **kwargs):
+ super(pvt_v2_b1, self).__init__(
+ patch_size=4, embed_dims=[64, 128, 320, 512], num_heads=[1, 2, 5, 8], mlp_ratios=[8, 8, 4, 4],
+ qkv_bias=True, norm_layer=partial(nn.LayerNorm, eps=1e-6), depths=[2, 2, 2, 2], sr_ratios=[8, 4, 2, 1],
+ drop_rate=0.0, drop_path_rate=0.1)
+
+## @register_model
+class pvt_v2_b2(PyramidVisionTransformerImpr):
+ def __init__(self, in_channels=3, **kwargs):
+ super(pvt_v2_b2, self).__init__(
+ patch_size=4, embed_dims=[64, 128, 320, 512], num_heads=[1, 2, 5, 8], mlp_ratios=[8, 8, 4, 4],
+ qkv_bias=True, norm_layer=partial(nn.LayerNorm, eps=1e-6), depths=[3, 4, 6, 3], sr_ratios=[8, 4, 2, 1],
+ drop_rate=0.0, drop_path_rate=0.1, in_channels=in_channels)
+
+## @register_model
+class pvt_v2_b3(PyramidVisionTransformerImpr):
+ def __init__(self, **kwargs):
+ super(pvt_v2_b3, self).__init__(
+ patch_size=4, embed_dims=[64, 128, 320, 512], num_heads=[1, 2, 5, 8], mlp_ratios=[8, 8, 4, 4],
+ qkv_bias=True, norm_layer=partial(nn.LayerNorm, eps=1e-6), depths=[3, 4, 18, 3], sr_ratios=[8, 4, 2, 1],
+ drop_rate=0.0, drop_path_rate=0.1)
+
+## @register_model
+class pvt_v2_b4(PyramidVisionTransformerImpr):
+ def __init__(self, **kwargs):
+ super(pvt_v2_b4, self).__init__(
+ patch_size=4, embed_dims=[64, 128, 320, 512], num_heads=[1, 2, 5, 8], mlp_ratios=[8, 8, 4, 4],
+ qkv_bias=True, norm_layer=partial(nn.LayerNorm, eps=1e-6), depths=[3, 8, 27, 3], sr_ratios=[8, 4, 2, 1],
+ drop_rate=0.0, drop_path_rate=0.1)
+
+
+## @register_model
+class pvt_v2_b5(PyramidVisionTransformerImpr):
+ def __init__(self, **kwargs):
+ super(pvt_v2_b5, self).__init__(
+ patch_size=4, embed_dims=[64, 128, 320, 512], num_heads=[1, 2, 5, 8], mlp_ratios=[4, 4, 4, 4],
+ qkv_bias=True, norm_layer=partial(nn.LayerNorm, eps=1e-6), depths=[3, 6, 40, 3], sr_ratios=[8, 4, 2, 1],
+ drop_rate=0.0, drop_path_rate=0.1)
diff --git a/py/BiRefNet_v2/models/backbones/swin_v1.py b/py/BiRefNet_v2/models/backbones/swin_v1.py
new file mode 100644
index 0000000..7739622
--- /dev/null
+++ b/py/BiRefNet_v2/models/backbones/swin_v1.py
@@ -0,0 +1,627 @@
+# --------------------------------------------------------
+# Swin Transformer
+# Copyright (c) 2021 Microsoft
+# Licensed under The MIT License [see LICENSE for details]
+# Written by Ze Liu, Yutong Lin, Yixuan Wei
+# --------------------------------------------------------
+
+import torch
+import torch.nn as nn
+import torch.nn.functional as F
+import torch.utils.checkpoint as checkpoint
+import numpy as np
+from timm.models.layers import DropPath, to_2tuple, trunc_normal_
+
+from ...config import Config
+
+
+config = Config()
+
+class Mlp(nn.Module):
+ """ Multilayer perceptron."""
+
+ def __init__(self, in_features, hidden_features=None, out_features=None, act_layer=nn.GELU, drop=0.):
+ super().__init__()
+ out_features = out_features or in_features
+ hidden_features = hidden_features or in_features
+ self.fc1 = nn.Linear(in_features, hidden_features)
+ self.act = act_layer()
+ self.fc2 = nn.Linear(hidden_features, out_features)
+ self.drop = nn.Dropout(drop)
+
+ def forward(self, x):
+ x = self.fc1(x)
+ x = self.act(x)
+ x = self.drop(x)
+ x = self.fc2(x)
+ x = self.drop(x)
+ return x
+
+
+def window_partition(x, window_size):
+ """
+ Args:
+ x: (B, H, W, C)
+ window_size (int): window size
+
+ Returns:
+ windows: (num_windows*B, window_size, window_size, C)
+ """
+ B, H, W, C = x.shape
+ x = x.view(B, H // window_size, window_size, W // window_size, window_size, C)
+ windows = x.permute(0, 1, 3, 2, 4, 5).contiguous().view(-1, window_size, window_size, C)
+ return windows
+
+
+def window_reverse(windows, window_size, H, W):
+ """
+ Args:
+ windows: (num_windows*B, window_size, window_size, C)
+ window_size (int): Window size
+ H (int): Height of image
+ W (int): Width of image
+
+ Returns:
+ x: (B, H, W, C)
+ """
+ B = int(windows.shape[0] / (H * W / window_size / window_size))
+ x = windows.view(B, H // window_size, W // window_size, window_size, window_size, -1)
+ x = x.permute(0, 1, 3, 2, 4, 5).contiguous().view(B, H, W, -1)
+ return x
+
+
+class WindowAttention(nn.Module):
+ """ Window based multi-head self attention (W-MSA) module with relative position bias.
+ It supports both of shifted and non-shifted window.
+
+ Args:
+ dim (int): Number of input channels.
+ window_size (tuple[int]): The height and width of the window.
+ num_heads (int): Number of attention heads.
+ qkv_bias (bool, optional): If True, add a learnable bias to query, key, value. Default: True
+ qk_scale (float | None, optional): Override default qk scale of head_dim ** -0.5 if set
+ attn_drop (float, optional): Dropout ratio of attention weight. Default: 0.0
+ proj_drop (float, optional): Dropout ratio of output. Default: 0.0
+ """
+
+ def __init__(self, dim, window_size, num_heads, qkv_bias=True, qk_scale=None, attn_drop=0., proj_drop=0.):
+
+ super().__init__()
+ self.dim = dim
+ self.window_size = window_size # Wh, Ww
+ self.num_heads = num_heads
+ head_dim = dim // num_heads
+ self.scale = qk_scale or head_dim ** -0.5
+
+ # define a parameter table of relative position bias
+ self.relative_position_bias_table = nn.Parameter(
+ torch.zeros((2 * window_size[0] - 1) * (2 * window_size[1] - 1), num_heads)) # 2*Wh-1 * 2*Ww-1, nH
+
+ # get pair-wise relative position index for each token inside the window
+ coords_h = torch.arange(self.window_size[0])
+ coords_w = torch.arange(self.window_size[1])
+ coords = torch.stack(torch.meshgrid([coords_h, coords_w], indexing='ij')) # 2, Wh, Ww
+ coords_flatten = torch.flatten(coords, 1) # 2, Wh*Ww
+ relative_coords = coords_flatten[:, :, None] - coords_flatten[:, None, :] # 2, Wh*Ww, Wh*Ww
+ relative_coords = relative_coords.permute(1, 2, 0).contiguous() # Wh*Ww, Wh*Ww, 2
+ relative_coords[:, :, 0] += self.window_size[0] - 1 # shift to start from 0
+ relative_coords[:, :, 1] += self.window_size[1] - 1
+ relative_coords[:, :, 0] *= 2 * self.window_size[1] - 1
+ relative_position_index = relative_coords.sum(-1) # Wh*Ww, Wh*Ww
+ self.register_buffer("relative_position_index", relative_position_index)
+
+ self.qkv = nn.Linear(dim, dim * 3, bias=qkv_bias)
+ self.attn_drop_prob = attn_drop
+ self.attn_drop = nn.Dropout(attn_drop)
+ self.proj = nn.Linear(dim, dim)
+ self.proj_drop = nn.Dropout(proj_drop)
+
+ trunc_normal_(self.relative_position_bias_table, std=.02)
+ self.softmax = nn.Softmax(dim=-1)
+
+ def forward(self, x, mask=None):
+ """ Forward function.
+
+ Args:
+ x: input features with shape of (num_windows*B, N, C)
+ mask: (0/-inf) mask with shape of (num_windows, Wh*Ww, Wh*Ww) or None
+ """
+ B_, N, C = x.shape
+ qkv = self.qkv(x).reshape(B_, N, 3, self.num_heads, C // self.num_heads).permute(2, 0, 3, 1, 4)
+ q, k, v = qkv[0], qkv[1], qkv[2] # make torchscript happy (cannot use tensor as tuple)
+
+ q = q * self.scale
+
+ if config.SDPA_enabled:
+ x = torch.nn.functional.scaled_dot_product_attention(
+ q, k, v,
+ attn_mask=None, dropout_p=self.attn_drop_prob, is_causal=False
+ ).transpose(1, 2).reshape(B_, N, C)
+ else:
+ attn = (q @ k.transpose(-2, -1))
+
+ relative_position_bias = self.relative_position_bias_table[self.relative_position_index.view(-1)].view(
+ self.window_size[0] * self.window_size[1], self.window_size[0] * self.window_size[1], -1) # Wh*Ww,Wh*Ww,nH
+ relative_position_bias = relative_position_bias.permute(2, 0, 1).contiguous() # nH, Wh*Ww, Wh*Ww
+ attn = attn + relative_position_bias.unsqueeze(0)
+
+ if mask is not None:
+ nW = mask.shape[0]
+ attn = attn.view(B_ // nW, nW, self.num_heads, N, N) + mask.unsqueeze(1).unsqueeze(0)
+ attn = attn.view(-1, self.num_heads, N, N)
+ attn = self.softmax(attn)
+ else:
+ attn = self.softmax(attn)
+
+ attn = self.attn_drop(attn)
+
+ x = (attn @ v).transpose(1, 2).reshape(B_, N, C)
+ x = self.proj(x)
+ x = self.proj_drop(x)
+ return x
+
+
+class SwinTransformerBlock(nn.Module):
+ """ Swin Transformer Block.
+
+ Args:
+ dim (int): Number of input channels.
+ num_heads (int): Number of attention heads.
+ window_size (int): Window size.
+ shift_size (int): Shift size for SW-MSA.
+ mlp_ratio (float): Ratio of mlp hidden dim to embedding dim.
+ qkv_bias (bool, optional): If True, add a learnable bias to query, key, value. Default: True
+ qk_scale (float | None, optional): Override default qk scale of head_dim ** -0.5 if set.
+ drop (float, optional): Dropout rate. Default: 0.0
+ attn_drop (float, optional): Attention dropout rate. Default: 0.0
+ drop_path (float, optional): Stochastic depth rate. Default: 0.0
+ act_layer (nn.Module, optional): Activation layer. Default: nn.GELU
+ norm_layer (nn.Module, optional): Normalization layer. Default: nn.LayerNorm
+ """
+
+ def __init__(self, dim, num_heads, window_size=7, shift_size=0,
+ mlp_ratio=4., qkv_bias=True, qk_scale=None, drop=0., attn_drop=0., drop_path=0.,
+ act_layer=nn.GELU, norm_layer=nn.LayerNorm):
+ super().__init__()
+ self.dim = dim
+ self.num_heads = num_heads
+ self.window_size = window_size
+ self.shift_size = shift_size
+ self.mlp_ratio = mlp_ratio
+ assert 0 <= self.shift_size < self.window_size, "shift_size must in 0-window_size"
+
+ self.norm1 = norm_layer(dim)
+ self.attn = WindowAttention(
+ dim, window_size=to_2tuple(self.window_size), num_heads=num_heads,
+ qkv_bias=qkv_bias, qk_scale=qk_scale, attn_drop=attn_drop, proj_drop=drop)
+
+ self.drop_path = DropPath(drop_path) if drop_path > 0. else nn.Identity()
+ self.norm2 = norm_layer(dim)
+ mlp_hidden_dim = int(dim * mlp_ratio)
+ self.mlp = Mlp(in_features=dim, hidden_features=mlp_hidden_dim, act_layer=act_layer, drop=drop)
+
+ self.H = None
+ self.W = None
+
+ def forward(self, x, mask_matrix):
+ """ Forward function.
+
+ Args:
+ x: Input feature, tensor size (B, H*W, C).
+ H, W: Spatial resolution of the input feature.
+ mask_matrix: Attention mask for cyclic shift.
+ """
+ B, L, C = x.shape
+ H, W = self.H, self.W
+ assert L == H * W, "input feature has wrong size"
+
+ shortcut = x
+ x = self.norm1(x)
+ x = x.view(B, H, W, C)
+
+ # pad feature maps to multiples of window size
+ pad_l = pad_t = 0
+ pad_r = (self.window_size - W % self.window_size) % self.window_size
+ pad_b = (self.window_size - H % self.window_size) % self.window_size
+ x = F.pad(x, (0, 0, pad_l, pad_r, pad_t, pad_b))
+ _, Hp, Wp, _ = x.shape
+
+ # cyclic shift
+ if self.shift_size > 0:
+ shifted_x = torch.roll(x, shifts=(-self.shift_size, -self.shift_size), dims=(1, 2))
+ attn_mask = mask_matrix
+ else:
+ shifted_x = x
+ attn_mask = None
+
+ # partition windows
+ x_windows = window_partition(shifted_x, self.window_size) # nW*B, window_size, window_size, C
+ x_windows = x_windows.view(-1, self.window_size * self.window_size, C) # nW*B, window_size*window_size, C
+
+ # W-MSA/SW-MSA
+ attn_windows = self.attn(x_windows, mask=attn_mask) # nW*B, window_size*window_size, C
+
+ # merge windows
+ attn_windows = attn_windows.view(-1, self.window_size, self.window_size, C)
+ shifted_x = window_reverse(attn_windows, self.window_size, Hp, Wp) # B H' W' C
+
+ # reverse cyclic shift
+ if self.shift_size > 0:
+ x = torch.roll(shifted_x, shifts=(self.shift_size, self.shift_size), dims=(1, 2))
+ else:
+ x = shifted_x
+
+ if pad_r > 0 or pad_b > 0:
+ x = x[:, :H, :W, :].contiguous()
+
+ x = x.view(B, H * W, C)
+
+ # FFN
+ x = shortcut + self.drop_path(x)
+ x = x + self.drop_path(self.mlp(self.norm2(x)))
+
+ return x
+
+
+class PatchMerging(nn.Module):
+ """ Patch Merging Layer
+
+ Args:
+ dim (int): Number of input channels.
+ norm_layer (nn.Module, optional): Normalization layer. Default: nn.LayerNorm
+ """
+ def __init__(self, dim, norm_layer=nn.LayerNorm):
+ super().__init__()
+ self.dim = dim
+ self.reduction = nn.Linear(4 * dim, 2 * dim, bias=False)
+ self.norm = norm_layer(4 * dim)
+
+ def forward(self, x, H, W):
+ """ Forward function.
+
+ Args:
+ x: Input feature, tensor size (B, H*W, C).
+ H, W: Spatial resolution of the input feature.
+ """
+ B, L, C = x.shape
+ assert L == H * W, "input feature has wrong size"
+
+ x = x.view(B, H, W, C)
+
+ # padding
+ pad_input = (H % 2 == 1) or (W % 2 == 1)
+ if pad_input:
+ x = F.pad(x, (0, 0, 0, W % 2, 0, H % 2))
+
+ x0 = x[:, 0::2, 0::2, :] # B H/2 W/2 C
+ x1 = x[:, 1::2, 0::2, :] # B H/2 W/2 C
+ x2 = x[:, 0::2, 1::2, :] # B H/2 W/2 C
+ x3 = x[:, 1::2, 1::2, :] # B H/2 W/2 C
+ x = torch.cat([x0, x1, x2, x3], -1) # B H/2 W/2 4*C
+ x = x.view(B, -1, 4 * C) # B H/2*W/2 4*C
+
+ x = self.norm(x)
+ x = self.reduction(x)
+
+ return x
+
+
+class BasicLayer(nn.Module):
+ """ A basic Swin Transformer layer for one stage.
+
+ Args:
+ dim (int): Number of feature channels
+ depth (int): Depths of this stage.
+ num_heads (int): Number of attention head.
+ window_size (int): Local window size. Default: 7.
+ mlp_ratio (float): Ratio of mlp hidden dim to embedding dim. Default: 4.
+ qkv_bias (bool, optional): If True, add a learnable bias to query, key, value. Default: True
+ qk_scale (float | None, optional): Override default qk scale of head_dim ** -0.5 if set.
+ drop (float, optional): Dropout rate. Default: 0.0
+ attn_drop (float, optional): Attention dropout rate. Default: 0.0
+ drop_path (float | tuple[float], optional): Stochastic depth rate. Default: 0.0
+ norm_layer (nn.Module, optional): Normalization layer. Default: nn.LayerNorm
+ downsample (nn.Module | None, optional): Downsample layer at the end of the layer. Default: None
+ use_checkpoint (bool): Whether to use checkpointing to save memory. Default: False.
+ """
+
+ def __init__(self,
+ dim,
+ depth,
+ num_heads,
+ window_size=7,
+ mlp_ratio=4.,
+ qkv_bias=True,
+ qk_scale=None,
+ drop=0.,
+ attn_drop=0.,
+ drop_path=0.,
+ norm_layer=nn.LayerNorm,
+ downsample=None,
+ use_checkpoint=False):
+ super().__init__()
+ self.window_size = window_size
+ self.shift_size = window_size // 2
+ self.depth = depth
+ self.use_checkpoint = use_checkpoint
+
+ # build blocks
+ self.blocks = nn.ModuleList([
+ SwinTransformerBlock(
+ dim=dim,
+ num_heads=num_heads,
+ window_size=window_size,
+ shift_size=0 if (i % 2 == 0) else window_size // 2,
+ mlp_ratio=mlp_ratio,
+ qkv_bias=qkv_bias,
+ qk_scale=qk_scale,
+ drop=drop,
+ attn_drop=attn_drop,
+ drop_path=drop_path[i] if isinstance(drop_path, list) else drop_path,
+ norm_layer=norm_layer)
+ for i in range(depth)])
+
+ # patch merging layer
+ if downsample is not None:
+ self.downsample = downsample(dim=dim, norm_layer=norm_layer)
+ else:
+ self.downsample = None
+
+ def forward(self, x, H, W):
+ """ Forward function.
+
+ Args:
+ x: Input feature, tensor size (B, H*W, C).
+ H, W: Spatial resolution of the input feature.
+ """
+
+ # calculate attention mask for SW-MSA
+ Hp = int(np.ceil(H / self.window_size)) * self.window_size
+ Wp = int(np.ceil(W / self.window_size)) * self.window_size
+ img_mask = torch.zeros((1, Hp, Wp, 1), device=x.device) # 1 Hp Wp 1
+ h_slices = (slice(0, -self.window_size),
+ slice(-self.window_size, -self.shift_size),
+ slice(-self.shift_size, None))
+ w_slices = (slice(0, -self.window_size),
+ slice(-self.window_size, -self.shift_size),
+ slice(-self.shift_size, None))
+ cnt = 0
+ for h in h_slices:
+ for w in w_slices:
+ img_mask[:, h, w, :] = cnt
+ cnt += 1
+
+ mask_windows = window_partition(img_mask, self.window_size) # nW, window_size, window_size, 1
+ mask_windows = mask_windows.view(-1, self.window_size * self.window_size)
+ attn_mask = mask_windows.unsqueeze(1) - mask_windows.unsqueeze(2)
+ attn_mask = attn_mask.masked_fill(attn_mask != 0, float(-100.0)).masked_fill(attn_mask == 0, float(0.0))
+
+ for blk in self.blocks:
+ blk.H, blk.W = H, W
+ if self.use_checkpoint:
+ x = checkpoint.checkpoint(blk, x, attn_mask)
+ else:
+ x = blk(x, attn_mask)
+ if self.downsample is not None:
+ x_down = self.downsample(x, H, W)
+ Wh, Ww = (H + 1) // 2, (W + 1) // 2
+ return x, H, W, x_down, Wh, Ww
+ else:
+ return x, H, W, x, H, W
+
+
+class PatchEmbed(nn.Module):
+ """ Image to Patch Embedding
+
+ Args:
+ patch_size (int): Patch token size. Default: 4.
+ in_channels (int): Number of input image channels. Default: 3.
+ embed_dim (int): Number of linear projection output channels. Default: 96.
+ norm_layer (nn.Module, optional): Normalization layer. Default: None
+ """
+
+ def __init__(self, patch_size=4, in_channels=3, embed_dim=96, norm_layer=None):
+ super().__init__()
+ patch_size = to_2tuple(patch_size)
+ self.patch_size = patch_size
+
+ self.in_channels = in_channels
+ self.embed_dim = embed_dim
+
+ self.proj = nn.Conv2d(in_channels, embed_dim, kernel_size=patch_size, stride=patch_size)
+ if norm_layer is not None:
+ self.norm = norm_layer(embed_dim)
+ else:
+ self.norm = None
+
+ def forward(self, x):
+ """Forward function."""
+ # padding
+ _, _, H, W = x.size()
+ if W % self.patch_size[1] != 0:
+ x = F.pad(x, (0, self.patch_size[1] - W % self.patch_size[1]))
+ if H % self.patch_size[0] != 0:
+ x = F.pad(x, (0, 0, 0, self.patch_size[0] - H % self.patch_size[0]))
+
+ x = self.proj(x) # B C Wh Ww
+ if self.norm is not None:
+ Wh, Ww = x.size(2), x.size(3)
+ x = x.flatten(2).transpose(1, 2)
+ x = self.norm(x)
+ x = x.transpose(1, 2).view(-1, self.embed_dim, Wh, Ww)
+
+ return x
+
+
+class SwinTransformer(nn.Module):
+ """ Swin Transformer backbone.
+ A PyTorch impl of : `Swin Transformer: Hierarchical Vision Transformer using Shifted Windows` -
+ https://arxiv.org/pdf/2103.14030
+
+ Args:
+ pretrain_img_size (int): Input image size for training the pretrained model,
+ used in absolute postion embedding. Default 224.
+ patch_size (int | tuple(int)): Patch size. Default: 4.
+ in_channels (int): Number of input image channels. Default: 3.
+ embed_dim (int): Number of linear projection output channels. Default: 96.
+ depths (tuple[int]): Depths of each Swin Transformer stage.
+ num_heads (tuple[int]): Number of attention head of each stage.
+ window_size (int): Window size. Default: 7.
+ mlp_ratio (float): Ratio of mlp hidden dim to embedding dim. Default: 4.
+ qkv_bias (bool): If True, add a learnable bias to query, key, value. Default: True
+ qk_scale (float): Override default qk scale of head_dim ** -0.5 if set.
+ drop_rate (float): Dropout rate.
+ attn_drop_rate (float): Attention dropout rate. Default: 0.
+ drop_path_rate (float): Stochastic depth rate. Default: 0.2.
+ norm_layer (nn.Module): Normalization layer. Default: nn.LayerNorm.
+ ape (bool): If True, add absolute position embedding to the patch embedding. Default: False.
+ patch_norm (bool): If True, add normalization after patch embedding. Default: True.
+ out_indices (Sequence[int]): Output from which stages.
+ frozen_stages (int): Stages to be frozen (stop grad and set eval mode).
+ -1 means not freezing any parameters.
+ use_checkpoint (bool): Whether to use checkpointing to save memory. Default: False.
+ """
+
+ def __init__(self,
+ pretrain_img_size=224,
+ patch_size=4,
+ in_channels=3,
+ embed_dim=96,
+ depths=[2, 2, 6, 2],
+ num_heads=[3, 6, 12, 24],
+ window_size=7,
+ mlp_ratio=4.,
+ qkv_bias=True,
+ qk_scale=None,
+ drop_rate=0.,
+ attn_drop_rate=0.,
+ drop_path_rate=0.2,
+ norm_layer=nn.LayerNorm,
+ ape=False,
+ patch_norm=True,
+ out_indices=(0, 1, 2, 3),
+ frozen_stages=-1,
+ use_checkpoint=False):
+ super().__init__()
+
+ self.pretrain_img_size = pretrain_img_size
+ self.num_layers = len(depths)
+ self.embed_dim = embed_dim
+ self.ape = ape
+ self.patch_norm = patch_norm
+ self.out_indices = out_indices
+ self.frozen_stages = frozen_stages
+
+ # split image into non-overlapping patches
+ self.patch_embed = PatchEmbed(
+ patch_size=patch_size, in_channels=in_channels, embed_dim=embed_dim,
+ norm_layer=norm_layer if self.patch_norm else None)
+
+ # absolute position embedding
+ if self.ape:
+ pretrain_img_size = to_2tuple(pretrain_img_size)
+ patch_size = to_2tuple(patch_size)
+ patches_resolution = [pretrain_img_size[0] // patch_size[0], pretrain_img_size[1] // patch_size[1]]
+
+ self.absolute_pos_embed = nn.Parameter(torch.zeros(1, embed_dim, patches_resolution[0], patches_resolution[1]))
+ trunc_normal_(self.absolute_pos_embed, std=.02)
+
+ self.pos_drop = nn.Dropout(p=drop_rate)
+
+ # stochastic depth
+ dpr = [x.item() for x in torch.linspace(0, drop_path_rate, sum(depths))] # stochastic depth decay rule
+
+ # build layers
+ self.layers = nn.ModuleList()
+ for i_layer in range(self.num_layers):
+ layer = BasicLayer(
+ dim=int(embed_dim * 2 ** i_layer),
+ depth=depths[i_layer],
+ num_heads=num_heads[i_layer],
+ window_size=window_size,
+ mlp_ratio=mlp_ratio,
+ qkv_bias=qkv_bias,
+ qk_scale=qk_scale,
+ drop=drop_rate,
+ attn_drop=attn_drop_rate,
+ drop_path=dpr[sum(depths[:i_layer]):sum(depths[:i_layer + 1])],
+ norm_layer=norm_layer,
+ downsample=PatchMerging if (i_layer < self.num_layers - 1) else None,
+ use_checkpoint=use_checkpoint)
+ self.layers.append(layer)
+
+ num_features = [int(embed_dim * 2 ** i) for i in range(self.num_layers)]
+ self.num_features = num_features
+
+ # add a norm layer for each output
+ for i_layer in out_indices:
+ layer = norm_layer(num_features[i_layer])
+ layer_name = f'norm{i_layer}'
+ self.add_module(layer_name, layer)
+
+ self._freeze_stages()
+
+ def _freeze_stages(self):
+ if self.frozen_stages >= 0:
+ self.patch_embed.eval()
+ for param in self.patch_embed.parameters():
+ param.requires_grad = False
+
+ if self.frozen_stages >= 1 and self.ape:
+ self.absolute_pos_embed.requires_grad = False
+
+ if self.frozen_stages >= 2:
+ self.pos_drop.eval()
+ for i in range(0, self.frozen_stages - 1):
+ m = self.layers[i]
+ m.eval()
+ for param in m.parameters():
+ param.requires_grad = False
+
+
+ def forward(self, x):
+ """Forward function."""
+ x = self.patch_embed(x)
+
+ Wh, Ww = x.size(2), x.size(3)
+ if self.ape:
+ # interpolate the position embedding to the corresponding size
+ absolute_pos_embed = F.interpolate(self.absolute_pos_embed, size=(Wh, Ww), mode='bicubic')
+ x = (x + absolute_pos_embed) # B Wh*Ww C
+
+ outs = []#x.contiguous()]
+ x = x.flatten(2).transpose(1, 2)
+ x = self.pos_drop(x)
+ for i in range(self.num_layers):
+ layer = self.layers[i]
+ x_out, H, W, x, Wh, Ww = layer(x, Wh, Ww)
+
+ if i in self.out_indices:
+ norm_layer = getattr(self, f'norm{i}')
+ x_out = norm_layer(x_out)
+
+ out = x_out.view(-1, H, W, self.num_features[i]).permute(0, 3, 1, 2).contiguous()
+ outs.append(out)
+
+ return tuple(outs)
+
+ def train(self, mode=True):
+ """Convert the model into training mode while keep layers freezed."""
+ super(SwinTransformer, self).train(mode)
+ self._freeze_stages()
+
+def swin_v1_t():
+ model = SwinTransformer(embed_dim=96, depths=[2, 2, 6, 2], num_heads=[3, 6, 12, 24], window_size=7)
+ return model
+
+def swin_v1_s():
+ model = SwinTransformer(embed_dim=96, depths=[2, 2, 18, 2], num_heads=[3, 6, 12, 24], window_size=7)
+ return model
+
+def swin_v1_b():
+ model = SwinTransformer(embed_dim=128, depths=[2, 2, 18, 2], num_heads=[4, 8, 16, 32], window_size=12)
+ return model
+
+def swin_v1_l():
+ model = SwinTransformer(embed_dim=192, depths=[2, 2, 18, 2], num_heads=[6, 12, 24, 48], window_size=12)
+ return model
diff --git a/py/BiRefNet_v2/models/birefnet.py b/py/BiRefNet_v2/models/birefnet.py
new file mode 100644
index 0000000..e3fe196
--- /dev/null
+++ b/py/BiRefNet_v2/models/birefnet.py
@@ -0,0 +1,286 @@
+import torch
+import torch.nn as nn
+import torch.nn.functional as F
+from kornia.filters import laplacian
+from huggingface_hub import PyTorchModelHubMixin
+
+from ..config import Config
+from ..dataset import class_labels_TR_sorted
+from ..models.backbones.build_backbone import build_backbone
+from ..models.modules.decoder_blocks import BasicDecBlk, ResBlk
+from ..models.modules.lateral_blocks import BasicLatBlk
+from ..models.modules.aspp import ASPP, ASPPDeformable
+from ..models.refinement.refiner import Refiner, RefinerPVTInChannels4, RefUNet
+from ..models.refinement.stem_layer import StemLayer
+
+
+class BiRefNet(
+ nn.Module,
+ PyTorchModelHubMixin,
+ library_name="birefnet",
+ repo_url="https://github.com/ZhengPeng7/BiRefNet",
+ tags=['Image Segmentation', 'Background Removal', 'Mask Generation', 'Dichotomous Image Segmentation', 'Camouflaged Object Detection', 'Salient Object Detection']
+):
+ def __init__(self, bb_pretrained=True):
+ super(BiRefNet, self).__init__()
+ self.config = Config()
+ self.epoch = 1
+ self.bb = build_backbone(self.config.bb, pretrained=bb_pretrained)
+
+ channels = self.config.lateral_channels_in_collection
+
+ if self.config.auxiliary_classification:
+ self.avgpool = nn.AdaptiveAvgPool2d((1, 1))
+ self.cls_head = nn.Sequential(
+ nn.Linear(channels[0], len(class_labels_TR_sorted))
+ )
+
+ if self.config.squeeze_block:
+ self.squeeze_module = nn.Sequential(*[
+ eval(self.config.squeeze_block.split('_x')[0])(channels[0]+sum(self.config.cxt), channels[0])
+ for _ in range(eval(self.config.squeeze_block.split('_x')[1]))
+ ])
+
+ self.decoder = Decoder(channels)
+
+ if self.config.ender:
+ self.dec_end = nn.Sequential(
+ nn.Conv2d(1, 16, 3, 1, 1),
+ nn.Conv2d(16, 1, 3, 1, 1),
+ nn.ReLU(inplace=True),
+ )
+
+ # refine patch-level segmentation
+ if self.config.refine:
+ if self.config.refine == 'itself':
+ self.stem_layer = StemLayer(in_channels=3+1, inter_channels=48, out_channels=3, norm_layer='BN' if self.config.batch_size > 1 else 'LN')
+ else:
+ self.refiner = eval('{}({})'.format(self.config.refine, 'in_channels=3+1'))
+
+ if self.config.freeze_bb:
+ # Freeze the backbone...
+ print(self.named_parameters())
+ for key, value in self.named_parameters():
+ if 'bb.' in key and 'refiner.' not in key:
+ value.requires_grad = False
+
+ def forward_enc(self, x):
+ if self.config.bb in ['vgg16', 'vgg16bn', 'resnet50']:
+ x1 = self.bb.conv1(x); x2 = self.bb.conv2(x1); x3 = self.bb.conv3(x2); x4 = self.bb.conv4(x3)
+ else:
+ x1, x2, x3, x4 = self.bb(x)
+ if self.config.mul_scl_ipt == 'cat':
+ B, C, H, W = x.shape
+ x1_, x2_, x3_, x4_ = self.bb(F.interpolate(x, size=(H//2, W//2), mode='bilinear', align_corners=True))
+ x1 = torch.cat([x1, F.interpolate(x1_, size=x1.shape[2:], mode='bilinear', align_corners=True)], dim=1)
+ x2 = torch.cat([x2, F.interpolate(x2_, size=x2.shape[2:], mode='bilinear', align_corners=True)], dim=1)
+ x3 = torch.cat([x3, F.interpolate(x3_, size=x3.shape[2:], mode='bilinear', align_corners=True)], dim=1)
+ x4 = torch.cat([x4, F.interpolate(x4_, size=x4.shape[2:], mode='bilinear', align_corners=True)], dim=1)
+ elif self.config.mul_scl_ipt == 'add':
+ B, C, H, W = x.shape
+ x1_, x2_, x3_, x4_ = self.bb(F.interpolate(x, size=(H//2, W//2), mode='bilinear', align_corners=True))
+ x1 = x1 + F.interpolate(x1_, size=x1.shape[2:], mode='bilinear', align_corners=True)
+ x2 = x2 + F.interpolate(x2_, size=x2.shape[2:], mode='bilinear', align_corners=True)
+ x3 = x3 + F.interpolate(x3_, size=x3.shape[2:], mode='bilinear', align_corners=True)
+ x4 = x4 + F.interpolate(x4_, size=x4.shape[2:], mode='bilinear', align_corners=True)
+ class_preds = self.cls_head(self.avgpool(x4).view(x4.shape[0], -1)) if self.training and self.config.auxiliary_classification else None
+ if self.config.cxt:
+ x4 = torch.cat(
+ (
+ *[
+ F.interpolate(x1, size=x4.shape[2:], mode='bilinear', align_corners=True),
+ F.interpolate(x2, size=x4.shape[2:], mode='bilinear', align_corners=True),
+ F.interpolate(x3, size=x4.shape[2:], mode='bilinear', align_corners=True),
+ ][-len(self.config.cxt):],
+ x4
+ ),
+ dim=1
+ )
+ return (x1, x2, x3, x4), class_preds
+
+ def forward_ori(self, x):
+ ########## Encoder ##########
+ (x1, x2, x3, x4), class_preds = self.forward_enc(x)
+ if self.config.squeeze_block:
+ x4 = self.squeeze_module(x4)
+ ########## Decoder ##########
+ features = [x, x1, x2, x3, x4]
+ if self.training and self.config.out_ref:
+ features.append(laplacian(torch.mean(x, dim=1).unsqueeze(1), kernel_size=5))
+ scaled_preds = self.decoder(features)
+ return scaled_preds, class_preds
+
+ def forward(self, x):
+ scaled_preds, class_preds = self.forward_ori(x)
+ class_preds_lst = [class_preds]
+ return [scaled_preds, class_preds_lst] if self.training else scaled_preds
+
+
+class Decoder(nn.Module):
+ def __init__(self, channels):
+ super(Decoder, self).__init__()
+ self.config = Config()
+ DecoderBlock = eval(self.config.dec_blk)
+ LateralBlock = eval(self.config.lat_blk)
+
+ if self.config.dec_ipt:
+ self.split = self.config.dec_ipt_split
+ N_dec_ipt = 64
+ DBlock = SimpleConvs
+ ic = 64
+ ipt_cha_opt = 1
+ self.ipt_blk5 = DBlock(2**10*3 if self.split else 3, [N_dec_ipt, channels[0]//8][ipt_cha_opt], inter_channels=ic)
+ self.ipt_blk4 = DBlock(2**8*3 if self.split else 3, [N_dec_ipt, channels[0]//8][ipt_cha_opt], inter_channels=ic)
+ self.ipt_blk3 = DBlock(2**6*3 if self.split else 3, [N_dec_ipt, channels[1]//8][ipt_cha_opt], inter_channels=ic)
+ self.ipt_blk2 = DBlock(2**4*3 if self.split else 3, [N_dec_ipt, channels[2]//8][ipt_cha_opt], inter_channels=ic)
+ self.ipt_blk1 = DBlock(2**0*3 if self.split else 3, [N_dec_ipt, channels[3]//8][ipt_cha_opt], inter_channels=ic)
+ else:
+ self.split = None
+
+ self.decoder_block4 = DecoderBlock(channels[0]+([N_dec_ipt, channels[0]//8][ipt_cha_opt] if self.config.dec_ipt else 0), channels[1])
+ self.decoder_block3 = DecoderBlock(channels[1]+([N_dec_ipt, channels[0]//8][ipt_cha_opt] if self.config.dec_ipt else 0), channels[2])
+ self.decoder_block2 = DecoderBlock(channels[2]+([N_dec_ipt, channels[1]//8][ipt_cha_opt] if self.config.dec_ipt else 0), channels[3])
+ self.decoder_block1 = DecoderBlock(channels[3]+([N_dec_ipt, channels[2]//8][ipt_cha_opt] if self.config.dec_ipt else 0), channels[3]//2)
+ self.conv_out1 = nn.Sequential(nn.Conv2d(channels[3]//2+([N_dec_ipt, channels[3]//8][ipt_cha_opt] if self.config.dec_ipt else 0), 1, 1, 1, 0))
+
+ self.lateral_block4 = LateralBlock(channels[1], channels[1])
+ self.lateral_block3 = LateralBlock(channels[2], channels[2])
+ self.lateral_block2 = LateralBlock(channels[3], channels[3])
+
+ if self.config.ms_supervision:
+ self.conv_ms_spvn_4 = nn.Conv2d(channels[1], 1, 1, 1, 0)
+ self.conv_ms_spvn_3 = nn.Conv2d(channels[2], 1, 1, 1, 0)
+ self.conv_ms_spvn_2 = nn.Conv2d(channels[3], 1, 1, 1, 0)
+
+ if self.config.out_ref:
+ _N = 16
+ self.gdt_convs_4 = nn.Sequential(nn.Conv2d(channels[1], _N, 3, 1, 1), nn.BatchNorm2d(_N) if self.config.batch_size > 1 else nn.Identity(), nn.ReLU(inplace=True))
+ self.gdt_convs_3 = nn.Sequential(nn.Conv2d(channels[2], _N, 3, 1, 1), nn.BatchNorm2d(_N) if self.config.batch_size > 1 else nn.Identity(), nn.ReLU(inplace=True))
+ self.gdt_convs_2 = nn.Sequential(nn.Conv2d(channels[3], _N, 3, 1, 1), nn.BatchNorm2d(_N) if self.config.batch_size > 1 else nn.Identity(), nn.ReLU(inplace=True))
+
+ self.gdt_convs_pred_4 = nn.Sequential(nn.Conv2d(_N, 1, 1, 1, 0))
+ self.gdt_convs_pred_3 = nn.Sequential(nn.Conv2d(_N, 1, 1, 1, 0))
+ self.gdt_convs_pred_2 = nn.Sequential(nn.Conv2d(_N, 1, 1, 1, 0))
+
+ self.gdt_convs_attn_4 = nn.Sequential(nn.Conv2d(_N, 1, 1, 1, 0))
+ self.gdt_convs_attn_3 = nn.Sequential(nn.Conv2d(_N, 1, 1, 1, 0))
+ self.gdt_convs_attn_2 = nn.Sequential(nn.Conv2d(_N, 1, 1, 1, 0))
+
+ def get_patches_batch(self, x, p):
+ _size_h, _size_w = p.shape[2:]
+ patches_batch = []
+ for idx in range(x.shape[0]):
+ columns_x = torch.split(x[idx], split_size_or_sections=_size_w, dim=-1)
+ patches_x = []
+ for column_x in columns_x:
+ patches_x += [p.unsqueeze(0) for p in torch.split(column_x, split_size_or_sections=_size_h, dim=-2)]
+ patch_sample = torch.cat(patches_x, dim=1)
+ patches_batch.append(patch_sample)
+ return torch.cat(patches_batch, dim=0)
+
+ def forward(self, features):
+ if self.training and self.config.out_ref:
+ outs_gdt_pred = []
+ outs_gdt_label = []
+ x, x1, x2, x3, x4, gdt_gt = features
+ else:
+ x, x1, x2, x3, x4 = features
+ outs = []
+
+ if self.config.dec_ipt:
+ patches_batch = self.get_patches_batch(x, x4) if self.split else x
+ x4 = torch.cat((x4, self.ipt_blk5(F.interpolate(patches_batch, size=x4.shape[2:], mode='bilinear', align_corners=True))), 1)
+ p4 = self.decoder_block4(x4)
+ m4 = self.conv_ms_spvn_4(p4) if self.config.ms_supervision and self.training else None
+ if self.config.out_ref:
+ p4_gdt = self.gdt_convs_4(p4)
+ if self.training:
+ # >> GT:
+ m4_dia = m4
+ gdt_label_main_4 = gdt_gt * F.interpolate(m4_dia, size=gdt_gt.shape[2:], mode='bilinear', align_corners=True)
+ outs_gdt_label.append(gdt_label_main_4)
+ # >> Pred:
+ gdt_pred_4 = self.gdt_convs_pred_4(p4_gdt)
+ outs_gdt_pred.append(gdt_pred_4)
+ gdt_attn_4 = self.gdt_convs_attn_4(p4_gdt).sigmoid()
+ # >> Finally:
+ p4 = p4 * gdt_attn_4
+ _p4 = F.interpolate(p4, size=x3.shape[2:], mode='bilinear', align_corners=True)
+ _p3 = _p4 + self.lateral_block4(x3)
+
+ if self.config.dec_ipt:
+ patches_batch = self.get_patches_batch(x, _p3) if self.split else x
+ _p3 = torch.cat((_p3, self.ipt_blk4(F.interpolate(patches_batch, size=x3.shape[2:], mode='bilinear', align_corners=True))), 1)
+ p3 = self.decoder_block3(_p3)
+ m3 = self.conv_ms_spvn_3(p3) if self.config.ms_supervision and self.training else None
+ if self.config.out_ref:
+ p3_gdt = self.gdt_convs_3(p3)
+ if self.training:
+ # >> GT:
+ # m3 --dilation--> m3_dia
+ # G_3^gt * m3_dia --> G_3^m, which is the label of gradient
+ m3_dia = m3
+ gdt_label_main_3 = gdt_gt * F.interpolate(m3_dia, size=gdt_gt.shape[2:], mode='bilinear', align_corners=True)
+ outs_gdt_label.append(gdt_label_main_3)
+ # >> Pred:
+ # p3 --conv--BN--> F_3^G, where F_3^G predicts the \hat{G_3} with xx
+ # F_3^G --sigmoid--> A_3^G
+ gdt_pred_3 = self.gdt_convs_pred_3(p3_gdt)
+ outs_gdt_pred.append(gdt_pred_3)
+ gdt_attn_3 = self.gdt_convs_attn_3(p3_gdt).sigmoid()
+ # >> Finally:
+ # p3 = p3 * A_3^G
+ p3 = p3 * gdt_attn_3
+ _p3 = F.interpolate(p3, size=x2.shape[2:], mode='bilinear', align_corners=True)
+ _p2 = _p3 + self.lateral_block3(x2)
+
+ if self.config.dec_ipt:
+ patches_batch = self.get_patches_batch(x, _p2) if self.split else x
+ _p2 = torch.cat((_p2, self.ipt_blk3(F.interpolate(patches_batch, size=x2.shape[2:], mode='bilinear', align_corners=True))), 1)
+ p2 = self.decoder_block2(_p2)
+ m2 = self.conv_ms_spvn_2(p2) if self.config.ms_supervision and self.training else None
+ if self.config.out_ref:
+ p2_gdt = self.gdt_convs_2(p2)
+ if self.training:
+ # >> GT:
+ m2_dia = m2
+ gdt_label_main_2 = gdt_gt * F.interpolate(m2_dia, size=gdt_gt.shape[2:], mode='bilinear', align_corners=True)
+ outs_gdt_label.append(gdt_label_main_2)
+ # >> Pred:
+ gdt_pred_2 = self.gdt_convs_pred_2(p2_gdt)
+ outs_gdt_pred.append(gdt_pred_2)
+ gdt_attn_2 = self.gdt_convs_attn_2(p2_gdt).sigmoid()
+ # >> Finally:
+ p2 = p2 * gdt_attn_2
+ _p2 = F.interpolate(p2, size=x1.shape[2:], mode='bilinear', align_corners=True)
+ _p1 = _p2 + self.lateral_block2(x1)
+
+ if self.config.dec_ipt:
+ patches_batch = self.get_patches_batch(x, _p1) if self.split else x
+ _p1 = torch.cat((_p1, self.ipt_blk2(F.interpolate(patches_batch, size=x1.shape[2:], mode='bilinear', align_corners=True))), 1)
+ _p1 = self.decoder_block1(_p1)
+ _p1 = F.interpolate(_p1, size=x.shape[2:], mode='bilinear', align_corners=True)
+
+ if self.config.dec_ipt:
+ patches_batch = self.get_patches_batch(x, _p1) if self.split else x
+ _p1 = torch.cat((_p1, self.ipt_blk1(F.interpolate(patches_batch, size=x.shape[2:], mode='bilinear', align_corners=True))), 1)
+ p1_out = self.conv_out1(_p1)
+
+ if self.config.ms_supervision and self.training:
+ outs.append(m4)
+ outs.append(m3)
+ outs.append(m2)
+ outs.append(p1_out)
+ return outs if not (self.config.out_ref and self.training) else ([outs_gdt_pred, outs_gdt_label], outs)
+
+
+class SimpleConvs(nn.Module):
+ def __init__(
+ self, in_channels: int, out_channels: int, inter_channels=64
+ ) -> None:
+ super().__init__()
+ self.conv1 = nn.Conv2d(in_channels, inter_channels, 3, 1, 1)
+ self.conv_out = nn.Conv2d(inter_channels, out_channels, 3, 1, 1)
+
+ def forward(self, x):
+ return self.conv_out(self.conv1(x))
diff --git a/py/BiRefNet_v2/models/modules/aspp.py b/py/BiRefNet_v2/models/modules/aspp.py
new file mode 100644
index 0000000..3c4f87e
--- /dev/null
+++ b/py/BiRefNet_v2/models/modules/aspp.py
@@ -0,0 +1,120 @@
+import torch
+import torch.nn as nn
+import torch.nn.functional as F
+
+from ...models.modules.deform_conv import DeformableConv2d
+from ...config import Config
+
+
+config = Config()
+
+
+class _ASPPModule(nn.Module):
+ def __init__(self, in_channels, planes, kernel_size, padding, dilation):
+ super(_ASPPModule, self).__init__()
+ self.atrous_conv = nn.Conv2d(in_channels, planes, kernel_size=kernel_size,
+ stride=1, padding=padding, dilation=dilation, bias=False)
+ self.bn = nn.BatchNorm2d(planes) if config.batch_size > 1 else nn.Identity()
+ self.relu = nn.ReLU(inplace=True)
+
+ def forward(self, x):
+ x = self.atrous_conv(x)
+ x = self.bn(x)
+
+ return self.relu(x)
+
+
+class ASPP(nn.Module):
+ def __init__(self, in_channels=64, out_channels=None, output_stride=16):
+ super(ASPP, self).__init__()
+ self.down_scale = 1
+ if out_channels is None:
+ out_channels = in_channels
+ self.in_channelster = 256 // self.down_scale
+ if output_stride == 16:
+ dilations = [1, 6, 12, 18]
+ elif output_stride == 8:
+ dilations = [1, 12, 24, 36]
+ else:
+ raise NotImplementedError
+
+ self.aspp1 = _ASPPModule(in_channels, self.in_channelster, 1, padding=0, dilation=dilations[0])
+ self.aspp2 = _ASPPModule(in_channels, self.in_channelster, 3, padding=dilations[1], dilation=dilations[1])
+ self.aspp3 = _ASPPModule(in_channels, self.in_channelster, 3, padding=dilations[2], dilation=dilations[2])
+ self.aspp4 = _ASPPModule(in_channels, self.in_channelster, 3, padding=dilations[3], dilation=dilations[3])
+
+ self.global_avg_pool = nn.Sequential(nn.AdaptiveAvgPool2d((1, 1)),
+ nn.Conv2d(in_channels, self.in_channelster, 1, stride=1, bias=False),
+ nn.BatchNorm2d(self.in_channelster) if config.batch_size > 1 else nn.Identity(),
+ nn.ReLU(inplace=True))
+ self.conv1 = nn.Conv2d(self.in_channelster * 5, out_channels, 1, bias=False)
+ self.bn1 = nn.BatchNorm2d(out_channels) if config.batch_size > 1 else nn.Identity()
+ self.relu = nn.ReLU(inplace=True)
+ self.dropout = nn.Dropout(0.5)
+
+ def forward(self, x):
+ x1 = self.aspp1(x)
+ x2 = self.aspp2(x)
+ x3 = self.aspp3(x)
+ x4 = self.aspp4(x)
+ x5 = self.global_avg_pool(x)
+ x5 = F.interpolate(x5, size=x1.size()[2:], mode='bilinear', align_corners=True)
+ x = torch.cat((x1, x2, x3, x4, x5), dim=1)
+
+ x = self.conv1(x)
+ x = self.bn1(x)
+ x = self.relu(x)
+
+ return self.dropout(x)
+
+
+##################### Deformable
+class _ASPPModuleDeformable(nn.Module):
+ def __init__(self, in_channels, planes, kernel_size, padding):
+ super(_ASPPModuleDeformable, self).__init__()
+ self.atrous_conv = DeformableConv2d(in_channels, planes, kernel_size=kernel_size,
+ stride=1, padding=padding, bias=False)
+ self.bn = nn.BatchNorm2d(planes) if config.batch_size > 1 else nn.Identity()
+ self.relu = nn.ReLU(inplace=True)
+
+ def forward(self, x):
+ x = self.atrous_conv(x)
+ x = self.bn(x)
+
+ return self.relu(x)
+
+
+class ASPPDeformable(nn.Module):
+ def __init__(self, in_channels, out_channels=None, parallel_block_sizes=[1, 3, 7]):
+ super(ASPPDeformable, self).__init__()
+ self.down_scale = 1
+ if out_channels is None:
+ out_channels = in_channels
+ self.in_channelster = 256 // self.down_scale
+
+ self.aspp1 = _ASPPModuleDeformable(in_channels, self.in_channelster, 1, padding=0)
+ self.aspp_deforms = nn.ModuleList([
+ _ASPPModuleDeformable(in_channels, self.in_channelster, conv_size, padding=int(conv_size//2)) for conv_size in parallel_block_sizes
+ ])
+
+ self.global_avg_pool = nn.Sequential(nn.AdaptiveAvgPool2d((1, 1)),
+ nn.Conv2d(in_channels, self.in_channelster, 1, stride=1, bias=False),
+ nn.BatchNorm2d(self.in_channelster) if config.batch_size > 1 else nn.Identity(),
+ nn.ReLU(inplace=True))
+ self.conv1 = nn.Conv2d(self.in_channelster * (2 + len(self.aspp_deforms)), out_channels, 1, bias=False)
+ self.bn1 = nn.BatchNorm2d(out_channels) if config.batch_size > 1 else nn.Identity()
+ self.relu = nn.ReLU(inplace=True)
+ self.dropout = nn.Dropout(0.5)
+
+ def forward(self, x):
+ x1 = self.aspp1(x)
+ x_aspp_deforms = [aspp_deform(x) for aspp_deform in self.aspp_deforms]
+ x5 = self.global_avg_pool(x)
+ x5 = F.interpolate(x5, size=x1.size()[2:], mode='bilinear', align_corners=True)
+ x = torch.cat((x1, *x_aspp_deforms, x5), dim=1)
+
+ x = self.conv1(x)
+ x = self.bn1(x)
+ x = self.relu(x)
+
+ return self.dropout(x)
diff --git a/py/BiRefNet_v2/models/modules/decoder_blocks.py b/py/BiRefNet_v2/models/modules/decoder_blocks.py
new file mode 100644
index 0000000..32a0b6a
--- /dev/null
+++ b/py/BiRefNet_v2/models/modules/decoder_blocks.py
@@ -0,0 +1,66 @@
+import torch
+import torch.nn as nn
+
+from ...models.modules.aspp import ASPP, ASPPDeformable
+from ...config import Config
+
+
+config = Config()
+
+
+class BasicDecBlk(nn.Module):
+ def __init__(self, in_channels=64, out_channels=64, inter_channels=64):
+ super(BasicDecBlk, self).__init__()
+ inter_channels = in_channels // 4 if config.dec_channels_inter == 'adap' else 64
+ self.conv_in = nn.Conv2d(in_channels, inter_channels, 3, 1, padding=1)
+ self.relu_in = nn.ReLU(inplace=True)
+ if config.dec_att == 'ASPP':
+ self.dec_att = ASPP(in_channels=inter_channels)
+ elif config.dec_att == 'ASPPDeformable':
+ self.dec_att = ASPPDeformable(in_channels=inter_channels)
+ self.conv_out = nn.Conv2d(inter_channels, out_channels, 3, 1, padding=1)
+ self.bn_in = nn.BatchNorm2d(inter_channels) if config.batch_size > 1 else nn.Identity()
+ self.bn_out = nn.BatchNorm2d(out_channels) if config.batch_size > 1 else nn.Identity()
+
+ def forward(self, x):
+ x = self.conv_in(x)
+ x = self.bn_in(x)
+ x = self.relu_in(x)
+ if hasattr(self, 'dec_att'):
+ x = self.dec_att(x)
+ x = self.conv_out(x)
+ x = self.bn_out(x)
+ return x
+
+
+class ResBlk(nn.Module):
+ def __init__(self, in_channels=64, out_channels=None, inter_channels=64):
+ super(ResBlk, self).__init__()
+ if out_channels is None:
+ out_channels = in_channels
+ inter_channels = in_channels // 4 if config.dec_channels_inter == 'adap' else 64
+
+ self.conv_in = nn.Conv2d(in_channels, inter_channels, 3, 1, padding=1)
+ self.bn_in = nn.BatchNorm2d(inter_channels) if config.batch_size > 1 else nn.Identity()
+ self.relu_in = nn.ReLU(inplace=True)
+
+ if config.dec_att == 'ASPP':
+ self.dec_att = ASPP(in_channels=inter_channels)
+ elif config.dec_att == 'ASPPDeformable':
+ self.dec_att = ASPPDeformable(in_channels=inter_channels)
+
+ self.conv_out = nn.Conv2d(inter_channels, out_channels, 3, 1, padding=1)
+ self.bn_out = nn.BatchNorm2d(out_channels) if config.batch_size > 1 else nn.Identity()
+
+ self.conv_resi = nn.Conv2d(in_channels, out_channels, 1, 1, 0)
+
+ def forward(self, x):
+ _x = self.conv_resi(x)
+ x = self.conv_in(x)
+ x = self.bn_in(x)
+ x = self.relu_in(x)
+ if hasattr(self, 'dec_att'):
+ x = self.dec_att(x)
+ x = self.conv_out(x)
+ x = self.bn_out(x)
+ return x + _x
diff --git a/py/BiRefNet_v2/models/modules/deform_conv.py b/py/BiRefNet_v2/models/modules/deform_conv.py
new file mode 100644
index 0000000..43f5e57
--- /dev/null
+++ b/py/BiRefNet_v2/models/modules/deform_conv.py
@@ -0,0 +1,66 @@
+import torch
+import torch.nn as nn
+from torchvision.ops import deform_conv2d
+
+
+class DeformableConv2d(nn.Module):
+ def __init__(self,
+ in_channels,
+ out_channels,
+ kernel_size=3,
+ stride=1,
+ padding=1,
+ bias=False):
+
+ super(DeformableConv2d, self).__init__()
+
+ assert type(kernel_size) == tuple or type(kernel_size) == int
+
+ kernel_size = kernel_size if type(kernel_size) == tuple else (kernel_size, kernel_size)
+ self.stride = stride if type(stride) == tuple else (stride, stride)
+ self.padding = padding
+
+ self.offset_conv = nn.Conv2d(in_channels,
+ 2 * kernel_size[0] * kernel_size[1],
+ kernel_size=kernel_size,
+ stride=stride,
+ padding=self.padding,
+ bias=True)
+
+ nn.init.constant_(self.offset_conv.weight, 0.)
+ nn.init.constant_(self.offset_conv.bias, 0.)
+
+ self.modulator_conv = nn.Conv2d(in_channels,
+ 1 * kernel_size[0] * kernel_size[1],
+ kernel_size=kernel_size,
+ stride=stride,
+ padding=self.padding,
+ bias=True)
+
+ nn.init.constant_(self.modulator_conv.weight, 0.)
+ nn.init.constant_(self.modulator_conv.bias, 0.)
+
+ self.regular_conv = nn.Conv2d(in_channels,
+ out_channels=out_channels,
+ kernel_size=kernel_size,
+ stride=stride,
+ padding=self.padding,
+ bias=bias)
+
+ def forward(self, x):
+ #h, w = x.shape[2:]
+ #max_offset = max(h, w)/4.
+
+ offset = self.offset_conv(x)#.clamp(-max_offset, max_offset)
+ modulator = 2. * torch.sigmoid(self.modulator_conv(x))
+
+ x = deform_conv2d(
+ input=x,
+ offset=offset,
+ weight=self.regular_conv.weight,
+ bias=self.regular_conv.bias,
+ padding=self.padding,
+ mask=modulator,
+ stride=self.stride,
+ )
+ return x
diff --git a/py/BiRefNet_v2/models/modules/lateral_blocks.py b/py/BiRefNet_v2/models/modules/lateral_blocks.py
new file mode 100644
index 0000000..de907ac
--- /dev/null
+++ b/py/BiRefNet_v2/models/modules/lateral_blocks.py
@@ -0,0 +1,21 @@
+import numpy as np
+import torch
+import torch.nn as nn
+import torch.nn.functional as F
+from functools import partial
+
+from ...config import Config
+
+
+config = Config()
+
+
+class BasicLatBlk(nn.Module):
+ def __init__(self, in_channels=64, out_channels=64, inter_channels=64):
+ super(BasicLatBlk, self).__init__()
+ inter_channels = in_channels // 4 if config.dec_channels_inter == 'adap' else 64
+ self.conv = nn.Conv2d(in_channels, out_channels, 1, 1, 0)
+
+ def forward(self, x):
+ x = self.conv(x)
+ return x
diff --git a/py/BiRefNet_v2/models/modules/mlp.py b/py/BiRefNet_v2/models/modules/mlp.py
new file mode 100644
index 0000000..a383459
--- /dev/null
+++ b/py/BiRefNet_v2/models/modules/mlp.py
@@ -0,0 +1,118 @@
+import torch
+import torch.nn as nn
+from functools import partial
+
+from timm.models.layers import DropPath, to_2tuple, trunc_normal_
+from timm.models import register_model
+
+import math
+
+
+class MLPLayer(nn.Module):
+ def __init__(self, in_features, hidden_features=None, out_features=None, act_layer=nn.GELU, drop=0.):
+ super().__init__()
+ out_features = out_features or in_features
+ hidden_features = hidden_features or in_features
+ self.fc1 = nn.Linear(in_features, hidden_features)
+ self.act = act_layer()
+ self.fc2 = nn.Linear(hidden_features, out_features)
+ self.drop = nn.Dropout(drop)
+
+ def forward(self, x):
+ x = self.fc1(x)
+ x = self.act(x)
+ x = self.drop(x)
+ x = self.fc2(x)
+ x = self.drop(x)
+ return x
+
+
+class Attention(nn.Module):
+ def __init__(self, dim, num_heads=8, qkv_bias=False, qk_scale=None, attn_drop=0., proj_drop=0., sr_ratio=1):
+ super().__init__()
+ assert dim % num_heads == 0, f"dim {dim} should be divided by num_heads {num_heads}."
+
+ self.dim = dim
+ self.num_heads = num_heads
+ head_dim = dim // num_heads
+ self.scale = qk_scale or head_dim ** -0.5
+
+ self.q = nn.Linear(dim, dim, bias=qkv_bias)
+ self.kv = nn.Linear(dim, dim * 2, bias=qkv_bias)
+ self.attn_drop = nn.Dropout(attn_drop)
+ self.proj = nn.Linear(dim, dim)
+ self.proj_drop = nn.Dropout(proj_drop)
+
+ self.sr_ratio = sr_ratio
+ if sr_ratio > 1:
+ self.sr = nn.Conv2d(dim, dim, kernel_size=sr_ratio, stride=sr_ratio)
+ self.norm = nn.LayerNorm(dim)
+
+ def forward(self, x, H, W):
+ B, N, C = x.shape
+ q = self.q(x).reshape(B, N, self.num_heads, C // self.num_heads).permute(0, 2, 1, 3)
+
+ if self.sr_ratio > 1:
+ x_ = x.permute(0, 2, 1).reshape(B, C, H, W)
+ x_ = self.sr(x_).reshape(B, C, -1).permute(0, 2, 1)
+ x_ = self.norm(x_)
+ kv = self.kv(x_).reshape(B, -1, 2, self.num_heads, C // self.num_heads).permute(2, 0, 3, 1, 4)
+ else:
+ kv = self.kv(x).reshape(B, -1, 2, self.num_heads, C // self.num_heads).permute(2, 0, 3, 1, 4)
+ k, v = kv[0], kv[1]
+
+ attn = (q @ k.transpose(-2, -1)) * self.scale
+ attn = attn.softmax(dim=-1)
+ attn = self.attn_drop(attn)
+
+ x = (attn @ v).transpose(1, 2).reshape(B, N, C)
+ x = self.proj(x)
+ x = self.proj_drop(x)
+ return x
+
+
+class Block(nn.Module):
+ def __init__(self, dim, num_heads, mlp_ratio=4., qkv_bias=False, qk_scale=None, drop=0., attn_drop=0.,
+ drop_path=0., act_layer=nn.GELU, norm_layer=nn.LayerNorm, sr_ratio=1):
+ super().__init__()
+ self.norm1 = norm_layer(dim)
+ self.attn = Attention(
+ dim,
+ num_heads=num_heads, qkv_bias=qkv_bias, qk_scale=qk_scale,
+ attn_drop=attn_drop, proj_drop=drop, sr_ratio=sr_ratio)
+ # NOTE: drop path for stochastic depth, we shall see if this is better than dropout here
+ self.drop_path = DropPath(drop_path) if drop_path > 0. else nn.Identity()
+ self.norm2 = norm_layer(dim)
+ mlp_hidden_dim = int(dim * mlp_ratio)
+ self.mlp = MLPLayer(in_features=dim, hidden_features=mlp_hidden_dim, act_layer=act_layer, drop=drop)
+
+ def forward(self, x, H, W):
+ x = x + self.drop_path(self.attn(self.norm1(x), H, W))
+ x = x + self.drop_path(self.mlp(self.norm2(x), H, W))
+ return x
+
+
+class OverlapPatchEmbed(nn.Module):
+ """ Image to Patch Embedding
+ """
+
+ def __init__(self, img_size=224, patch_size=7, stride=4, in_channels=3, embed_dim=768):
+ super().__init__()
+ img_size = to_2tuple(img_size)
+ patch_size = to_2tuple(patch_size)
+
+ self.img_size = img_size
+ self.patch_size = patch_size
+ self.H, self.W = img_size[0] // patch_size[0], img_size[1] // patch_size[1]
+ self.num_patches = self.H * self.W
+ self.proj = nn.Conv2d(in_channels, embed_dim, kernel_size=patch_size, stride=stride,
+ padding=(patch_size[0] // 2, patch_size[1] // 2))
+ self.norm = nn.LayerNorm(embed_dim)
+
+ def forward(self, x):
+ x = self.proj(x)
+ _, _, H, W = x.shape
+ x = x.flatten(2).transpose(1, 2)
+ x = self.norm(x)
+ return x, H, W
+
diff --git a/py/BiRefNet_v2/models/modules/prompt_encoder.py b/py/BiRefNet_v2/models/modules/prompt_encoder.py
new file mode 100644
index 0000000..23ce18c
--- /dev/null
+++ b/py/BiRefNet_v2/models/modules/prompt_encoder.py
@@ -0,0 +1,222 @@
+import numpy as np
+import torch
+import torch.nn as nn
+from typing import Any, Optional, Tuple, Type
+
+
+class PromptEncoder(nn.Module):
+ def __init__(
+ self,
+ embed_dim=256,
+ image_embedding_size=1024,
+ input_image_size=(1024, 1024),
+ mask_in_chans=16,
+ activation=nn.GELU
+ ) -> None:
+ super().__init__()
+ """
+ Codes are partially from SAM: https://github.com/facebookresearch/segment-anything/blob/6fdee8f2727f4506cfbbe553e23b895e27956588/segment_anything/modeling/prompt_encoder.py.
+
+ Arguments:
+ embed_dim (int): The prompts' embedding dimension
+ image_embedding_size (tuple(int, int)): The spatial size of the
+ image embedding, as (H, W).
+ input_image_size (int): The padded size of the image as input
+ to the image encoder, as (H, W).
+ mask_in_chans (int): The number of hidden channels used for
+ encoding input masks.
+ activation (nn.Module): The activation to use when encoding
+ input masks.
+ """
+ super().__init__()
+ self.embed_dim = embed_dim
+ self.input_image_size = input_image_size
+ self.image_embedding_size = image_embedding_size
+ self.pe_layer = PositionEmbeddingRandom(embed_dim // 2)
+
+ self.num_point_embeddings: int = 4 # pos/neg point + 2 box corners
+ point_embeddings = [nn.Embedding(1, embed_dim) for i in range(self.num_point_embeddings)]
+ self.point_embeddings = nn.ModuleList(point_embeddings)
+ self.not_a_point_embed = nn.Embedding(1, embed_dim)
+
+ self.mask_input_size = (4 * image_embedding_size[0], 4 * image_embedding_size[1])
+ self.mask_downscaling = nn.Sequential(
+ nn.Conv2d(1, mask_in_chans // 4, kernel_size=2, stride=2),
+ LayerNorm2d(mask_in_chans // 4),
+ activation(),
+ nn.Conv2d(mask_in_chans // 4, mask_in_chans, kernel_size=2, stride=2),
+ LayerNorm2d(mask_in_chans),
+ activation(),
+ nn.Conv2d(mask_in_chans, embed_dim, kernel_size=1),
+ )
+ self.no_mask_embed = nn.Embedding(1, embed_dim)
+
+ def get_dense_pe(self) -> torch.Tensor:
+ """
+ Returns the positional encoding used to encode point prompts,
+ applied to a dense set of points the shape of the image encoding.
+
+ Returns:
+ torch.Tensor: Positional encoding with shape
+ 1x(embed_dim)x(embedding_h)x(embedding_w)
+ """
+ return self.pe_layer(self.image_embedding_size).unsqueeze(0)
+
+ def _embed_points(
+ self,
+ points: torch.Tensor,
+ labels: torch.Tensor,
+ pad: bool,
+ ) -> torch.Tensor:
+ """Embeds point prompts."""
+ points = points + 0.5 # Shift to center of pixel
+ if pad:
+ padding_point = torch.zeros((points.shape[0], 1, 2), device=points.device)
+ padding_label = -torch.ones((labels.shape[0], 1), device=labels.device)
+ points = torch.cat([points, padding_point], dim=1)
+ labels = torch.cat([labels, padding_label], dim=1)
+ point_embedding = self.pe_layer.forward_with_coords(points, self.input_image_size)
+ point_embedding[labels == -1] = 0.0
+ point_embedding[labels == -1] += self.not_a_point_embed.weight
+ point_embedding[labels == 0] += self.point_embeddings[0].weight
+ point_embedding[labels == 1] += self.point_embeddings[1].weight
+ return point_embedding
+
+ def _embed_boxes(self, boxes: torch.Tensor) -> torch.Tensor:
+ """Embeds box prompts."""
+ boxes = boxes + 0.5 # Shift to center of pixel
+ coords = boxes.reshape(-1, 2, 2)
+ corner_embedding = self.pe_layer.forward_with_coords(coords, self.input_image_size)
+ corner_embedding[:, 0, :] += self.point_embeddings[2].weight
+ corner_embedding[:, 1, :] += self.point_embeddings[3].weight
+ return corner_embedding
+
+ def _embed_masks(self, masks: torch.Tensor) -> torch.Tensor:
+ """Embeds mask inputs."""
+ mask_embedding = self.mask_downscaling(masks)
+ return mask_embedding
+
+ def _get_batch_size(
+ self,
+ points: Optional[Tuple[torch.Tensor, torch.Tensor]],
+ boxes: Optional[torch.Tensor],
+ masks: Optional[torch.Tensor],
+ ) -> int:
+ """
+ Gets the batch size of the output given the batch size of the input prompts.
+ """
+ if points is not None:
+ return points[0].shape[0]
+ elif boxes is not None:
+ return boxes.shape[0]
+ elif masks is not None:
+ return masks.shape[0]
+ else:
+ return 1
+
+ def _get_device(self) -> torch.device:
+ return self.point_embeddings[0].weight.device
+
+ def forward(
+ self,
+ points: Optional[Tuple[torch.Tensor, torch.Tensor]],
+ boxes: Optional[torch.Tensor],
+ masks: Optional[torch.Tensor],
+ ) -> Tuple[torch.Tensor, torch.Tensor]:
+ """
+ Embeds different types of prompts, returning both sparse and dense
+ embeddings.
+
+ Arguments:
+ points (tuple(torch.Tensor, torch.Tensor) or none): point coordinates
+ and labels to embed.
+ boxes (torch.Tensor or none): boxes to embed
+ masks (torch.Tensor or none): masks to embed
+
+ Returns:
+ torch.Tensor: sparse embeddings for the points and boxes, with shape
+ BxNx(embed_dim), where N is determined by the number of input points
+ and boxes.
+ torch.Tensor: dense embeddings for the masks, in the shape
+ Bx(embed_dim)x(embed_H)x(embed_W)
+ """
+ bs = self._get_batch_size(points, boxes, masks)
+ sparse_embeddings = torch.empty((bs, 0, self.embed_dim), device=self._get_device())
+ if points is not None:
+ coords, labels = points
+ point_embeddings = self._embed_points(coords, labels, pad=(boxes is None))
+ sparse_embeddings = torch.cat([sparse_embeddings, point_embeddings], dim=1)
+ if boxes is not None:
+ box_embeddings = self._embed_boxes(boxes)
+ sparse_embeddings = torch.cat([sparse_embeddings, box_embeddings], dim=1)
+
+ if masks is not None:
+ dense_embeddings = self._embed_masks(masks)
+ else:
+ dense_embeddings = self.no_mask_embed.weight.reshape(1, -1, 1, 1).expand(
+ bs, -1, self.image_embedding_size[0], self.image_embedding_size[1]
+ )
+
+ return sparse_embeddings, dense_embeddings
+
+
+class PositionEmbeddingRandom(nn.Module):
+ """
+ Positional encoding using random spatial frequencies.
+ """
+
+ def __init__(self, num_pos_feats: int = 64, scale: Optional[float] = None) -> None:
+ super().__init__()
+ if scale is None or scale <= 0.0:
+ scale = 1.0
+ self.register_buffer(
+ "positional_encoding_gaussian_matrix",
+ scale * torch.randn((2, num_pos_feats)),
+ )
+
+ def _pe_encoding(self, coords: torch.Tensor) -> torch.Tensor:
+ """Positionally encode points that are normalized to [0,1]."""
+ # assuming coords are in [0, 1]^2 square and have d_1 x ... x d_n x 2 shape
+ coords = 2 * coords - 1
+ coords = coords @ self.positional_encoding_gaussian_matrix
+ coords = 2 * np.pi * coords
+ # outputs d_1 x ... x d_n x C shape
+ return torch.cat([torch.sin(coords), torch.cos(coords)], dim=-1)
+
+ def forward(self, size: Tuple[int, int]) -> torch.Tensor:
+ """Generate positional encoding for a grid of the specified size."""
+ h, w = size
+ device: Any = self.positional_encoding_gaussian_matrix.device
+ grid = torch.ones((h, w), device=device, dtype=torch.float32)
+ y_embed = grid.cumsum(dim=0) - 0.5
+ x_embed = grid.cumsum(dim=1) - 0.5
+ y_embed = y_embed / h
+ x_embed = x_embed / w
+
+ pe = self._pe_encoding(torch.stack([x_embed, y_embed], dim=-1))
+ return pe.permute(2, 0, 1) # C x H x W
+
+ def forward_with_coords(
+ self, coords_input: torch.Tensor, image_size: Tuple[int, int]
+ ) -> torch.Tensor:
+ """Positionally encode points that are not normalized to [0,1]."""
+ coords = coords_input.clone()
+ coords[:, :, 0] = coords[:, :, 0] / image_size[1]
+ coords[:, :, 1] = coords[:, :, 1] / image_size[0]
+ return self._pe_encoding(coords.to(torch.float)) # B x N x C
+
+
+class LayerNorm2d(nn.Module):
+ def __init__(self, num_channels: int, eps: float = 1e-6) -> None:
+ super().__init__()
+ self.weight = nn.Parameter(torch.ones(num_channels))
+ self.bias = nn.Parameter(torch.zeros(num_channels))
+ self.eps = eps
+
+ def forward(self, x: torch.Tensor) -> torch.Tensor:
+ u = x.mean(1, keepdim=True)
+ s = (x - u).pow(2).mean(1, keepdim=True)
+ x = (x - u) / torch.sqrt(s + self.eps)
+ x = self.weight[:, None, None] * x + self.bias[:, None, None]
+ return x
+
diff --git a/py/BiRefNet_v2/models/modules/utils.py b/py/BiRefNet_v2/models/modules/utils.py
new file mode 100644
index 0000000..59bd912
--- /dev/null
+++ b/py/BiRefNet_v2/models/modules/utils.py
@@ -0,0 +1,54 @@
+import torch.nn as nn
+
+
+def build_act_layer(act_layer):
+ if act_layer == 'ReLU':
+ return nn.ReLU(inplace=True)
+ elif act_layer == 'SiLU':
+ return nn.SiLU(inplace=True)
+ elif act_layer == 'GELU':
+ return nn.GELU()
+
+ raise NotImplementedError(f'build_act_layer does not support {act_layer}')
+
+
+def build_norm_layer(dim,
+ norm_layer,
+ in_format='channels_last',
+ out_format='channels_last',
+ eps=1e-6):
+ layers = []
+ if norm_layer == 'BN':
+ if in_format == 'channels_last':
+ layers.append(to_channels_first())
+ layers.append(nn.BatchNorm2d(dim))
+ if out_format == 'channels_last':
+ layers.append(to_channels_last())
+ elif norm_layer == 'LN':
+ if in_format == 'channels_first':
+ layers.append(to_channels_last())
+ layers.append(nn.LayerNorm(dim, eps=eps))
+ if out_format == 'channels_first':
+ layers.append(to_channels_first())
+ else:
+ raise NotImplementedError(
+ f'build_norm_layer does not support {norm_layer}')
+ return nn.Sequential(*layers)
+
+
+class to_channels_first(nn.Module):
+
+ def __init__(self):
+ super().__init__()
+
+ def forward(self, x):
+ return x.permute(0, 3, 1, 2)
+
+
+class to_channels_last(nn.Module):
+
+ def __init__(self):
+ super().__init__()
+
+ def forward(self, x):
+ return x.permute(0, 2, 3, 1)
diff --git a/py/BiRefNet_v2/models/refinement/refiner.py b/py/BiRefNet_v2/models/refinement/refiner.py
new file mode 100644
index 0000000..f63ad28
--- /dev/null
+++ b/py/BiRefNet_v2/models/refinement/refiner.py
@@ -0,0 +1,252 @@
+import torch
+import torch.nn as nn
+from collections import OrderedDict
+import torch
+import torch.nn as nn
+import torch.nn.functional as F
+from torchvision.models import vgg16, vgg16_bn
+from torchvision.models import resnet50
+
+from ...config import Config
+from ...dataset import class_labels_TR_sorted
+from ...models.backbones.build_backbone import build_backbone
+from ...models.modules.decoder_blocks import BasicDecBlk
+from ...models.modules.lateral_blocks import BasicLatBlk
+from ...models.refinement.stem_layer import StemLayer
+
+
+class RefinerPVTInChannels4(nn.Module):
+ def __init__(self, in_channels=3+1):
+ super(RefinerPVTInChannels4, self).__init__()
+ self.config = Config()
+ self.epoch = 1
+ self.bb = build_backbone(self.config.bb, params_settings='in_channels=4')
+
+ lateral_channels_in_collection = {
+ 'vgg16': [512, 256, 128, 64], 'vgg16bn': [512, 256, 128, 64], 'resnet50': [1024, 512, 256, 64],
+ 'pvt_v2_b2': [512, 320, 128, 64], 'pvt_v2_b5': [512, 320, 128, 64],
+ 'swin_v1_b': [1024, 512, 256, 128], 'swin_v1_l': [1536, 768, 384, 192],
+ }
+ channels = lateral_channels_in_collection[self.config.bb]
+ self.squeeze_module = BasicDecBlk(channels[0], channels[0])
+
+ self.decoder = Decoder(channels)
+
+ if 0:
+ for key, value in self.named_parameters():
+ if 'bb.' in key:
+ value.requires_grad = False
+
+ def forward(self, x):
+ if isinstance(x, list):
+ x = torch.cat(x, dim=1)
+ ########## Encoder ##########
+ if self.config.bb in ['vgg16', 'vgg16bn', 'resnet50']:
+ x1 = self.bb.conv1(x)
+ x2 = self.bb.conv2(x1)
+ x3 = self.bb.conv3(x2)
+ x4 = self.bb.conv4(x3)
+ else:
+ x1, x2, x3, x4 = self.bb(x)
+
+ x4 = self.squeeze_module(x4)
+
+ ########## Decoder ##########
+
+ features = [x, x1, x2, x3, x4]
+ scaled_preds = self.decoder(features)
+
+ return scaled_preds
+
+
+class Refiner(nn.Module):
+ def __init__(self, in_channels=3+1):
+ super(Refiner, self).__init__()
+ self.config = Config()
+ self.epoch = 1
+ self.stem_layer = StemLayer(in_channels=in_channels, inter_channels=48, out_channels=3, norm_layer='BN' if self.config.batch_size > 1 else 'LN')
+ self.bb = build_backbone(self.config.bb)
+
+ lateral_channels_in_collection = {
+ 'vgg16': [512, 256, 128, 64], 'vgg16bn': [512, 256, 128, 64], 'resnet50': [1024, 512, 256, 64],
+ 'pvt_v2_b2': [512, 320, 128, 64], 'pvt_v2_b5': [512, 320, 128, 64],
+ 'swin_v1_b': [1024, 512, 256, 128], 'swin_v1_l': [1536, 768, 384, 192],
+ }
+ channels = lateral_channels_in_collection[self.config.bb]
+ self.squeeze_module = BasicDecBlk(channels[0], channels[0])
+
+ self.decoder = Decoder(channels)
+
+ if 0:
+ for key, value in self.named_parameters():
+ if 'bb.' in key:
+ value.requires_grad = False
+
+ def forward(self, x):
+ if isinstance(x, list):
+ x = torch.cat(x, dim=1)
+ x = self.stem_layer(x)
+ ########## Encoder ##########
+ if self.config.bb in ['vgg16', 'vgg16bn', 'resnet50']:
+ x1 = self.bb.conv1(x)
+ x2 = self.bb.conv2(x1)
+ x3 = self.bb.conv3(x2)
+ x4 = self.bb.conv4(x3)
+ else:
+ x1, x2, x3, x4 = self.bb(x)
+
+ x4 = self.squeeze_module(x4)
+
+ ########## Decoder ##########
+
+ features = [x, x1, x2, x3, x4]
+ scaled_preds = self.decoder(features)
+
+ return scaled_preds
+
+
+class Decoder(nn.Module):
+ def __init__(self, channels):
+ super(Decoder, self).__init__()
+ self.config = Config()
+ DecoderBlock = eval('BasicDecBlk')
+ LateralBlock = eval('BasicLatBlk')
+
+ self.decoder_block4 = DecoderBlock(channels[0], channels[1])
+ self.decoder_block3 = DecoderBlock(channels[1], channels[2])
+ self.decoder_block2 = DecoderBlock(channels[2], channels[3])
+ self.decoder_block1 = DecoderBlock(channels[3], channels[3]//2)
+
+ self.lateral_block4 = LateralBlock(channels[1], channels[1])
+ self.lateral_block3 = LateralBlock(channels[2], channels[2])
+ self.lateral_block2 = LateralBlock(channels[3], channels[3])
+
+ if self.config.ms_supervision:
+ self.conv_ms_spvn_4 = nn.Conv2d(channels[1], 1, 1, 1, 0)
+ self.conv_ms_spvn_3 = nn.Conv2d(channels[2], 1, 1, 1, 0)
+ self.conv_ms_spvn_2 = nn.Conv2d(channels[3], 1, 1, 1, 0)
+ self.conv_out1 = nn.Sequential(nn.Conv2d(channels[3]//2, 1, 1, 1, 0))
+
+ def forward(self, features):
+ x, x1, x2, x3, x4 = features
+ outs = []
+ p4 = self.decoder_block4(x4)
+ _p4 = F.interpolate(p4, size=x3.shape[2:], mode='bilinear', align_corners=True)
+ _p3 = _p4 + self.lateral_block4(x3)
+
+ p3 = self.decoder_block3(_p3)
+ _p3 = F.interpolate(p3, size=x2.shape[2:], mode='bilinear', align_corners=True)
+ _p2 = _p3 + self.lateral_block3(x2)
+
+ p2 = self.decoder_block2(_p2)
+ _p2 = F.interpolate(p2, size=x1.shape[2:], mode='bilinear', align_corners=True)
+ _p1 = _p2 + self.lateral_block2(x1)
+
+ _p1 = self.decoder_block1(_p1)
+ _p1 = F.interpolate(_p1, size=x.shape[2:], mode='bilinear', align_corners=True)
+ p1_out = self.conv_out1(_p1)
+
+ if self.config.ms_supervision:
+ outs.append(self.conv_ms_spvn_4(p4))
+ outs.append(self.conv_ms_spvn_3(p3))
+ outs.append(self.conv_ms_spvn_2(p2))
+ outs.append(p1_out)
+ return outs
+
+
+class RefUNet(nn.Module):
+ # Refinement
+ def __init__(self, in_channels=3+1):
+ super(RefUNet, self).__init__()
+ self.encoder_1 = nn.Sequential(
+ nn.Conv2d(in_channels, 64, 3, 1, 1),
+ nn.Conv2d(64, 64, 3, 1, 1),
+ nn.BatchNorm2d(64),
+ nn.ReLU(inplace=True)
+ )
+
+ self.encoder_2 = nn.Sequential(
+ nn.MaxPool2d(2, 2, ceil_mode=True),
+ nn.Conv2d(64, 64, 3, 1, 1),
+ nn.BatchNorm2d(64),
+ nn.ReLU(inplace=True)
+ )
+
+ self.encoder_3 = nn.Sequential(
+ nn.MaxPool2d(2, 2, ceil_mode=True),
+ nn.Conv2d(64, 64, 3, 1, 1),
+ nn.BatchNorm2d(64),
+ nn.ReLU(inplace=True)
+ )
+
+ self.encoder_4 = nn.Sequential(
+ nn.MaxPool2d(2, 2, ceil_mode=True),
+ nn.Conv2d(64, 64, 3, 1, 1),
+ nn.BatchNorm2d(64),
+ nn.ReLU(inplace=True)
+ )
+
+ self.pool4 = nn.MaxPool2d(2, 2, ceil_mode=True)
+ #####
+ self.decoder_5 = nn.Sequential(
+ nn.Conv2d(64, 64, 3, 1, 1),
+ nn.BatchNorm2d(64),
+ nn.ReLU(inplace=True)
+ )
+ #####
+ self.decoder_4 = nn.Sequential(
+ nn.Conv2d(128, 64, 3, 1, 1),
+ nn.BatchNorm2d(64),
+ nn.ReLU(inplace=True)
+ )
+
+ self.decoder_3 = nn.Sequential(
+ nn.Conv2d(128, 64, 3, 1, 1),
+ nn.BatchNorm2d(64),
+ nn.ReLU(inplace=True)
+ )
+
+ self.decoder_2 = nn.Sequential(
+ nn.Conv2d(128, 64, 3, 1, 1),
+ nn.BatchNorm2d(64),
+ nn.ReLU(inplace=True)
+ )
+
+ self.decoder_1 = nn.Sequential(
+ nn.Conv2d(128, 64, 3, 1, 1),
+ nn.BatchNorm2d(64),
+ nn.ReLU(inplace=True)
+ )
+
+ self.conv_d0 = nn.Conv2d(64, 1, 3, 1, 1)
+
+ self.upscore2 = nn.Upsample(scale_factor=2, mode='bilinear', align_corners=True)
+
+ def forward(self, x):
+ outs = []
+ if isinstance(x, list):
+ x = torch.cat(x, dim=1)
+ hx = x
+
+ hx1 = self.encoder_1(hx)
+ hx2 = self.encoder_2(hx1)
+ hx3 = self.encoder_3(hx2)
+ hx4 = self.encoder_4(hx3)
+
+ hx = self.decoder_5(self.pool4(hx4))
+ hx = torch.cat((self.upscore2(hx), hx4), 1)
+
+ d4 = self.decoder_4(hx)
+ hx = torch.cat((self.upscore2(d4), hx3), 1)
+
+ d3 = self.decoder_3(hx)
+ hx = torch.cat((self.upscore2(d3), hx2), 1)
+
+ d2 = self.decoder_2(hx)
+ hx = torch.cat((self.upscore2(d2), hx1), 1)
+
+ d1 = self.decoder_1(hx)
+
+ x = self.conv_d0(d1)
+ outs.append(x)
+ return outs
diff --git a/py/BiRefNet_v2/models/refinement/stem_layer.py b/py/BiRefNet_v2/models/refinement/stem_layer.py
new file mode 100644
index 0000000..8dd0a0d
--- /dev/null
+++ b/py/BiRefNet_v2/models/refinement/stem_layer.py
@@ -0,0 +1,45 @@
+import torch.nn as nn
+from ...models.modules.utils import build_act_layer, build_norm_layer
+
+
+class StemLayer(nn.Module):
+ r""" Stem layer of InternImage
+ Args:
+ in_channels (int): number of input channels
+ out_channels (int): number of output channels
+ act_layer (str): activation layer
+ norm_layer (str): normalization layer
+ """
+
+ def __init__(self,
+ in_channels=3+1,
+ inter_channels=48,
+ out_channels=96,
+ act_layer='GELU',
+ norm_layer='BN'):
+ super().__init__()
+ self.conv1 = nn.Conv2d(in_channels,
+ inter_channels,
+ kernel_size=3,
+ stride=1,
+ padding=1)
+ self.norm1 = build_norm_layer(
+ inter_channels, norm_layer, 'channels_first', 'channels_first'
+ )
+ self.act = build_act_layer(act_layer)
+ self.conv2 = nn.Conv2d(inter_channels,
+ out_channels,
+ kernel_size=3,
+ stride=1,
+ padding=1)
+ self.norm2 = build_norm_layer(
+ out_channels, norm_layer, 'channels_first', 'channels_first'
+ )
+
+ def forward(self, x):
+ x = self.conv1(x)
+ x = self.norm1(x)
+ x = self.act(x)
+ x = self.conv2(x)
+ x = self.norm2(x)
+ return x
diff --git a/py/BiRefNet_v2/requirements.txt b/py/BiRefNet_v2/requirements.txt
new file mode 100644
index 0000000..546ffa3
--- /dev/null
+++ b/py/BiRefNet_v2/requirements.txt
@@ -0,0 +1,15 @@
+--extra-index-url https://download.pytorch.org/whl/cu118
+torch==2.0.1
+--extra-index-url https://download.pytorch.org/whl/cu118
+torchvision==0.15.2
+numpy<2
+opencv-python
+timm
+scipy
+scikit-image
+kornia
+
+tqdm
+prettytable
+
+huggingface_hub
diff --git a/py/BiRefNet_v2/rm_cache.sh b/py/BiRefNet_v2/rm_cache.sh
new file mode 100644
index 0000000..5e75b92
--- /dev/null
+++ b/py/BiRefNet_v2/rm_cache.sh
@@ -0,0 +1,20 @@
+#!/bin/bash
+rm -rf __pycache__ */__pycache__
+
+# Val
+rm -r tmp*
+
+# Train
+rm slurm*
+rm -r ckpt
+rm nohup.out*
+
+# Eval
+rm -r evaluation/eval-*
+rm -r tmp*
+rm -r e_logs/
+
+# System
+rm core-*-python-*
+
+clear
diff --git a/py/BiRefNet_v2/sub.sh b/py/BiRefNet_v2/sub.sh
new file mode 100644
index 0000000..9e216b9
--- /dev/null
+++ b/py/BiRefNet_v2/sub.sh
@@ -0,0 +1,17 @@
+#!/bin/sh
+# Example: ./sub.sh tmp_proj 0,1,2,3 3 --> Use 0,1,2,3 for training, release GPUs, use GPU:3 for inference.
+
+# module load gcc/11.2.0 cuda/11.8 cudnn/8.6.0_cu11x && cpu_core_num=6
+module load compilers/cuda/11.8 compilers/gcc/12.2.0 cudnn/8.4.0.27_cuda11.x && cpu_core_num=32
+
+export PYTHONUNBUFFERED=1
+
+method=${1:-"BSL"}
+devices=${2:-0}
+gpu_num=$(($(echo ${devices%%,} | grep -o "," | wc -l)+1))
+
+sbatch --nodes=1 -p vip_gpu_ailab -A ai4bio \
+ --gres=gpu:${gpu_num} --ntasks-per-node=1 --cpus-per-task=$((gpu_num*cpu_core_num)) \
+ ./train_test.sh ${method} ${devices}
+
+hostname
diff --git a/py/BiRefNet_v2/test.sh b/py/BiRefNet_v2/test.sh
new file mode 100644
index 0000000..66a6149
--- /dev/null
+++ b/py/BiRefNet_v2/test.sh
@@ -0,0 +1,29 @@
+devices=${1:-0}
+pred_root=${2:-e_preds}
+
+# Inference
+
+CUDA_VISIBLE_DEVICES=${devices} python inference.py --pred_root ${pred_root}
+
+echo Inference finished at $(date)
+
+# Evaluation
+log_dir=e_logs && mkdir ${log_dir}
+
+task=$(python3 config.py)
+case "${task}" in
+ "DIS5K") testsets='DIS-VD,DIS-TE1,DIS-TE2,DIS-TE3,DIS-TE4' ;;
+ "COD") testsets='CHAMELEON,NC4K,TE-CAMO,TE-COD10K' ;;
+ "HRSOD") testsets='DAVIS-S,TE-HRSOD,TE-UHRSD,DUT-OMRON,TE-DUTS' ;;
+ "General") testsets='DIS-VD' ;;
+ "Matting") testsets='TE-P3M-500-P' ;;
+esac
+testsets=(`echo ${testsets} | tr ',' ' '`) && testsets=${testsets[@]}
+
+for testset in ${testsets}; do
+ python eval_existingOnes.py --pred_root ${pred_root} --data_lst ${testset} > ${log_dir}/eval_${testset}.out
+ # nohup python eval_existingOnes.py --pred_root ${pred_root} --data_lst ${testset} > ${log_dir}/eval_${testset}.out 2>&1 &
+done
+
+
+echo Evaluation started at $(date)
diff --git a/py/BiRefNet_v2/train.py b/py/BiRefNet_v2/train.py
new file mode 100644
index 0000000..8b47b54
--- /dev/null
+++ b/py/BiRefNet_v2/train.py
@@ -0,0 +1,333 @@
+import os
+import datetime
+import argparse
+import torch
+import torch.nn as nn
+import torch.optim as optim
+from torch.autograd import Variable
+
+from .config import Config
+from .loss import PixLoss, ClsLoss
+from .dataset import MyData
+from .models.birefnet import BiRefNet
+from .utils import Logger, AverageMeter, set_seed, check_state_dict
+
+from torch.utils.data.distributed import DistributedSampler
+from torch.nn.parallel import DistributedDataParallel as DDP
+from torch.distributed import init_process_group, destroy_process_group, get_rank
+from torch.cuda import amp
+
+
+parser = argparse.ArgumentParser(description='')
+parser.add_argument('--resume', default=None, type=str, help='path to latest checkpoint')
+parser.add_argument('--epochs', default=120, type=int)
+parser.add_argument('--trainset', default='DIS5K', type=str, help="Options: 'DIS5K'")
+parser.add_argument('--ckpt_dir', default=None, help='Temporary folder')
+parser.add_argument('--testsets', default='DIS-VD+DIS-TE1+DIS-TE2+DIS-TE3+DIS-TE4', type=str)
+parser.add_argument('--dist', default=False, type=lambda x: x == 'True')
+args = parser.parse_args()
+
+
+config = Config()
+if config.rand_seed:
+ set_seed(config.rand_seed)
+
+if config.use_fp16:
+ # Half Precision
+ scaler = amp.GradScaler(enabled=config.use_fp16)
+
+# DDP
+to_be_distributed = args.dist
+if to_be_distributed:
+ init_process_group(backend="nccl", timeout=datetime.timedelta(seconds=3600*10))
+ device = int(os.environ["LOCAL_RANK"])
+else:
+ device = config.device
+
+epoch_st = 1
+# make dir for ckpt
+os.makedirs(args.ckpt_dir, exist_ok=True)
+
+# Init log file
+logger = Logger(os.path.join(args.ckpt_dir, "log.txt"))
+logger_loss_idx = 1
+
+# log model and optimizer params
+# logger.info("Model details:"); logger.info(model)
+logger.info("datasets: load_all={}, compile={}.".format(config.load_all, config.compile))
+logger.info("Other hyperparameters:"); logger.info(args)
+print('batch size:', config.batch_size)
+
+
+if os.path.exists(os.path.join(config.data_root_dir, config.task, args.testsets.strip('+').split('+')[0])):
+ args.testsets = args.testsets.strip('+').split('+')
+else:
+ args.testsets = []
+
+# Init model
+def prepare_dataloader(dataset: torch.utils.data.Dataset, batch_size: int, to_be_distributed=False, is_train=True):
+ if to_be_distributed:
+ return torch.utils.data.DataLoader(
+ dataset=dataset, batch_size=batch_size, num_workers=min(config.num_workers, batch_size), pin_memory=True,
+ shuffle=False, sampler=DistributedSampler(dataset), drop_last=True
+ )
+ else:
+ return torch.utils.data.DataLoader(
+ dataset=dataset, batch_size=batch_size, num_workers=min(config.num_workers, batch_size, 0), pin_memory=True,
+ shuffle=is_train, drop_last=True
+ )
+
+
+def init_data_loaders(to_be_distributed):
+ # Prepare dataset
+ train_loader = prepare_dataloader(
+ MyData(datasets=config.training_set, image_size=config.size, is_train=True),
+ config.batch_size, to_be_distributed=to_be_distributed, is_train=True
+ )
+ print(len(train_loader), "batches of train dataloader {} have been created.".format(config.training_set))
+ test_loaders = {}
+ for testset in args.testsets:
+ _data_loader_test = prepare_dataloader(
+ MyData(datasets=testset, image_size=config.size, is_train=False),
+ config.batch_size_valid, is_train=False
+ )
+ print(len(_data_loader_test), "batches of valid dataloader {} have been created.".format(testset))
+ test_loaders[testset] = _data_loader_test
+ return train_loader, test_loaders
+
+
+def init_models_optimizers(epochs, to_be_distributed):
+ model = BiRefNet(bb_pretrained=True)
+ if args.resume:
+ if os.path.isfile(args.resume):
+ logger.info("=> loading checkpoint '{}'".format(args.resume))
+ state_dict = torch.load(args.resume, map_location='cpu')
+ state_dict = check_state_dict(state_dict)
+ model.load_state_dict(state_dict)
+ global epoch_st
+ epoch_st = int(args.resume.rstrip('.pth').split('epoch_')[-1]) + 1
+ else:
+ logger.info("=> no checkpoint found at '{}'".format(args.resume))
+ if to_be_distributed:
+ model = model.to(device)
+ model = DDP(model, device_ids=[device])
+ else:
+ model = model.to(device)
+ if config.compile:
+ model = torch.compile(model, mode=['default', 'reduce-overhead', 'max-autotune'][0])
+ if config.precisionHigh:
+ torch.set_float32_matmul_precision('high')
+
+
+ # Setting optimizer
+ if config.optimizer == 'AdamW':
+ optimizer = optim.AdamW(params=model.parameters(), lr=config.lr, weight_decay=1e-2)
+ elif config.optimizer == 'Adam':
+ optimizer = optim.Adam(params=model.parameters(), lr=config.lr, weight_decay=0)
+ lr_scheduler = torch.optim.lr_scheduler.MultiStepLR(
+ optimizer,
+ milestones=[lde if lde > 0 else epochs + lde + 1 for lde in config.lr_decay_epochs],
+ gamma=config.lr_decay_rate
+ )
+ logger.info("Optimizer details:"); logger.info(optimizer)
+ logger.info("Scheduler details:"); logger.info(lr_scheduler)
+
+ return model, optimizer, lr_scheduler
+
+
+class Trainer:
+ def __init__(
+ self, data_loaders, model_opt_lrsch,
+ ):
+ self.model, self.optimizer, self.lr_scheduler = model_opt_lrsch
+ self.train_loader, self.test_loaders = data_loaders
+ if config.out_ref:
+ self.criterion_gdt = nn.BCELoss() if not config.use_fp16 else nn.BCEWithLogitsLoss()
+
+ # Setting Losses
+ self.pix_loss = PixLoss()
+ self.cls_loss = ClsLoss()
+
+ # Others
+ self.loss_log = AverageMeter()
+ if config.lambda_adv_g:
+ self.optimizer_d, self.lr_scheduler_d, self.disc, self.adv_criterion = self._load_adv_components()
+ self.disc_update_for_odd = 0
+
+ def _load_adv_components(self):
+ # AIL
+ from loss import Discriminator
+ disc = Discriminator(channels=3, img_size=config.size)
+ if to_be_distributed:
+ disc = disc.to(device)
+ disc = DDP(disc, device_ids=[device], broadcast_buffers=False)
+ else:
+ disc = disc.to(device)
+ if config.compile:
+ disc = torch.compile(disc, mode=['default', 'reduce-overhead', 'max-autotune'][0])
+ adv_criterion = nn.BCELoss() if not config.use_fp16 else nn.BCEWithLogitsLoss()
+ if config.optimizer == 'AdamW':
+ optimizer_d = optim.AdamW(params=disc.parameters(), lr=config.lr, weight_decay=1e-2)
+ elif config.optimizer == 'Adam':
+ optimizer_d = optim.Adam(params=disc.parameters(), lr=config.lr, weight_decay=0)
+ lr_scheduler_d = torch.optim.lr_scheduler.MultiStepLR(
+ optimizer_d,
+ milestones=[lde if lde > 0 else args.epochs + lde + 1 for lde in config.lr_decay_epochs],
+ gamma=config.lr_decay_rate
+ )
+ return optimizer_d, lr_scheduler_d, disc, adv_criterion
+
+ def _train_batch(self, batch):
+ inputs = batch[0].to(device)
+ gts = batch[1].to(device)
+ class_labels = batch[2].to(device)
+ if config.use_fp16:
+ with amp.autocast(enabled=config.use_fp16):
+ scaled_preds, class_preds_lst = self.model(inputs)
+ if config.out_ref:
+ (outs_gdt_pred, outs_gdt_label), scaled_preds = scaled_preds
+ for _idx, (_gdt_pred, _gdt_label) in enumerate(zip(outs_gdt_pred, outs_gdt_label)):
+ _gdt_pred = nn.functional.interpolate(_gdt_pred, size=_gdt_label.shape[2:], mode='bilinear', align_corners=True)#.sigmoid()
+ # _gdt_label = _gdt_label.sigmoid()
+ loss_gdt = self.criterion_gdt(_gdt_pred, _gdt_label) if _idx == 0 else self.criterion_gdt(_gdt_pred, _gdt_label) + loss_gdt
+ # self.loss_dict['loss_gdt'] = loss_gdt.item()
+ if None in class_preds_lst:
+ loss_cls = 0.
+ else:
+ loss_cls = self.cls_loss(class_preds_lst, class_labels) * 1.0
+ self.loss_dict['loss_cls'] = loss_cls.item()
+
+ # Loss
+ loss_pix = self.pix_loss(scaled_preds, torch.clamp(gts, 0, 1)) * 1.0
+ self.loss_dict['loss_pix'] = loss_pix.item()
+ # since there may be several losses for sal, the lambdas for them (lambdas_pix) are inside the loss.py
+ loss = loss_pix + loss_cls
+ if config.out_ref:
+ loss = loss + loss_gdt * 1.0
+
+ if config.lambda_adv_g:
+ # gen
+ valid = Variable(torch.cuda.FloatTensor(scaled_preds[-1].shape[0], 1).fill_(1.0), requires_grad=False).to(device)
+ adv_loss_g = self.adv_criterion(self.disc(scaled_preds[-1] * inputs), valid) * config.lambda_adv_g
+ loss += adv_loss_g
+ self.loss_dict['loss_adv'] = adv_loss_g.item()
+ self.disc_update_for_odd += 1
+ # self.loss_log.update(loss.item(), inputs.size(0))
+ # self.optimizer.zero_grad()
+ # loss.backward()
+ # self.optimizer.step()
+ self.optimizer.zero_grad()
+ scaler.scale(loss).backward()
+ scaler.step(self.optimizer)
+ scaler.update()
+
+ if config.lambda_adv_g and self.disc_update_for_odd % 2 == 0:
+ # disc
+ fake = Variable(torch.cuda.FloatTensor(scaled_preds[-1].shape[0], 1).fill_(0.0), requires_grad=False).to(device)
+ adv_loss_real = self.adv_criterion(self.disc(gts * inputs), valid)
+ adv_loss_fake = self.adv_criterion(self.disc(scaled_preds[-1].detach() * inputs.detach()), fake)
+ adv_loss_d = (adv_loss_real + adv_loss_fake) / 2 * config.lambda_adv_d
+ self.loss_dict['loss_adv_d'] = adv_loss_d.item()
+ # self.optimizer_d.zero_grad()
+ # adv_loss_d.backward()
+ # self.optimizer_d.step()
+ self.optimizer_d.zero_grad()
+ scaler.scale(adv_loss_d).backward()
+ scaler.step(self.optimizer_d)
+ scaler.update()
+ else:
+ scaled_preds, class_preds_lst = self.model(inputs)
+ if config.out_ref:
+ (outs_gdt_pred, outs_gdt_label), scaled_preds = scaled_preds
+ for _idx, (_gdt_pred, _gdt_label) in enumerate(zip(outs_gdt_pred, outs_gdt_label)):
+ _gdt_pred = nn.functional.interpolate(_gdt_pred, size=_gdt_label.shape[2:], mode='bilinear', align_corners=True).sigmoid()
+ _gdt_label = _gdt_label.sigmoid()
+ loss_gdt = self.criterion_gdt(_gdt_pred, _gdt_label) if _idx == 0 else self.criterion_gdt(_gdt_pred, _gdt_label) + loss_gdt
+ # self.loss_dict['loss_gdt'] = loss_gdt.item()
+ if None in class_preds_lst:
+ loss_cls = 0.
+ else:
+ loss_cls = self.cls_loss(class_preds_lst, class_labels) * 1.0
+ self.loss_dict['loss_cls'] = loss_cls.item()
+
+ # Loss
+ loss_pix = self.pix_loss(scaled_preds, torch.clamp(gts, 0, 1)) * 1.0
+ self.loss_dict['loss_pix'] = loss_pix.item()
+ # since there may be several losses for sal, the lambdas for them (lambdas_pix) are inside the loss.py
+ loss = loss_pix + loss_cls
+ if config.out_ref:
+ loss = loss + loss_gdt * 1.0
+
+ if config.lambda_adv_g:
+ # gen
+ valid = Variable(torch.cuda.FloatTensor(scaled_preds[-1].shape[0], 1).fill_(1.0), requires_grad=False).to(device)
+ adv_loss_g = self.adv_criterion(self.disc(scaled_preds[-1] * inputs), valid) * config.lambda_adv_g
+ loss += adv_loss_g
+ self.loss_dict['loss_adv'] = adv_loss_g.item()
+ self.disc_update_for_odd += 1
+ self.loss_log.update(loss.item(), inputs.size(0))
+ self.optimizer.zero_grad()
+ loss.backward()
+ self.optimizer.step()
+
+ if config.lambda_adv_g and self.disc_update_for_odd % 2 == 0:
+ # disc
+ fake = Variable(torch.cuda.FloatTensor(scaled_preds[-1].shape[0], 1).fill_(0.0), requires_grad=False).to(device)
+ adv_loss_real = self.adv_criterion(self.disc(gts * inputs), valid)
+ adv_loss_fake = self.adv_criterion(self.disc(scaled_preds[-1].detach() * inputs.detach()), fake)
+ adv_loss_d = (adv_loss_real + adv_loss_fake) / 2 * config.lambda_adv_d
+ self.loss_dict['loss_adv_d'] = adv_loss_d.item()
+ self.optimizer_d.zero_grad()
+ adv_loss_d.backward()
+ self.optimizer_d.step()
+
+ def train_epoch(self, epoch):
+ global logger_loss_idx
+ self.model.train()
+ self.loss_dict = {}
+ if epoch > args.epochs + config.finetune_last_epochs[1]:
+ for k in self.pix_loss.lambdas_pix_last.keys():
+ if k.lower() == config.finetune_last_epochs[0].lower():
+ self.pix_loss.lambdas_pix_last[k] = config.lambdas_pix_last[k] * 0.5
+ else:
+ self.pix_loss.lambdas_pix_last[k] = 0
+
+ for batch_idx, batch in enumerate(self.train_loader):
+ self._train_batch(batch)
+ # Logger
+ if batch_idx % 20 == 0:
+ info_progress = 'Epoch[{0}/{1}] Iter[{2}/{3}].'.format(epoch, args.epochs, batch_idx, len(self.train_loader))
+ info_loss = 'Training Losses'
+ for loss_name, loss_value in self.loss_dict.items():
+ info_loss += ', {}: {:.3f}'.format(loss_name, loss_value)
+ logger.info(' '.join((info_progress, info_loss)))
+ info_loss = '@==Final== Epoch[{0}/{1}] Training Loss: {loss.avg:.3f} '.format(epoch, args.epochs, loss=self.loss_log)
+ logger.info(info_loss)
+
+ self.lr_scheduler.step()
+ if config.lambda_adv_g:
+ self.lr_scheduler_d.step()
+ return self.loss_log.avg
+
+
+def main():
+
+ trainer = Trainer(
+ data_loaders=init_data_loaders(to_be_distributed),
+ model_opt_lrsch=init_models_optimizers(args.epochs, to_be_distributed)
+ )
+
+ for epoch in range(epoch_st, args.epochs+1):
+ train_loss = trainer.train_epoch(epoch)
+ # Save checkpoint
+ # DDP
+ if epoch >= args.epochs - config.save_last and epoch % config.save_step == 0:
+ torch.save(
+ trainer.model.module.state_dict() if to_be_distributed else trainer.model.state_dict(),
+ os.path.join(args.ckpt_dir, 'epoch_{}.pth'.format(epoch))
+ )
+ if to_be_distributed:
+ destroy_process_group()
+
+if __name__ == '__main__':
+ main()
diff --git a/py/BiRefNet_v2/train.sh b/py/BiRefNet_v2/train.sh
new file mode 100644
index 0000000..78421d8
--- /dev/null
+++ b/py/BiRefNet_v2/train.sh
@@ -0,0 +1,42 @@
+#!/bin/bash
+# Run script
+# Settings of training & test for different tasks.
+method="$1"
+task=$(python3 config.py)
+case "${task}" in
+ "DIS5K") epochs=600 && val_last=50 && step=5 ;;
+ "COD") epochs=150 && val_last=50 && step=5 ;;
+ "HRSOD") epochs=150 && val_last=50 && step=5 ;;
+ "General") epochs=250 && val_last=20 && step=2 ;;
+ "Matting") epochs=100 && val_last=20 && step=2 ;;
+esac
+testsets=NO # Non-existing folder to skip.
+# testsets=TE-COD10K # for COD
+
+# Train
+devices=$2
+nproc_per_node=$(echo ${devices%%,} | grep -o "," | wc -l)
+
+to_be_distributed=`echo ${nproc_per_node} | awk '{if($e > 0) print "True"; else print "False";}'`
+
+echo Training started at $(date)
+if [ ${to_be_distributed} == "True" ]
+then
+ # Adapt the nproc_per_node by the number of GPUs. Give 8989 as the default value of master_port.
+ echo "Multi-GPU mode received..."
+ CUDA_VISIBLE_DEVICES=${devices} \
+ torchrun --nproc_per_node $((nproc_per_node+1)) --master_port=${3:-8999} \
+ train.py --ckpt_dir ckpt/${method} --epochs ${epochs} \
+ --testsets ${testsets} \
+ --dist ${to_be_distributed} \
+ --resume xx/xx-epoch_244.pth
+else
+ echo "Single-GPU mode received..."
+ CUDA_VISIBLE_DEVICES=${devices} \
+ python train.py --ckpt_dir ckpt/${method} --epochs ${epochs} \
+ --testsets ${testsets} \
+ --dist ${to_be_distributed} \
+ --resume xx/xx-epoch_244.pth
+fi
+
+echo Training finished at $(date)
diff --git a/py/BiRefNet_v2/train_test.sh b/py/BiRefNet_v2/train_test.sh
new file mode 100644
index 0000000..e9d3a26
--- /dev/null
+++ b/py/BiRefNet_v2/train_test.sh
@@ -0,0 +1,11 @@
+#!/bin/sh
+
+method=${1:-"BSL"}
+devices=${2:-"0,1,2,3,4,5,6,7"}
+
+bash train.sh ${method} ${devices}
+
+devices_test=${3:-0}
+bash test.sh ${devices_test}
+
+hostname
diff --git a/py/BiRefNet_v2/tutorials/BiRefNet_inference.ipynb b/py/BiRefNet_v2/tutorials/BiRefNet_inference.ipynb
new file mode 100644
index 0000000..4173711
--- /dev/null
+++ b/py/BiRefNet_v2/tutorials/BiRefNet_inference.ipynb
@@ -0,0 +1,1575 @@
+{
+ "cells": [
+ {
+ "cell_type": "markdown",
+ "metadata": {},
+ "source": [
+ "### Online Colab Demo: https://colab.research.google.com/drive/14Dqg7oeBkFEtchaHLNpig2BcdkZEogba\n",
+ "### Hugging Face Spaces Demo: https://huggingface.co/spaces/ZhengPeng7/BiRefNet_demo"
+ ]
+ },
+ {
+ "cell_type": "code",
+ "execution_count": null,
+ "metadata": {
+ "colab": {
+ "base_uri": "https://localhost:8080/",
+ "height": 391,
+ "referenced_widgets": [
+ "7d19deaab4c845eea4705567bdc65d60",
+ "a8941bdff0984189be91fab5bfe1c52c",
+ "b6be81c6cc1e4608a88c785c448bfaa4",
+ "3b0ddb32ffa442aab3b02a22432bb233",
+ "688a14cd34704e4ea2e261f6619449ee",
+ "8a54f72e65a24d7a93852aca8ecae0a2",
+ "7745af46aa694f5b8ff0ec8fbd72025f",
+ "4e5fe4296291455f88ae0e5f257c398e",
+ "60d3c15d4b944546992d3517bb73912a",
+ "a124da94c5b143a29d9420eba3954859",
+ "96485cf384484a709e4c32b74e0223af",
+ "7bef3ce58df040c986a74eb5334f42c9",
+ "1848689bf6a14235b508647d1bdb7f69",
+ "055ec5d0d4da4d1b8df33ddb22d6b55b",
+ "3c9b1a3ed2f64ed58a29083280098292",
+ "b2b95dd8b75d4625ad1bd02a461f304d",
+ "044920317fb64cd088d49f334930b886",
+ "a8416e76428b46c092c7afbf7129b5f9",
+ "ebf0b02bc7734ebea9a846233a1a6ec4",
+ "fb3e80729a214bc5993b62ee04d7c58d",
+ "cf459ec049624ba29f54cceeb1469785",
+ "db64a13268ab452db1241d06f16d3d6a",
+ "80a53775d54e47b2922031cc4cd00548",
+ "e13013a9d20843bca6da8fe9f0fdb644",
+ "6f45e75849f749a1b5fad6e8f7879c8f",
+ "0ef4b41f8d0141859b7767d112596198",
+ "fecc8e75d8d643759123eaf2ff30fa2a",
+ "56eb4e5291404f3289196d526fef524f",
+ "7fb491d5b8f34a2e842518c1e4ea4906",
+ "e02fc794fb3645bda6b802277a3e5c1a",
+ "f9a719112f53400782f97d0683862a5f",
+ "4480d4d37c2648e7bb5635109b5d4715",
+ "67ce0f56b38941749bdb580a830598c2",
+ "6667788f1fa44603a4a04ba8aa5e1cc9",
+ "4d162f2efc364cf793709f0b9c9dc888",
+ "e078f8d23e8449e7a9b5771e342febeb",
+ "a330873af8cf45eca95780c50260dbb9",
+ "61746c0aaf96444391d1a86ea7223cb0",
+ "2ad78a6887eb4a369bac75e2436758b9",
+ "78d3ab060eb64941bcba51fa3b32f493",
+ "1f1a0c0dc2b74d56b68a779ef944416d",
+ "1546b1d7383c4a25b2ab3782ff204cad",
+ "6588cde28032444980c4e314d6b9a648",
+ "3e8d434d8e524c8c9f3fe3359df4b327"
+ ]
+ },
+ "id": "7lFgKfPS8Icy",
+ "outputId": "2f00b063-86bf-4ba8-fa5e-38d2f5a66462"
+ },
+ "outputs": [],
+ "source": [
+ "# Imports\n",
+ "from PIL import Image\n",
+ "import torch\n",
+ "from torchvision import transforms\n",
+ "from IPython.display import display\n",
+ "\n",
+ "import sys\n",
+ "sys.path.insert(0, \"../\")\n",
+ "from models.birefnet import BiRefNet\n",
+ "\n",
+ "\n",
+ "# Load Model\n",
+ "# Option 2 and Option 3 is better for local running -- we can modify codes locally.\n",
+ "\n",
+ "# # # Option 1: loading BiRefNet with weights:\n",
+ "# from transformers import AutoModelForImageSegmentation\n",
+ "# birefnet = AutoModelForImageSegmentation.from_pretrained('zhengpeng7/BiRefNet', trust_remote_code=True)\n",
+ "\n",
+ "# Option-2: loading weights with BiReNet codes:\n",
+ "birefnet = BiRefNet.from_pretrained(\n",
+ " [\n",
+ " 'zhengpeng7/BiRefNet',\n",
+ " 'zhengpeng7/BiRefNet-portrait',\n",
+ " 'zhengpeng7/BiRefNet-legacy', 'zhengpeng7/BiRefNet-DIS5K-TR_TEs', 'zhengpeng7/BiRefNet-DIS5K', 'zhengpeng7/BiRefNet-HRSOD', 'zhengpeng7/BiRefNet-COD',\n",
+ " 'zhengpeng7/BiRefNet_lite', # Modify the `bb` in `config.py` to `swin_v1_tiny`.\n",
+ " ][0]\n",
+ ")\n",
+ "\n",
+ "# # Option-3: Loading model and weights from local disk:\n",
+ "# from utils import check_state_dict\n",
+ "\n",
+ "# birefnet = BiRefNet(bb_pretrained=False)\n",
+ "# state_dict = torch.load('../BiRefNet-general-epoch_244.pth', map_location='cpu')\n",
+ "# state_dict = check_state_dict(state_dict)\n",
+ "# birefnet.load_state_dict(state_dict)\n",
+ "\n",
+ "device = 'cuda' if torch.cuda.is_available() else 'cpu'\n",
+ "\n",
+ "torch.set_float32_matmul_precision(['high', 'highest'][0])\n",
+ "\n",
+ "birefnet.to(device)\n",
+ "birefnet.eval()\n",
+ "print('BiRefNet is ready to use.')\n",
+ "\n",
+ "# Input Data\n",
+ "transform_image = transforms.Compose([\n",
+ " transforms.Resize((1024, 1024)),\n",
+ " transforms.ToTensor(),\n",
+ " transforms.Normalize([0.485, 0.456, 0.406], [0.229, 0.224, 0.225])\n",
+ "])"
+ ]
+ },
+ {
+ "cell_type": "code",
+ "execution_count": null,
+ "metadata": {
+ "colab": {
+ "base_uri": "https://localhost:8080/",
+ "height": 1000
+ },
+ "id": "PECYekO53hrR",
+ "outputId": "73f47406-9d92-48b1-fe74-abbb5b83c7a8"
+ },
+ "outputs": [],
+ "source": [
+ "import os\n",
+ "from glob import glob\n",
+ "from image_proc import refine_foreground\n",
+ "\n",
+ "src_dir = '../images_todo'\n",
+ "image_paths = glob(os.path.join(src_dir, '*'))\n",
+ "dst_dir = '../predictions'\n",
+ "os.makedirs(dst_dir, exist_ok=True)\n",
+ "for image_path in image_paths:\n",
+ " print('Processing {} ...'.format(image_path))\n",
+ " image = Image.open(image_path)\n",
+ " input_images = transform_image(image).unsqueeze(0).to(device)\n",
+ "\n",
+ " # Prediction\n",
+ " with torch.no_grad():\n",
+ " preds = birefnet(input_images)[-1].sigmoid().cpu()\n",
+ " pred = preds[0].squeeze()\n",
+ "\n",
+ " # Show Results\n",
+ " pred_pil = transforms.ToPILImage()(pred)\n",
+ " pred_pil.resize(image.size).save(image_path.replace(src_dir, dst_dir))\n",
+ "\n",
+ " # Visualize the last sample:\n",
+ " # Scale proportionally with max length to 1024 for faster showing\n",
+ " scale_ratio = 1024 / max(image.size)\n",
+ " scaled_size = (int(image.size[0] * scale_ratio), int(image.size[1] * scale_ratio))\n",
+ "\n",
+ " image_masked = refine_foreground(image, pred_pil)\n",
+ " image_masked.putalpha(pred_pil.resize(image.size))\n",
+ "\n",
+ "display(image.resize(scaled_size))\n",
+ "display(pred_pil.resize(scaled_size))\n",
+ "display(image_masked.resize(scaled_size))"
+ ]
+ },
+ {
+ "cell_type": "code",
+ "execution_count": null,
+ "metadata": {},
+ "outputs": [],
+ "source": []
+ }
+ ],
+ "metadata": {
+ "accelerator": "GPU",
+ "colab": {
+ "gpuType": "T4",
+ "provenance": []
+ },
+ "kernelspec": {
+ "display_name": "Python 3",
+ "name": "python3"
+ },
+ "language_info": {
+ "codemirror_mode": {
+ "name": "ipython",
+ "version": 3
+ },
+ "file_extension": ".py",
+ "mimetype": "text/x-python",
+ "name": "python",
+ "nbconvert_exporter": "python",
+ "pygments_lexer": "ipython3",
+ "version": "3.9.19"
+ },
+ "widgets": {
+ "application/vnd.jupyter.widget-state+json": {
+ "044920317fb64cd088d49f334930b886": {
+ "model_module": "@jupyter-widgets/base",
+ "model_module_version": "1.2.0",
+ "model_name": "LayoutModel",
+ "state": {
+ "_model_module": "@jupyter-widgets/base",
+ "_model_module_version": "1.2.0",
+ "_model_name": "LayoutModel",
+ "_view_count": null,
+ "_view_module": "@jupyter-widgets/base",
+ "_view_module_version": "1.2.0",
+ "_view_name": "LayoutView",
+ "align_content": null,
+ "align_items": null,
+ "align_self": null,
+ "border": null,
+ "bottom": null,
+ "display": null,
+ "flex": null,
+ "flex_flow": null,
+ "grid_area": null,
+ "grid_auto_columns": null,
+ "grid_auto_flow": null,
+ "grid_auto_rows": null,
+ "grid_column": null,
+ "grid_gap": null,
+ "grid_row": null,
+ "grid_template_areas": null,
+ "grid_template_columns": null,
+ "grid_template_rows": null,
+ "height": null,
+ "justify_content": null,
+ "justify_items": null,
+ "left": null,
+ "margin": null,
+ "max_height": null,
+ "max_width": null,
+ "min_height": null,
+ "min_width": null,
+ "object_fit": null,
+ "object_position": null,
+ "order": null,
+ "overflow": null,
+ "overflow_x": null,
+ "overflow_y": null,
+ "padding": null,
+ "right": null,
+ "top": null,
+ "visibility": null,
+ "width": null
+ }
+ },
+ "055ec5d0d4da4d1b8df33ddb22d6b55b": {
+ "model_module": "@jupyter-widgets/controls",
+ "model_module_version": "1.5.0",
+ "model_name": "FloatProgressModel",
+ "state": {
+ "_dom_classes": [],
+ "_model_module": "@jupyter-widgets/controls",
+ "_model_module_version": "1.5.0",
+ "_model_name": "FloatProgressModel",
+ "_view_count": null,
+ "_view_module": "@jupyter-widgets/controls",
+ "_view_module_version": "1.5.0",
+ "_view_name": "ProgressView",
+ "bar_style": "success",
+ "description": "",
+ "description_tooltip": null,
+ "layout": "IPY_MODEL_ebf0b02bc7734ebea9a846233a1a6ec4",
+ "max": 298,
+ "min": 0,
+ "orientation": "horizontal",
+ "style": "IPY_MODEL_fb3e80729a214bc5993b62ee04d7c58d",
+ "value": 298
+ }
+ },
+ "0ef4b41f8d0141859b7767d112596198": {
+ "model_module": "@jupyter-widgets/controls",
+ "model_module_version": "1.5.0",
+ "model_name": "HTMLModel",
+ "state": {
+ "_dom_classes": [],
+ "_model_module": "@jupyter-widgets/controls",
+ "_model_module_version": "1.5.0",
+ "_model_name": "HTMLModel",
+ "_view_count": null,
+ "_view_module": "@jupyter-widgets/controls",
+ "_view_module_version": "1.5.0",
+ "_view_name": "HTMLView",
+ "description": "",
+ "description_tooltip": null,
+ "layout": "IPY_MODEL_4480d4d37c2648e7bb5635109b5d4715",
+ "placeholder": "",
+ "style": "IPY_MODEL_67ce0f56b38941749bdb580a830598c2",
+ "value": " 91.3k/91.3k [00:00<00:00, 1.97MB/s]"
+ }
+ },
+ "1546b1d7383c4a25b2ab3782ff204cad": {
+ "model_module": "@jupyter-widgets/controls",
+ "model_module_version": "1.5.0",
+ "model_name": "ProgressStyleModel",
+ "state": {
+ "_model_module": "@jupyter-widgets/controls",
+ "_model_module_version": "1.5.0",
+ "_model_name": "ProgressStyleModel",
+ "_view_count": null,
+ "_view_module": "@jupyter-widgets/base",
+ "_view_module_version": "1.2.0",
+ "_view_name": "StyleView",
+ "bar_color": null,
+ "description_width": ""
+ }
+ },
+ "1848689bf6a14235b508647d1bdb7f69": {
+ "model_module": "@jupyter-widgets/controls",
+ "model_module_version": "1.5.0",
+ "model_name": "HTMLModel",
+ "state": {
+ "_dom_classes": [],
+ "_model_module": "@jupyter-widgets/controls",
+ "_model_module_version": "1.5.0",
+ "_model_name": "HTMLModel",
+ "_view_count": null,
+ "_view_module": "@jupyter-widgets/controls",
+ "_view_module_version": "1.5.0",
+ "_view_name": "HTMLView",
+ "description": "",
+ "description_tooltip": null,
+ "layout": "IPY_MODEL_044920317fb64cd088d49f334930b886",
+ "placeholder": "",
+ "style": "IPY_MODEL_a8416e76428b46c092c7afbf7129b5f9",
+ "value": "BiRefNet_config.py: 100%"
+ }
+ },
+ "1f1a0c0dc2b74d56b68a779ef944416d": {
+ "model_module": "@jupyter-widgets/base",
+ "model_module_version": "1.2.0",
+ "model_name": "LayoutModel",
+ "state": {
+ "_model_module": "@jupyter-widgets/base",
+ "_model_module_version": "1.2.0",
+ "_model_name": "LayoutModel",
+ "_view_count": null,
+ "_view_module": "@jupyter-widgets/base",
+ "_view_module_version": "1.2.0",
+ "_view_name": "LayoutView",
+ "align_content": null,
+ "align_items": null,
+ "align_self": null,
+ "border": null,
+ "bottom": null,
+ "display": null,
+ "flex": null,
+ "flex_flow": null,
+ "grid_area": null,
+ "grid_auto_columns": null,
+ "grid_auto_flow": null,
+ "grid_auto_rows": null,
+ "grid_column": null,
+ "grid_gap": null,
+ "grid_row": null,
+ "grid_template_areas": null,
+ "grid_template_columns": null,
+ "grid_template_rows": null,
+ "height": null,
+ "justify_content": null,
+ "justify_items": null,
+ "left": null,
+ "margin": null,
+ "max_height": null,
+ "max_width": null,
+ "min_height": null,
+ "min_width": null,
+ "object_fit": null,
+ "object_position": null,
+ "order": null,
+ "overflow": null,
+ "overflow_x": null,
+ "overflow_y": null,
+ "padding": null,
+ "right": null,
+ "top": null,
+ "visibility": null,
+ "width": null
+ }
+ },
+ "2ad78a6887eb4a369bac75e2436758b9": {
+ "model_module": "@jupyter-widgets/base",
+ "model_module_version": "1.2.0",
+ "model_name": "LayoutModel",
+ "state": {
+ "_model_module": "@jupyter-widgets/base",
+ "_model_module_version": "1.2.0",
+ "_model_name": "LayoutModel",
+ "_view_count": null,
+ "_view_module": "@jupyter-widgets/base",
+ "_view_module_version": "1.2.0",
+ "_view_name": "LayoutView",
+ "align_content": null,
+ "align_items": null,
+ "align_self": null,
+ "border": null,
+ "bottom": null,
+ "display": null,
+ "flex": null,
+ "flex_flow": null,
+ "grid_area": null,
+ "grid_auto_columns": null,
+ "grid_auto_flow": null,
+ "grid_auto_rows": null,
+ "grid_column": null,
+ "grid_gap": null,
+ "grid_row": null,
+ "grid_template_areas": null,
+ "grid_template_columns": null,
+ "grid_template_rows": null,
+ "height": null,
+ "justify_content": null,
+ "justify_items": null,
+ "left": null,
+ "margin": null,
+ "max_height": null,
+ "max_width": null,
+ "min_height": null,
+ "min_width": null,
+ "object_fit": null,
+ "object_position": null,
+ "order": null,
+ "overflow": null,
+ "overflow_x": null,
+ "overflow_y": null,
+ "padding": null,
+ "right": null,
+ "top": null,
+ "visibility": null,
+ "width": null
+ }
+ },
+ "3b0ddb32ffa442aab3b02a22432bb233": {
+ "model_module": "@jupyter-widgets/controls",
+ "model_module_version": "1.5.0",
+ "model_name": "HTMLModel",
+ "state": {
+ "_dom_classes": [],
+ "_model_module": "@jupyter-widgets/controls",
+ "_model_module_version": "1.5.0",
+ "_model_name": "HTMLModel",
+ "_view_count": null,
+ "_view_module": "@jupyter-widgets/controls",
+ "_view_module_version": "1.5.0",
+ "_view_name": "HTMLView",
+ "description": "",
+ "description_tooltip": null,
+ "layout": "IPY_MODEL_a124da94c5b143a29d9420eba3954859",
+ "placeholder": "",
+ "style": "IPY_MODEL_96485cf384484a709e4c32b74e0223af",
+ "value": " 413/413 [00:00<00:00, 5.97kB/s]"
+ }
+ },
+ "3c9b1a3ed2f64ed58a29083280098292": {
+ "model_module": "@jupyter-widgets/controls",
+ "model_module_version": "1.5.0",
+ "model_name": "HTMLModel",
+ "state": {
+ "_dom_classes": [],
+ "_model_module": "@jupyter-widgets/controls",
+ "_model_module_version": "1.5.0",
+ "_model_name": "HTMLModel",
+ "_view_count": null,
+ "_view_module": "@jupyter-widgets/controls",
+ "_view_module_version": "1.5.0",
+ "_view_name": "HTMLView",
+ "description": "",
+ "description_tooltip": null,
+ "layout": "IPY_MODEL_cf459ec049624ba29f54cceeb1469785",
+ "placeholder": "",
+ "style": "IPY_MODEL_db64a13268ab452db1241d06f16d3d6a",
+ "value": " 298/298 [00:00<00:00, 9.24kB/s]"
+ }
+ },
+ "3e8d434d8e524c8c9f3fe3359df4b327": {
+ "model_module": "@jupyter-widgets/controls",
+ "model_module_version": "1.5.0",
+ "model_name": "DescriptionStyleModel",
+ "state": {
+ "_model_module": "@jupyter-widgets/controls",
+ "_model_module_version": "1.5.0",
+ "_model_name": "DescriptionStyleModel",
+ "_view_count": null,
+ "_view_module": "@jupyter-widgets/base",
+ "_view_module_version": "1.2.0",
+ "_view_name": "StyleView",
+ "description_width": ""
+ }
+ },
+ "4480d4d37c2648e7bb5635109b5d4715": {
+ "model_module": "@jupyter-widgets/base",
+ "model_module_version": "1.2.0",
+ "model_name": "LayoutModel",
+ "state": {
+ "_model_module": "@jupyter-widgets/base",
+ "_model_module_version": "1.2.0",
+ "_model_name": "LayoutModel",
+ "_view_count": null,
+ "_view_module": "@jupyter-widgets/base",
+ "_view_module_version": "1.2.0",
+ "_view_name": "LayoutView",
+ "align_content": null,
+ "align_items": null,
+ "align_self": null,
+ "border": null,
+ "bottom": null,
+ "display": null,
+ "flex": null,
+ "flex_flow": null,
+ "grid_area": null,
+ "grid_auto_columns": null,
+ "grid_auto_flow": null,
+ "grid_auto_rows": null,
+ "grid_column": null,
+ "grid_gap": null,
+ "grid_row": null,
+ "grid_template_areas": null,
+ "grid_template_columns": null,
+ "grid_template_rows": null,
+ "height": null,
+ "justify_content": null,
+ "justify_items": null,
+ "left": null,
+ "margin": null,
+ "max_height": null,
+ "max_width": null,
+ "min_height": null,
+ "min_width": null,
+ "object_fit": null,
+ "object_position": null,
+ "order": null,
+ "overflow": null,
+ "overflow_x": null,
+ "overflow_y": null,
+ "padding": null,
+ "right": null,
+ "top": null,
+ "visibility": null,
+ "width": null
+ }
+ },
+ "4d162f2efc364cf793709f0b9c9dc888": {
+ "model_module": "@jupyter-widgets/controls",
+ "model_module_version": "1.5.0",
+ "model_name": "HTMLModel",
+ "state": {
+ "_dom_classes": [],
+ "_model_module": "@jupyter-widgets/controls",
+ "_model_module_version": "1.5.0",
+ "_model_name": "HTMLModel",
+ "_view_count": null,
+ "_view_module": "@jupyter-widgets/controls",
+ "_view_module_version": "1.5.0",
+ "_view_name": "HTMLView",
+ "description": "",
+ "description_tooltip": null,
+ "layout": "IPY_MODEL_2ad78a6887eb4a369bac75e2436758b9",
+ "placeholder": "",
+ "style": "IPY_MODEL_78d3ab060eb64941bcba51fa3b32f493",
+ "value": "model.safetensors: 100%"
+ }
+ },
+ "4e5fe4296291455f88ae0e5f257c398e": {
+ "model_module": "@jupyter-widgets/base",
+ "model_module_version": "1.2.0",
+ "model_name": "LayoutModel",
+ "state": {
+ "_model_module": "@jupyter-widgets/base",
+ "_model_module_version": "1.2.0",
+ "_model_name": "LayoutModel",
+ "_view_count": null,
+ "_view_module": "@jupyter-widgets/base",
+ "_view_module_version": "1.2.0",
+ "_view_name": "LayoutView",
+ "align_content": null,
+ "align_items": null,
+ "align_self": null,
+ "border": null,
+ "bottom": null,
+ "display": null,
+ "flex": null,
+ "flex_flow": null,
+ "grid_area": null,
+ "grid_auto_columns": null,
+ "grid_auto_flow": null,
+ "grid_auto_rows": null,
+ "grid_column": null,
+ "grid_gap": null,
+ "grid_row": null,
+ "grid_template_areas": null,
+ "grid_template_columns": null,
+ "grid_template_rows": null,
+ "height": null,
+ "justify_content": null,
+ "justify_items": null,
+ "left": null,
+ "margin": null,
+ "max_height": null,
+ "max_width": null,
+ "min_height": null,
+ "min_width": null,
+ "object_fit": null,
+ "object_position": null,
+ "order": null,
+ "overflow": null,
+ "overflow_x": null,
+ "overflow_y": null,
+ "padding": null,
+ "right": null,
+ "top": null,
+ "visibility": null,
+ "width": null
+ }
+ },
+ "56eb4e5291404f3289196d526fef524f": {
+ "model_module": "@jupyter-widgets/base",
+ "model_module_version": "1.2.0",
+ "model_name": "LayoutModel",
+ "state": {
+ "_model_module": "@jupyter-widgets/base",
+ "_model_module_version": "1.2.0",
+ "_model_name": "LayoutModel",
+ "_view_count": null,
+ "_view_module": "@jupyter-widgets/base",
+ "_view_module_version": "1.2.0",
+ "_view_name": "LayoutView",
+ "align_content": null,
+ "align_items": null,
+ "align_self": null,
+ "border": null,
+ "bottom": null,
+ "display": null,
+ "flex": null,
+ "flex_flow": null,
+ "grid_area": null,
+ "grid_auto_columns": null,
+ "grid_auto_flow": null,
+ "grid_auto_rows": null,
+ "grid_column": null,
+ "grid_gap": null,
+ "grid_row": null,
+ "grid_template_areas": null,
+ "grid_template_columns": null,
+ "grid_template_rows": null,
+ "height": null,
+ "justify_content": null,
+ "justify_items": null,
+ "left": null,
+ "margin": null,
+ "max_height": null,
+ "max_width": null,
+ "min_height": null,
+ "min_width": null,
+ "object_fit": null,
+ "object_position": null,
+ "order": null,
+ "overflow": null,
+ "overflow_x": null,
+ "overflow_y": null,
+ "padding": null,
+ "right": null,
+ "top": null,
+ "visibility": null,
+ "width": null
+ }
+ },
+ "60d3c15d4b944546992d3517bb73912a": {
+ "model_module": "@jupyter-widgets/controls",
+ "model_module_version": "1.5.0",
+ "model_name": "ProgressStyleModel",
+ "state": {
+ "_model_module": "@jupyter-widgets/controls",
+ "_model_module_version": "1.5.0",
+ "_model_name": "ProgressStyleModel",
+ "_view_count": null,
+ "_view_module": "@jupyter-widgets/base",
+ "_view_module_version": "1.2.0",
+ "_view_name": "StyleView",
+ "bar_color": null,
+ "description_width": ""
+ }
+ },
+ "61746c0aaf96444391d1a86ea7223cb0": {
+ "model_module": "@jupyter-widgets/base",
+ "model_module_version": "1.2.0",
+ "model_name": "LayoutModel",
+ "state": {
+ "_model_module": "@jupyter-widgets/base",
+ "_model_module_version": "1.2.0",
+ "_model_name": "LayoutModel",
+ "_view_count": null,
+ "_view_module": "@jupyter-widgets/base",
+ "_view_module_version": "1.2.0",
+ "_view_name": "LayoutView",
+ "align_content": null,
+ "align_items": null,
+ "align_self": null,
+ "border": null,
+ "bottom": null,
+ "display": null,
+ "flex": null,
+ "flex_flow": null,
+ "grid_area": null,
+ "grid_auto_columns": null,
+ "grid_auto_flow": null,
+ "grid_auto_rows": null,
+ "grid_column": null,
+ "grid_gap": null,
+ "grid_row": null,
+ "grid_template_areas": null,
+ "grid_template_columns": null,
+ "grid_template_rows": null,
+ "height": null,
+ "justify_content": null,
+ "justify_items": null,
+ "left": null,
+ "margin": null,
+ "max_height": null,
+ "max_width": null,
+ "min_height": null,
+ "min_width": null,
+ "object_fit": null,
+ "object_position": null,
+ "order": null,
+ "overflow": null,
+ "overflow_x": null,
+ "overflow_y": null,
+ "padding": null,
+ "right": null,
+ "top": null,
+ "visibility": null,
+ "width": null
+ }
+ },
+ "6588cde28032444980c4e314d6b9a648": {
+ "model_module": "@jupyter-widgets/base",
+ "model_module_version": "1.2.0",
+ "model_name": "LayoutModel",
+ "state": {
+ "_model_module": "@jupyter-widgets/base",
+ "_model_module_version": "1.2.0",
+ "_model_name": "LayoutModel",
+ "_view_count": null,
+ "_view_module": "@jupyter-widgets/base",
+ "_view_module_version": "1.2.0",
+ "_view_name": "LayoutView",
+ "align_content": null,
+ "align_items": null,
+ "align_self": null,
+ "border": null,
+ "bottom": null,
+ "display": null,
+ "flex": null,
+ "flex_flow": null,
+ "grid_area": null,
+ "grid_auto_columns": null,
+ "grid_auto_flow": null,
+ "grid_auto_rows": null,
+ "grid_column": null,
+ "grid_gap": null,
+ "grid_row": null,
+ "grid_template_areas": null,
+ "grid_template_columns": null,
+ "grid_template_rows": null,
+ "height": null,
+ "justify_content": null,
+ "justify_items": null,
+ "left": null,
+ "margin": null,
+ "max_height": null,
+ "max_width": null,
+ "min_height": null,
+ "min_width": null,
+ "object_fit": null,
+ "object_position": null,
+ "order": null,
+ "overflow": null,
+ "overflow_x": null,
+ "overflow_y": null,
+ "padding": null,
+ "right": null,
+ "top": null,
+ "visibility": null,
+ "width": null
+ }
+ },
+ "6667788f1fa44603a4a04ba8aa5e1cc9": {
+ "model_module": "@jupyter-widgets/controls",
+ "model_module_version": "1.5.0",
+ "model_name": "HBoxModel",
+ "state": {
+ "_dom_classes": [],
+ "_model_module": "@jupyter-widgets/controls",
+ "_model_module_version": "1.5.0",
+ "_model_name": "HBoxModel",
+ "_view_count": null,
+ "_view_module": "@jupyter-widgets/controls",
+ "_view_module_version": "1.5.0",
+ "_view_name": "HBoxView",
+ "box_style": "",
+ "children": [
+ "IPY_MODEL_4d162f2efc364cf793709f0b9c9dc888",
+ "IPY_MODEL_e078f8d23e8449e7a9b5771e342febeb",
+ "IPY_MODEL_a330873af8cf45eca95780c50260dbb9"
+ ],
+ "layout": "IPY_MODEL_61746c0aaf96444391d1a86ea7223cb0"
+ }
+ },
+ "67ce0f56b38941749bdb580a830598c2": {
+ "model_module": "@jupyter-widgets/controls",
+ "model_module_version": "1.5.0",
+ "model_name": "DescriptionStyleModel",
+ "state": {
+ "_model_module": "@jupyter-widgets/controls",
+ "_model_module_version": "1.5.0",
+ "_model_name": "DescriptionStyleModel",
+ "_view_count": null,
+ "_view_module": "@jupyter-widgets/base",
+ "_view_module_version": "1.2.0",
+ "_view_name": "StyleView",
+ "description_width": ""
+ }
+ },
+ "688a14cd34704e4ea2e261f6619449ee": {
+ "model_module": "@jupyter-widgets/base",
+ "model_module_version": "1.2.0",
+ "model_name": "LayoutModel",
+ "state": {
+ "_model_module": "@jupyter-widgets/base",
+ "_model_module_version": "1.2.0",
+ "_model_name": "LayoutModel",
+ "_view_count": null,
+ "_view_module": "@jupyter-widgets/base",
+ "_view_module_version": "1.2.0",
+ "_view_name": "LayoutView",
+ "align_content": null,
+ "align_items": null,
+ "align_self": null,
+ "border": null,
+ "bottom": null,
+ "display": null,
+ "flex": null,
+ "flex_flow": null,
+ "grid_area": null,
+ "grid_auto_columns": null,
+ "grid_auto_flow": null,
+ "grid_auto_rows": null,
+ "grid_column": null,
+ "grid_gap": null,
+ "grid_row": null,
+ "grid_template_areas": null,
+ "grid_template_columns": null,
+ "grid_template_rows": null,
+ "height": null,
+ "justify_content": null,
+ "justify_items": null,
+ "left": null,
+ "margin": null,
+ "max_height": null,
+ "max_width": null,
+ "min_height": null,
+ "min_width": null,
+ "object_fit": null,
+ "object_position": null,
+ "order": null,
+ "overflow": null,
+ "overflow_x": null,
+ "overflow_y": null,
+ "padding": null,
+ "right": null,
+ "top": null,
+ "visibility": null,
+ "width": null
+ }
+ },
+ "6f45e75849f749a1b5fad6e8f7879c8f": {
+ "model_module": "@jupyter-widgets/controls",
+ "model_module_version": "1.5.0",
+ "model_name": "FloatProgressModel",
+ "state": {
+ "_dom_classes": [],
+ "_model_module": "@jupyter-widgets/controls",
+ "_model_module_version": "1.5.0",
+ "_model_name": "FloatProgressModel",
+ "_view_count": null,
+ "_view_module": "@jupyter-widgets/controls",
+ "_view_module_version": "1.5.0",
+ "_view_name": "ProgressView",
+ "bar_style": "success",
+ "description": "",
+ "description_tooltip": null,
+ "layout": "IPY_MODEL_e02fc794fb3645bda6b802277a3e5c1a",
+ "max": 91316,
+ "min": 0,
+ "orientation": "horizontal",
+ "style": "IPY_MODEL_f9a719112f53400782f97d0683862a5f",
+ "value": 91316
+ }
+ },
+ "7745af46aa694f5b8ff0ec8fbd72025f": {
+ "model_module": "@jupyter-widgets/controls",
+ "model_module_version": "1.5.0",
+ "model_name": "DescriptionStyleModel",
+ "state": {
+ "_model_module": "@jupyter-widgets/controls",
+ "_model_module_version": "1.5.0",
+ "_model_name": "DescriptionStyleModel",
+ "_view_count": null,
+ "_view_module": "@jupyter-widgets/base",
+ "_view_module_version": "1.2.0",
+ "_view_name": "StyleView",
+ "description_width": ""
+ }
+ },
+ "78d3ab060eb64941bcba51fa3b32f493": {
+ "model_module": "@jupyter-widgets/controls",
+ "model_module_version": "1.5.0",
+ "model_name": "DescriptionStyleModel",
+ "state": {
+ "_model_module": "@jupyter-widgets/controls",
+ "_model_module_version": "1.5.0",
+ "_model_name": "DescriptionStyleModel",
+ "_view_count": null,
+ "_view_module": "@jupyter-widgets/base",
+ "_view_module_version": "1.2.0",
+ "_view_name": "StyleView",
+ "description_width": ""
+ }
+ },
+ "7bef3ce58df040c986a74eb5334f42c9": {
+ "model_module": "@jupyter-widgets/controls",
+ "model_module_version": "1.5.0",
+ "model_name": "HBoxModel",
+ "state": {
+ "_dom_classes": [],
+ "_model_module": "@jupyter-widgets/controls",
+ "_model_module_version": "1.5.0",
+ "_model_name": "HBoxModel",
+ "_view_count": null,
+ "_view_module": "@jupyter-widgets/controls",
+ "_view_module_version": "1.5.0",
+ "_view_name": "HBoxView",
+ "box_style": "",
+ "children": [
+ "IPY_MODEL_1848689bf6a14235b508647d1bdb7f69",
+ "IPY_MODEL_055ec5d0d4da4d1b8df33ddb22d6b55b",
+ "IPY_MODEL_3c9b1a3ed2f64ed58a29083280098292"
+ ],
+ "layout": "IPY_MODEL_b2b95dd8b75d4625ad1bd02a461f304d"
+ }
+ },
+ "7d19deaab4c845eea4705567bdc65d60": {
+ "model_module": "@jupyter-widgets/controls",
+ "model_module_version": "1.5.0",
+ "model_name": "HBoxModel",
+ "state": {
+ "_dom_classes": [],
+ "_model_module": "@jupyter-widgets/controls",
+ "_model_module_version": "1.5.0",
+ "_model_name": "HBoxModel",
+ "_view_count": null,
+ "_view_module": "@jupyter-widgets/controls",
+ "_view_module_version": "1.5.0",
+ "_view_name": "HBoxView",
+ "box_style": "",
+ "children": [
+ "IPY_MODEL_a8941bdff0984189be91fab5bfe1c52c",
+ "IPY_MODEL_b6be81c6cc1e4608a88c785c448bfaa4",
+ "IPY_MODEL_3b0ddb32ffa442aab3b02a22432bb233"
+ ],
+ "layout": "IPY_MODEL_688a14cd34704e4ea2e261f6619449ee"
+ }
+ },
+ "7fb491d5b8f34a2e842518c1e4ea4906": {
+ "model_module": "@jupyter-widgets/controls",
+ "model_module_version": "1.5.0",
+ "model_name": "DescriptionStyleModel",
+ "state": {
+ "_model_module": "@jupyter-widgets/controls",
+ "_model_module_version": "1.5.0",
+ "_model_name": "DescriptionStyleModel",
+ "_view_count": null,
+ "_view_module": "@jupyter-widgets/base",
+ "_view_module_version": "1.2.0",
+ "_view_name": "StyleView",
+ "description_width": ""
+ }
+ },
+ "80a53775d54e47b2922031cc4cd00548": {
+ "model_module": "@jupyter-widgets/controls",
+ "model_module_version": "1.5.0",
+ "model_name": "HBoxModel",
+ "state": {
+ "_dom_classes": [],
+ "_model_module": "@jupyter-widgets/controls",
+ "_model_module_version": "1.5.0",
+ "_model_name": "HBoxModel",
+ "_view_count": null,
+ "_view_module": "@jupyter-widgets/controls",
+ "_view_module_version": "1.5.0",
+ "_view_name": "HBoxView",
+ "box_style": "",
+ "children": [
+ "IPY_MODEL_e13013a9d20843bca6da8fe9f0fdb644",
+ "IPY_MODEL_6f45e75849f749a1b5fad6e8f7879c8f",
+ "IPY_MODEL_0ef4b41f8d0141859b7767d112596198"
+ ],
+ "layout": "IPY_MODEL_fecc8e75d8d643759123eaf2ff30fa2a"
+ }
+ },
+ "8a54f72e65a24d7a93852aca8ecae0a2": {
+ "model_module": "@jupyter-widgets/base",
+ "model_module_version": "1.2.0",
+ "model_name": "LayoutModel",
+ "state": {
+ "_model_module": "@jupyter-widgets/base",
+ "_model_module_version": "1.2.0",
+ "_model_name": "LayoutModel",
+ "_view_count": null,
+ "_view_module": "@jupyter-widgets/base",
+ "_view_module_version": "1.2.0",
+ "_view_name": "LayoutView",
+ "align_content": null,
+ "align_items": null,
+ "align_self": null,
+ "border": null,
+ "bottom": null,
+ "display": null,
+ "flex": null,
+ "flex_flow": null,
+ "grid_area": null,
+ "grid_auto_columns": null,
+ "grid_auto_flow": null,
+ "grid_auto_rows": null,
+ "grid_column": null,
+ "grid_gap": null,
+ "grid_row": null,
+ "grid_template_areas": null,
+ "grid_template_columns": null,
+ "grid_template_rows": null,
+ "height": null,
+ "justify_content": null,
+ "justify_items": null,
+ "left": null,
+ "margin": null,
+ "max_height": null,
+ "max_width": null,
+ "min_height": null,
+ "min_width": null,
+ "object_fit": null,
+ "object_position": null,
+ "order": null,
+ "overflow": null,
+ "overflow_x": null,
+ "overflow_y": null,
+ "padding": null,
+ "right": null,
+ "top": null,
+ "visibility": null,
+ "width": null
+ }
+ },
+ "96485cf384484a709e4c32b74e0223af": {
+ "model_module": "@jupyter-widgets/controls",
+ "model_module_version": "1.5.0",
+ "model_name": "DescriptionStyleModel",
+ "state": {
+ "_model_module": "@jupyter-widgets/controls",
+ "_model_module_version": "1.5.0",
+ "_model_name": "DescriptionStyleModel",
+ "_view_count": null,
+ "_view_module": "@jupyter-widgets/base",
+ "_view_module_version": "1.2.0",
+ "_view_name": "StyleView",
+ "description_width": ""
+ }
+ },
+ "a124da94c5b143a29d9420eba3954859": {
+ "model_module": "@jupyter-widgets/base",
+ "model_module_version": "1.2.0",
+ "model_name": "LayoutModel",
+ "state": {
+ "_model_module": "@jupyter-widgets/base",
+ "_model_module_version": "1.2.0",
+ "_model_name": "LayoutModel",
+ "_view_count": null,
+ "_view_module": "@jupyter-widgets/base",
+ "_view_module_version": "1.2.0",
+ "_view_name": "LayoutView",
+ "align_content": null,
+ "align_items": null,
+ "align_self": null,
+ "border": null,
+ "bottom": null,
+ "display": null,
+ "flex": null,
+ "flex_flow": null,
+ "grid_area": null,
+ "grid_auto_columns": null,
+ "grid_auto_flow": null,
+ "grid_auto_rows": null,
+ "grid_column": null,
+ "grid_gap": null,
+ "grid_row": null,
+ "grid_template_areas": null,
+ "grid_template_columns": null,
+ "grid_template_rows": null,
+ "height": null,
+ "justify_content": null,
+ "justify_items": null,
+ "left": null,
+ "margin": null,
+ "max_height": null,
+ "max_width": null,
+ "min_height": null,
+ "min_width": null,
+ "object_fit": null,
+ "object_position": null,
+ "order": null,
+ "overflow": null,
+ "overflow_x": null,
+ "overflow_y": null,
+ "padding": null,
+ "right": null,
+ "top": null,
+ "visibility": null,
+ "width": null
+ }
+ },
+ "a330873af8cf45eca95780c50260dbb9": {
+ "model_module": "@jupyter-widgets/controls",
+ "model_module_version": "1.5.0",
+ "model_name": "HTMLModel",
+ "state": {
+ "_dom_classes": [],
+ "_model_module": "@jupyter-widgets/controls",
+ "_model_module_version": "1.5.0",
+ "_model_name": "HTMLModel",
+ "_view_count": null,
+ "_view_module": "@jupyter-widgets/controls",
+ "_view_module_version": "1.5.0",
+ "_view_name": "HTMLView",
+ "description": "",
+ "description_tooltip": null,
+ "layout": "IPY_MODEL_6588cde28032444980c4e314d6b9a648",
+ "placeholder": "",
+ "style": "IPY_MODEL_3e8d434d8e524c8c9f3fe3359df4b327",
+ "value": " 885M/885M [00:05<00:00, 192MB/s]"
+ }
+ },
+ "a8416e76428b46c092c7afbf7129b5f9": {
+ "model_module": "@jupyter-widgets/controls",
+ "model_module_version": "1.5.0",
+ "model_name": "DescriptionStyleModel",
+ "state": {
+ "_model_module": "@jupyter-widgets/controls",
+ "_model_module_version": "1.5.0",
+ "_model_name": "DescriptionStyleModel",
+ "_view_count": null,
+ "_view_module": "@jupyter-widgets/base",
+ "_view_module_version": "1.2.0",
+ "_view_name": "StyleView",
+ "description_width": ""
+ }
+ },
+ "a8941bdff0984189be91fab5bfe1c52c": {
+ "model_module": "@jupyter-widgets/controls",
+ "model_module_version": "1.5.0",
+ "model_name": "HTMLModel",
+ "state": {
+ "_dom_classes": [],
+ "_model_module": "@jupyter-widgets/controls",
+ "_model_module_version": "1.5.0",
+ "_model_name": "HTMLModel",
+ "_view_count": null,
+ "_view_module": "@jupyter-widgets/controls",
+ "_view_module_version": "1.5.0",
+ "_view_name": "HTMLView",
+ "description": "",
+ "description_tooltip": null,
+ "layout": "IPY_MODEL_8a54f72e65a24d7a93852aca8ecae0a2",
+ "placeholder": "",
+ "style": "IPY_MODEL_7745af46aa694f5b8ff0ec8fbd72025f",
+ "value": "config.json: 100%"
+ }
+ },
+ "b2b95dd8b75d4625ad1bd02a461f304d": {
+ "model_module": "@jupyter-widgets/base",
+ "model_module_version": "1.2.0",
+ "model_name": "LayoutModel",
+ "state": {
+ "_model_module": "@jupyter-widgets/base",
+ "_model_module_version": "1.2.0",
+ "_model_name": "LayoutModel",
+ "_view_count": null,
+ "_view_module": "@jupyter-widgets/base",
+ "_view_module_version": "1.2.0",
+ "_view_name": "LayoutView",
+ "align_content": null,
+ "align_items": null,
+ "align_self": null,
+ "border": null,
+ "bottom": null,
+ "display": null,
+ "flex": null,
+ "flex_flow": null,
+ "grid_area": null,
+ "grid_auto_columns": null,
+ "grid_auto_flow": null,
+ "grid_auto_rows": null,
+ "grid_column": null,
+ "grid_gap": null,
+ "grid_row": null,
+ "grid_template_areas": null,
+ "grid_template_columns": null,
+ "grid_template_rows": null,
+ "height": null,
+ "justify_content": null,
+ "justify_items": null,
+ "left": null,
+ "margin": null,
+ "max_height": null,
+ "max_width": null,
+ "min_height": null,
+ "min_width": null,
+ "object_fit": null,
+ "object_position": null,
+ "order": null,
+ "overflow": null,
+ "overflow_x": null,
+ "overflow_y": null,
+ "padding": null,
+ "right": null,
+ "top": null,
+ "visibility": null,
+ "width": null
+ }
+ },
+ "b6be81c6cc1e4608a88c785c448bfaa4": {
+ "model_module": "@jupyter-widgets/controls",
+ "model_module_version": "1.5.0",
+ "model_name": "FloatProgressModel",
+ "state": {
+ "_dom_classes": [],
+ "_model_module": "@jupyter-widgets/controls",
+ "_model_module_version": "1.5.0",
+ "_model_name": "FloatProgressModel",
+ "_view_count": null,
+ "_view_module": "@jupyter-widgets/controls",
+ "_view_module_version": "1.5.0",
+ "_view_name": "ProgressView",
+ "bar_style": "success",
+ "description": "",
+ "description_tooltip": null,
+ "layout": "IPY_MODEL_4e5fe4296291455f88ae0e5f257c398e",
+ "max": 413,
+ "min": 0,
+ "orientation": "horizontal",
+ "style": "IPY_MODEL_60d3c15d4b944546992d3517bb73912a",
+ "value": 413
+ }
+ },
+ "cf459ec049624ba29f54cceeb1469785": {
+ "model_module": "@jupyter-widgets/base",
+ "model_module_version": "1.2.0",
+ "model_name": "LayoutModel",
+ "state": {
+ "_model_module": "@jupyter-widgets/base",
+ "_model_module_version": "1.2.0",
+ "_model_name": "LayoutModel",
+ "_view_count": null,
+ "_view_module": "@jupyter-widgets/base",
+ "_view_module_version": "1.2.0",
+ "_view_name": "LayoutView",
+ "align_content": null,
+ "align_items": null,
+ "align_self": null,
+ "border": null,
+ "bottom": null,
+ "display": null,
+ "flex": null,
+ "flex_flow": null,
+ "grid_area": null,
+ "grid_auto_columns": null,
+ "grid_auto_flow": null,
+ "grid_auto_rows": null,
+ "grid_column": null,
+ "grid_gap": null,
+ "grid_row": null,
+ "grid_template_areas": null,
+ "grid_template_columns": null,
+ "grid_template_rows": null,
+ "height": null,
+ "justify_content": null,
+ "justify_items": null,
+ "left": null,
+ "margin": null,
+ "max_height": null,
+ "max_width": null,
+ "min_height": null,
+ "min_width": null,
+ "object_fit": null,
+ "object_position": null,
+ "order": null,
+ "overflow": null,
+ "overflow_x": null,
+ "overflow_y": null,
+ "padding": null,
+ "right": null,
+ "top": null,
+ "visibility": null,
+ "width": null
+ }
+ },
+ "db64a13268ab452db1241d06f16d3d6a": {
+ "model_module": "@jupyter-widgets/controls",
+ "model_module_version": "1.5.0",
+ "model_name": "DescriptionStyleModel",
+ "state": {
+ "_model_module": "@jupyter-widgets/controls",
+ "_model_module_version": "1.5.0",
+ "_model_name": "DescriptionStyleModel",
+ "_view_count": null,
+ "_view_module": "@jupyter-widgets/base",
+ "_view_module_version": "1.2.0",
+ "_view_name": "StyleView",
+ "description_width": ""
+ }
+ },
+ "e02fc794fb3645bda6b802277a3e5c1a": {
+ "model_module": "@jupyter-widgets/base",
+ "model_module_version": "1.2.0",
+ "model_name": "LayoutModel",
+ "state": {
+ "_model_module": "@jupyter-widgets/base",
+ "_model_module_version": "1.2.0",
+ "_model_name": "LayoutModel",
+ "_view_count": null,
+ "_view_module": "@jupyter-widgets/base",
+ "_view_module_version": "1.2.0",
+ "_view_name": "LayoutView",
+ "align_content": null,
+ "align_items": null,
+ "align_self": null,
+ "border": null,
+ "bottom": null,
+ "display": null,
+ "flex": null,
+ "flex_flow": null,
+ "grid_area": null,
+ "grid_auto_columns": null,
+ "grid_auto_flow": null,
+ "grid_auto_rows": null,
+ "grid_column": null,
+ "grid_gap": null,
+ "grid_row": null,
+ "grid_template_areas": null,
+ "grid_template_columns": null,
+ "grid_template_rows": null,
+ "height": null,
+ "justify_content": null,
+ "justify_items": null,
+ "left": null,
+ "margin": null,
+ "max_height": null,
+ "max_width": null,
+ "min_height": null,
+ "min_width": null,
+ "object_fit": null,
+ "object_position": null,
+ "order": null,
+ "overflow": null,
+ "overflow_x": null,
+ "overflow_y": null,
+ "padding": null,
+ "right": null,
+ "top": null,
+ "visibility": null,
+ "width": null
+ }
+ },
+ "e078f8d23e8449e7a9b5771e342febeb": {
+ "model_module": "@jupyter-widgets/controls",
+ "model_module_version": "1.5.0",
+ "model_name": "FloatProgressModel",
+ "state": {
+ "_dom_classes": [],
+ "_model_module": "@jupyter-widgets/controls",
+ "_model_module_version": "1.5.0",
+ "_model_name": "FloatProgressModel",
+ "_view_count": null,
+ "_view_module": "@jupyter-widgets/controls",
+ "_view_module_version": "1.5.0",
+ "_view_name": "ProgressView",
+ "bar_style": "success",
+ "description": "",
+ "description_tooltip": null,
+ "layout": "IPY_MODEL_1f1a0c0dc2b74d56b68a779ef944416d",
+ "max": 884878856,
+ "min": 0,
+ "orientation": "horizontal",
+ "style": "IPY_MODEL_1546b1d7383c4a25b2ab3782ff204cad",
+ "value": 884878856
+ }
+ },
+ "e13013a9d20843bca6da8fe9f0fdb644": {
+ "model_module": "@jupyter-widgets/controls",
+ "model_module_version": "1.5.0",
+ "model_name": "HTMLModel",
+ "state": {
+ "_dom_classes": [],
+ "_model_module": "@jupyter-widgets/controls",
+ "_model_module_version": "1.5.0",
+ "_model_name": "HTMLModel",
+ "_view_count": null,
+ "_view_module": "@jupyter-widgets/controls",
+ "_view_module_version": "1.5.0",
+ "_view_name": "HTMLView",
+ "description": "",
+ "description_tooltip": null,
+ "layout": "IPY_MODEL_56eb4e5291404f3289196d526fef524f",
+ "placeholder": "",
+ "style": "IPY_MODEL_7fb491d5b8f34a2e842518c1e4ea4906",
+ "value": "birefnet.py: 100%"
+ }
+ },
+ "ebf0b02bc7734ebea9a846233a1a6ec4": {
+ "model_module": "@jupyter-widgets/base",
+ "model_module_version": "1.2.0",
+ "model_name": "LayoutModel",
+ "state": {
+ "_model_module": "@jupyter-widgets/base",
+ "_model_module_version": "1.2.0",
+ "_model_name": "LayoutModel",
+ "_view_count": null,
+ "_view_module": "@jupyter-widgets/base",
+ "_view_module_version": "1.2.0",
+ "_view_name": "LayoutView",
+ "align_content": null,
+ "align_items": null,
+ "align_self": null,
+ "border": null,
+ "bottom": null,
+ "display": null,
+ "flex": null,
+ "flex_flow": null,
+ "grid_area": null,
+ "grid_auto_columns": null,
+ "grid_auto_flow": null,
+ "grid_auto_rows": null,
+ "grid_column": null,
+ "grid_gap": null,
+ "grid_row": null,
+ "grid_template_areas": null,
+ "grid_template_columns": null,
+ "grid_template_rows": null,
+ "height": null,
+ "justify_content": null,
+ "justify_items": null,
+ "left": null,
+ "margin": null,
+ "max_height": null,
+ "max_width": null,
+ "min_height": null,
+ "min_width": null,
+ "object_fit": null,
+ "object_position": null,
+ "order": null,
+ "overflow": null,
+ "overflow_x": null,
+ "overflow_y": null,
+ "padding": null,
+ "right": null,
+ "top": null,
+ "visibility": null,
+ "width": null
+ }
+ },
+ "f9a719112f53400782f97d0683862a5f": {
+ "model_module": "@jupyter-widgets/controls",
+ "model_module_version": "1.5.0",
+ "model_name": "ProgressStyleModel",
+ "state": {
+ "_model_module": "@jupyter-widgets/controls",
+ "_model_module_version": "1.5.0",
+ "_model_name": "ProgressStyleModel",
+ "_view_count": null,
+ "_view_module": "@jupyter-widgets/base",
+ "_view_module_version": "1.2.0",
+ "_view_name": "StyleView",
+ "bar_color": null,
+ "description_width": ""
+ }
+ },
+ "fb3e80729a214bc5993b62ee04d7c58d": {
+ "model_module": "@jupyter-widgets/controls",
+ "model_module_version": "1.5.0",
+ "model_name": "ProgressStyleModel",
+ "state": {
+ "_model_module": "@jupyter-widgets/controls",
+ "_model_module_version": "1.5.0",
+ "_model_name": "ProgressStyleModel",
+ "_view_count": null,
+ "_view_module": "@jupyter-widgets/base",
+ "_view_module_version": "1.2.0",
+ "_view_name": "StyleView",
+ "bar_color": null,
+ "description_width": ""
+ }
+ },
+ "fecc8e75d8d643759123eaf2ff30fa2a": {
+ "model_module": "@jupyter-widgets/base",
+ "model_module_version": "1.2.0",
+ "model_name": "LayoutModel",
+ "state": {
+ "_model_module": "@jupyter-widgets/base",
+ "_model_module_version": "1.2.0",
+ "_model_name": "LayoutModel",
+ "_view_count": null,
+ "_view_module": "@jupyter-widgets/base",
+ "_view_module_version": "1.2.0",
+ "_view_name": "LayoutView",
+ "align_content": null,
+ "align_items": null,
+ "align_self": null,
+ "border": null,
+ "bottom": null,
+ "display": null,
+ "flex": null,
+ "flex_flow": null,
+ "grid_area": null,
+ "grid_auto_columns": null,
+ "grid_auto_flow": null,
+ "grid_auto_rows": null,
+ "grid_column": null,
+ "grid_gap": null,
+ "grid_row": null,
+ "grid_template_areas": null,
+ "grid_template_columns": null,
+ "grid_template_rows": null,
+ "height": null,
+ "justify_content": null,
+ "justify_items": null,
+ "left": null,
+ "margin": null,
+ "max_height": null,
+ "max_width": null,
+ "min_height": null,
+ "min_width": null,
+ "object_fit": null,
+ "object_position": null,
+ "order": null,
+ "overflow": null,
+ "overflow_x": null,
+ "overflow_y": null,
+ "padding": null,
+ "right": null,
+ "top": null,
+ "visibility": null,
+ "width": null
+ }
+ }
+ }
+ }
+ },
+ "nbformat": 4,
+ "nbformat_minor": 0
+}
diff --git a/py/BiRefNet_v2/tutorials/BiRefNet_pth2onnx.ipynb b/py/BiRefNet_v2/tutorials/BiRefNet_pth2onnx.ipynb
new file mode 100644
index 0000000..c087b08
--- /dev/null
+++ b/py/BiRefNet_v2/tutorials/BiRefNet_pth2onnx.ipynb
@@ -0,0 +1,312 @@
+{
+ "cells": [
+ {
+ "cell_type": "markdown",
+ "metadata": {
+ "id": "LTj2A0RUQFNo"
+ },
+ "source": [
+ "# Convert our BiRefNet weights to onnx format.\n",
+ "\n",
+ "> This colab file is modified from [Kazuhito00](https://github.com/Kazuhito00)'s nice work.\n",
+ "\n",
+ "> Repo: https://github.com/Kazuhito00/BiRefNet-ONNX-Sample \n",
+ "> Original Colab: https://colab.research.google.com/github/Kazuhito00/BiRefNet-ONNX-Sample/blob/main/Convert2ONNX.ipynb\n",
+ "\n",
+ "+ Currently, Colab with 12.7GB RAM / 15GB GPU Mem cannot hold the transformation of BiRefNet in default setting. So, I take BiRefNet with swin_v1_tiny backbone as an example."
+ ]
+ },
+ {
+ "cell_type": "markdown",
+ "metadata": {},
+ "source": [
+ "### Online Colab version: https://colab.research.google.com/drive/1z6OruR52LOvDDpnp516F-N4EyPGrp5om"
+ ]
+ },
+ {
+ "cell_type": "code",
+ "execution_count": null,
+ "metadata": {
+ "id": "781JHjLJmveh"
+ },
+ "outputs": [],
+ "source": [
+ "import torch\n",
+ "\n",
+ "\n",
+ "weights_file = 'BiRefNet-general-bb_swin_v1_tiny-epoch_232.pth' # https://github.com/ZhengPeng7/BiRefNet/releases/download/v1/BiRefNet-general-bb_swin_v1_tiny-epoch_232.pth\n",
+ "device = 'cuda' if torch.cuda.is_available() else 'cpu'"
+ ]
+ },
+ {
+ "cell_type": "code",
+ "execution_count": null,
+ "metadata": {},
+ "outputs": [],
+ "source": [
+ "with open('config.py') as fp:\n",
+ " file_lines = fp.read()\n",
+ "if 'swin_v1_tiny' in weights_file:\n",
+ " print('Set `swin_v1_tiny` as the backbone.')\n",
+ " file_lines = file_lines.replace(\n",
+ " '''\n",
+ " 'pvt_v2_b2', 'pvt_v2_b5', # 9-bs10, 10-bs5\n",
+ " ][6]\n",
+ " ''',\n",
+ " '''\n",
+ " 'pvt_v2_b2', 'pvt_v2_b5', # 9-bs10, 10-bs5\n",
+ " ][3]\n",
+ " ''',\n",
+ " )\n",
+ " with open('config.py', mode=\"w\") as fp:\n",
+ " fp.write(file_lines)\n",
+ "else:\n",
+ " file_lines = file_lines.replace(\n",
+ " '''\n",
+ " 'pvt_v2_b2', 'pvt_v2_b5', # 9-bs10, 10-bs5\n",
+ " ][3]\n",
+ " ''',\n",
+ " '''\n",
+ " 'pvt_v2_b2', 'pvt_v2_b5', # 9-bs10, 10-bs5\n",
+ " ][6]\n",
+ " ''',\n",
+ " )\n",
+ " with open('config.py', mode=\"w\") as fp:\n",
+ " fp.write(file_lines)"
+ ]
+ },
+ {
+ "cell_type": "code",
+ "execution_count": null,
+ "metadata": {
+ "id": "7lFgKfPS8Icy"
+ },
+ "outputs": [],
+ "source": [
+ "from utils import check_state_dict\n",
+ "from models.birefnet import BiRefNet\n",
+ "\n",
+ "\n",
+ "birefnet = BiRefNet(bb_pretrained=False)\n",
+ "state_dict = torch.load('./{}'.format(weights_file), map_location=device)\n",
+ "state_dict = check_state_dict(state_dict)\n",
+ "birefnet.load_state_dict(state_dict)\n",
+ "\n",
+ "torch.set_float32_matmul_precision(['high', 'highest'][0])\n",
+ "\n",
+ "birefnet.to(device)\n",
+ "_ = birefnet.eval()"
+ ]
+ },
+ {
+ "cell_type": "markdown",
+ "metadata": {
+ "id": "JVgJAdgxQVJW"
+ },
+ "source": [
+ "# Process deform_conv2d in the conversion to ONNX"
+ ]
+ },
+ {
+ "cell_type": "code",
+ "execution_count": null,
+ "metadata": {
+ "id": "vJiZv0L75kTe"
+ },
+ "outputs": [],
+ "source": [
+ "from torchvision.ops.deform_conv import DeformConv2d\n",
+ "import deform_conv2d_onnx_exporter\n",
+ "\n",
+ "# register deform_conv2d operator\n",
+ "deform_conv2d_onnx_exporter.register_deform_conv2d_onnx_op()\n",
+ "\n",
+ "def convert_to_onnx(net, file_name='output.onnx', input_shape=(1024, 1024), device=device):\n",
+ " input = torch.randn(1, 3, input_shape[0], input_shape[1]).to(device)\n",
+ "\n",
+ " input_layer_names = ['input_image']\n",
+ " output_layer_names = ['output_image']\n",
+ "\n",
+ " torch.onnx.export(\n",
+ " net,\n",
+ " input,\n",
+ " file_name,\n",
+ " verbose=False,\n",
+ " opset_version=17,\n",
+ " input_names=input_layer_names,\n",
+ " output_names=output_layer_names,\n",
+ " )\n",
+ "convert_to_onnx(birefnet, weights_file.replace('.pth', '.onnx'), input_shape=(1024, 1024), device=device)"
+ ]
+ },
+ {
+ "cell_type": "markdown",
+ "metadata": {
+ "id": "-eU-g40P1zS-"
+ },
+ "source": [
+ "# Load ONNX weights and do the inference."
+ ]
+ },
+ {
+ "cell_type": "code",
+ "execution_count": null,
+ "metadata": {
+ "id": "LZ4HVqcoDvto"
+ },
+ "outputs": [],
+ "source": [
+ "from PIL import Image\n",
+ "from torchvision import transforms\n",
+ "\n",
+ "\n",
+ "transform_image = transforms.Compose([\n",
+ " transforms.Resize((1024, 1024)),\n",
+ " transforms.ToTensor(),\n",
+ " transforms.Normalize([0.485, 0.456, 0.406], [0.229, 0.224, 0.225])\n",
+ "])\n",
+ "\n",
+ "imagepath = './Helicopter-HR.jpg'\n",
+ "image = Image.open(imagepath)\n",
+ "input_images = transform_image(image).unsqueeze(0).to(device)\n",
+ "input_images_numpy = input_images.cpu().numpy()"
+ ]
+ },
+ {
+ "cell_type": "code",
+ "execution_count": null,
+ "metadata": {
+ "id": "rwzdKX1EfYkd"
+ },
+ "outputs": [],
+ "source": [
+ "import onnxruntime\n",
+ "import matplotlib.pyplot as plt\n",
+ "\n",
+ "\n",
+ "providers = ['CPUExecutionProvider'] if device == 'cpu' else ['CUDAExecutionProvider']\n",
+ "onnx_session = onnxruntime.InferenceSession(\n",
+ " weights_file.replace('.pth', '.onnx'),\n",
+ " providers=providers\n",
+ ")\n",
+ "input_name = onnx_session.get_inputs()[0].name\n",
+ "print(onnxruntime.get_device(), onnx_session.get_providers())"
+ ]
+ },
+ {
+ "cell_type": "code",
+ "execution_count": null,
+ "metadata": {
+ "id": "DJVtxZUZum4-"
+ },
+ "outputs": [],
+ "source": [
+ "from time import time\n",
+ "import matplotlib.pyplot as plt\n",
+ "\n",
+ "time_st = time()\n",
+ "pred_onnx = torch.tensor(\n",
+ " onnx_session.run(None, {input_name: input_images_numpy if device == 'cpu' else input_images_numpy})[-1]\n",
+ ").squeeze(0).sigmoid().cpu()\n",
+ "print(time() - time_st)\n",
+ "\n",
+ "plt.imshow(pred_onnx.squeeze(), cmap='gray'); plt.show()"
+ ]
+ },
+ {
+ "cell_type": "code",
+ "execution_count": null,
+ "metadata": {},
+ "outputs": [],
+ "source": [
+ "with torch.no_grad():\n",
+ " preds = birefnet(input_images)[-1].sigmoid().cpu()\n",
+ "plt.imshow(preds.squeeze(), cmap='gray'); plt.show()"
+ ]
+ },
+ {
+ "cell_type": "code",
+ "execution_count": null,
+ "metadata": {},
+ "outputs": [],
+ "source": [
+ "diff = abs(preds - pred_onnx)\n",
+ "print('sum(diff):', diff.sum())\n",
+ "plt.imshow((diff).squeeze(), cmap='gray'); plt.show()"
+ ]
+ },
+ {
+ "cell_type": "markdown",
+ "metadata": {
+ "id": "qzYHflt92Bjd"
+ },
+ "source": [
+ "# Efficiency Comparison between .pth and .onnx"
+ ]
+ },
+ {
+ "cell_type": "code",
+ "execution_count": null,
+ "metadata": {
+ "colab": {
+ "base_uri": "https://localhost:8080/"
+ },
+ "id": "A5IYfT-uzphA",
+ "outputId": "2999e345-950e-41b3-ddd3-9f58a71a3f21"
+ },
+ "outputs": [],
+ "source": [
+ "%%timeit\n",
+ "with torch.no_grad():\n",
+ " preds = birefnet(input_images)[-1].sigmoid().cpu()"
+ ]
+ },
+ {
+ "cell_type": "code",
+ "execution_count": null,
+ "metadata": {
+ "id": "G0Ul4rfNg1za"
+ },
+ "outputs": [],
+ "source": [
+ "%%timeit\n",
+ "pred_onnx = torch.tensor(\n",
+ " onnx_session.run(None, {input_name: input_images_numpy})[-1]\n",
+ ").squeeze(0).sigmoid().cpu()"
+ ]
+ },
+ {
+ "cell_type": "code",
+ "execution_count": null,
+ "metadata": {},
+ "outputs": [],
+ "source": []
+ }
+ ],
+ "metadata": {
+ "accelerator": "GPU",
+ "colab": {
+ "gpuType": "T4",
+ "provenance": []
+ },
+ "kernelspec": {
+ "display_name": "Python 3 (ipykernel)",
+ "language": "python",
+ "name": "python3"
+ },
+ "language_info": {
+ "codemirror_mode": {
+ "name": "ipython",
+ "version": 3
+ },
+ "file_extension": ".py",
+ "mimetype": "text/x-python",
+ "name": "python",
+ "nbconvert_exporter": "python",
+ "pygments_lexer": "ipython3",
+ "version": "3.10.14"
+ }
+ },
+ "nbformat": 4,
+ "nbformat_minor": 4
+}
diff --git a/py/BiRefNet_v2/utils.py b/py/BiRefNet_v2/utils.py
new file mode 100644
index 0000000..1b43754
--- /dev/null
+++ b/py/BiRefNet_v2/utils.py
@@ -0,0 +1,97 @@
+import logging
+import os
+import torch
+from torchvision import transforms
+import numpy as np
+import random
+import cv2
+from PIL import Image
+
+
+def path_to_image(path, size=(1024, 1024), color_type=['rgb', 'gray'][0]):
+ if color_type.lower() == 'rgb':
+ image = cv2.imread(path)
+ elif color_type.lower() == 'gray':
+ image = cv2.imread(path, cv2.IMREAD_GRAYSCALE)
+ else:
+ print('Select the color_type to return, either to RGB or gray image.')
+ return
+ if size:
+ image = cv2.resize(image, size, interpolation=cv2.INTER_LINEAR)
+ if color_type.lower() == 'rgb':
+ image = Image.fromarray(cv2.cvtColor(image, cv2.COLOR_BGR2RGB)).convert('RGB')
+ else:
+ image = Image.fromarray(image).convert('L')
+ return image
+
+
+
+def check_state_dict(state_dict, unwanted_prefix='_orig_mod.'):
+ for k, v in list(state_dict.items()):
+ if k.startswith(unwanted_prefix):
+ state_dict[k[len(unwanted_prefix):]] = state_dict.pop(k)
+ return state_dict
+
+
+def generate_smoothed_gt(gts):
+ epsilon = 0.001
+ new_gts = (1-epsilon)*gts+epsilon/2
+ return new_gts
+
+
+class Logger():
+ def __init__(self, path="log.txt"):
+ self.logger = logging.getLogger('BiRefNet')
+ self.file_handler = logging.FileHandler(path, "w")
+ self.stdout_handler = logging.StreamHandler()
+ self.stdout_handler.setFormatter(logging.Formatter('%(asctime)s %(levelname)s %(message)s'))
+ self.file_handler.setFormatter(logging.Formatter('%(asctime)s %(levelname)s %(message)s'))
+ self.logger.addHandler(self.file_handler)
+ self.logger.addHandler(self.stdout_handler)
+ self.logger.setLevel(logging.INFO)
+ self.logger.propagate = False
+
+ def info(self, txt):
+ self.logger.info(txt)
+
+ def close(self):
+ self.file_handler.close()
+ self.stdout_handler.close()
+
+
+class AverageMeter(object):
+ """Computes and stores the average and current value"""
+ def __init__(self):
+ self.reset()
+
+ def reset(self):
+ self.val = 0.0
+ self.avg = 0.0
+ self.sum = 0.0
+ self.count = 0.0
+
+ def update(self, val, n=1):
+ self.val = val
+ self.sum += val * n
+ self.count += n
+ self.avg = self.sum / self.count
+
+
+def save_checkpoint(state, path, filename="latest.pth"):
+ torch.save(state, os.path.join(path, filename))
+
+
+def save_tensor_img(tenor_im, path):
+ im = tenor_im.cpu().clone()
+ im = im.squeeze(0)
+ tensor2pil = transforms.ToPILImage()
+ im = tensor2pil(im)
+ im.save(path)
+
+
+def set_seed(seed):
+ torch.manual_seed(seed)
+ torch.cuda.manual_seed_all(seed)
+ np.random.seed(seed)
+ random.seed(seed)
+ torch.backends.cudnn.deterministic = True
diff --git a/py/Qwen_image2prompt.py b/py/Qwen_image2prompt.py
new file mode 100644
index 0000000..533e7fb
--- /dev/null
+++ b/py/Qwen_image2prompt.py
@@ -0,0 +1,60 @@
+# layerstyle advance
+
+import os.path
+from pathlib import Path
+import torch
+from PIL import Image
+import math
+from torchvision.transforms import ToPILImage
+import folder_paths
+from .imagefunc import files_for_uform_gen2_qwen, StopOnTokens, UformGen2QwenChat, clear_memory, log
+
+NODE_NAME = "QWenImage2Prompt"
+# Example of integrating UformGen2QwenChat into a node-like structure
+class QWenImage2Prompt:
+
+ @classmethod
+ def INPUT_TYPES(cls):
+ return {
+ "required": {
+ "image": ("IMAGE",),
+ "question": ("STRING", {"multiline": False, "default": "describe this image",},),
+ },
+ }
+
+ RETURN_TYPES = ("STRING",)
+ RETURN_NAMES = ("text",)
+ FUNCTION = "uform_gen2_qwen_chat"
+ CATEGORY = '😺dzNodes/LayerUtility/Prompt'
+
+ def uform_gen2_qwen_chat(self, image, question):
+ chat_model = UformGen2QwenChat()
+ history = [] # Example empty history
+ pil_image = ToPILImage()(image[0].permute(2, 0, 1))
+ width, height = pil_image.size
+ ratio = width / height
+ if width * height > 1024 * 1024:
+ target_width = math.sqrt(ratio * 1024 * 1024)
+ target_height = target_width / ratio
+ target_width = int(target_width)
+ target_height = int(target_height)
+ pil_image = pil_image.resize((target_width, target_height), Image.LANCZOS)
+ temp_path = files_for_uform_gen2_qwen / "temp.png"
+ pil_image.save(temp_path)
+ question = f"{question} but output no more then 80 words."
+ response = chat_model.chat_response(question, history, temp_path)
+
+ # Cleanup
+ del chat_model
+ clear_memory()
+ ret_text = response.split("assistant\n", 1)[1]
+ log(f"{NODE_NAME} Processed, Question: {question}, Response: {ret_text} ", message_type='finish')
+ return (ret_text, )
+
+NODE_CLASS_MAPPINGS = {
+ "LayerUtility: QWenImage2Prompt": QWenImage2Prompt
+}
+
+NODE_DISPLAY_NAME_MAPPINGS = {
+ "LayerUtility: QWenImage2Prompt": "LayerUtility: QWenImage2Prompt(Advance)"
+}
\ No newline at end of file
diff --git a/py/birefnet_legacy.py b/py/birefnet_legacy.py
new file mode 100644
index 0000000..02943c5
--- /dev/null
+++ b/py/birefnet_legacy.py
@@ -0,0 +1,83 @@
+# layerstyle advance
+
+from .imagefunc import *
+
+import torch.nn as nn
+from torchvision import transforms
+from .BiRefNet_legacy.baseline import BiRefNet
+from .BiRefNet_legacy.config import Config
+
+class BiRefNet_img_processor:
+ def __init__(self, config):
+ self.config = config
+ self.data_size = (config.size, config.size)
+ self.transform_image = transforms.Compose([
+ transforms.Resize(self.data_size),
+ transforms.ToTensor(),
+ transforms.Normalize([0.485, 0.456, 0.406], [0.229, 0.224, 0.225]),
+ ])
+
+ def __call__(self, _image: np.array):
+ _image_rs = cv2.resize(_image, (self.config.size, self.config.size), interpolation=cv2.INTER_LINEAR)
+ _image_rs = Image.fromarray(np.uint8(_image_rs*255)).convert('RGB')
+ image = self.transform_image(_image_rs)
+ return image
+
+class BiRefNetRemoveBackground:
+ def __init__(self):
+ self.ready = False
+
+ def load(self, weight_path, device):
+ # load model
+ self.model = BiRefNet()
+ state_dict = torch.load(weight_path, map_location='cpu')
+ unwanted_prefix = '_orig_mod.'
+ for k, v in list(state_dict.items()):
+ if k.startswith(unwanted_prefix):
+ state_dict[k[len(unwanted_prefix):]] = state_dict.pop(k)
+ self.model.load_state_dict(state_dict)
+ self.model = self.model.to(device)
+ self.model.eval()
+ # load processor
+ self.processor = BiRefNet_img_processor(Config())
+ self.ready = True
+
+
+ def generate_mask(self, image:Image) -> Image:
+
+ if torch.backends.mps.is_available():
+ device = "mps"
+ elif torch.cuda.is_available():
+ device = "cuda"
+ else:
+ device = "cpu"
+
+ if not self.ready:
+ model_folder_name = 'BiRefNet'
+ model_name = 'BiRefNet-ep480.pth'
+ model_file_path = ""
+ try:
+ model_file_path = os.path.join(
+ os.path.normpath(folder_paths.folder_names_and_paths[model_folder_name][0][0]), model_name)
+ except:
+ pass
+ if not os.path.exists(model_file_path):
+ model_file_path = os.path.join(folder_paths.models_dir, model_folder_name, model_name)
+ self.load(model_file_path, device=device)
+
+ i = pil2tensor(image)
+ orig_image = image.convert('RGB')
+ np_image = i.squeeze().numpy()
+ img = self.processor(np_image)
+ inputs = img[None, ...].to(device)
+ with torch.no_grad():
+ scaled_preds = self.model(inputs)[-1].sigmoid()
+ _mask = nn.functional.interpolate(scaled_preds[0].unsqueeze(0),
+ size=np_image.shape[:2],
+ mode='bilinear',
+ align_corners=True
+ )[0]
+
+ brightness_image = ImageEnhance.Brightness(tensor2pil(_mask))
+
+ return brightness_image.enhance(factor=1.01)
diff --git a/py/birefnet_ultra.py b/py/birefnet_ultra.py
new file mode 100644
index 0000000..8fbcfc3
--- /dev/null
+++ b/py/birefnet_ultra.py
@@ -0,0 +1,83 @@
+# layerstyle advance
+
+from .imagefunc import *
+
+NODE_NAME = 'BiRefNetUltra'
+
+class BiRefNetUltra:
+
+ @classmethod
+ def INPUT_TYPES(cls):
+
+ method_list = ['VITMatte', 'VITMatte(local)', 'PyMatting', 'GuidedFilter', ]
+ device_list = ['cuda','cpu']
+ return {
+ "required": {
+ "image": ("IMAGE",),
+ "detail_method": (method_list,),
+ "detail_erode": ("INT", {"default": 6, "min": 1, "max": 255, "step": 1}),
+ "detail_dilate": ("INT", {"default": 6, "min": 1, "max": 255, "step": 1}),
+ "black_point": ("FLOAT", {"default": 0.01, "min": 0.01, "max": 0.98, "step": 0.01, "display": "slider"}),
+ "white_point": ("FLOAT", {"default": 0.99, "min": 0.02, "max": 0.99, "step": 0.01, "display": "slider"}),
+ "process_detail": ("BOOLEAN", {"default": True}),
+ "device": (device_list,),
+ "max_megapixels": ("FLOAT", {"default": 2.0, "min": 1, "max": 999, "step": 0.1}),
+ },
+ "optional": {
+ }
+ }
+
+ RETURN_TYPES = ("IMAGE", "MASK", )
+ RETURN_NAMES = ("image", "mask", )
+ FUNCTION = "birefnet_ultra"
+ CATEGORY = '😺dzNodes/LayerMask'
+
+ def birefnet_ultra(self, image, detail_method, detail_erode, detail_dilate,
+ black_point, white_point, process_detail, device, max_megapixels):
+ ret_images = []
+ ret_masks = []
+
+ if detail_method == 'VITMatte(local)':
+ local_files_only = True
+ else:
+ local_files_only = False
+
+ from .birefnet_legacy import BiRefNetRemoveBackground
+ birefnetrmbg = BiRefNetRemoveBackground()
+
+ for i in image:
+ i = torch.unsqueeze(i, 0)
+ orig_image = tensor2pil(i).convert('RGB')
+
+ _mask = birefnetrmbg.generate_mask(orig_image)
+ _mask = image2mask(_mask)
+
+ detail_range = detail_erode + detail_dilate
+
+ if process_detail:
+ if detail_method == 'GuidedFilter':
+ _mask = guided_filter_alpha(i, _mask, detail_range // 6 + 1)
+ _mask = tensor2pil(histogram_remap(_mask, black_point, white_point))
+ elif detail_method == 'PyMatting':
+ _mask = tensor2pil(mask_edge_detail(i, _mask, detail_range // 8 + 1, black_point, white_point))
+ else:
+ _trimap = generate_VITMatte_trimap(_mask, detail_erode, detail_dilate)
+ _mask = generate_VITMatte(orig_image, _trimap, local_files_only=local_files_only, device=device, max_megapixels=max_megapixels)
+ _mask = tensor2pil(histogram_remap(pil2tensor(_mask), black_point, white_point))
+ else:
+ _mask = tensor2pil(_mask)
+
+ ret_image = RGB2RGBA(orig_image, _mask.convert('L'))
+ ret_images.append(pil2tensor(ret_image))
+ ret_masks.append(image2mask(_mask))
+
+ log(f"{NODE_NAME} Processed {len(ret_masks)} image(s).", message_type='finish')
+ return (torch.cat(ret_images, dim=0), torch.cat(ret_masks, dim=0),)
+
+NODE_CLASS_MAPPINGS = {
+ "LayerMask: BiRefNetUltra": BiRefNetUltra,
+}
+
+NODE_DISPLAY_NAME_MAPPINGS = {
+ "LayerMask: BiRefNetUltra": "LayerMask: BiRefNetUltra(Advance)",
+}
diff --git a/py/birefnet_ultra_v2.py b/py/birefnet_ultra_v2.py
new file mode 100644
index 0000000..596ebc4
--- /dev/null
+++ b/py/birefnet_ultra_v2.py
@@ -0,0 +1,215 @@
+# layerstyle advance
+
+import os
+import sys
+import torch
+from torchvision import transforms
+from transformers import AutoModelForImageSegmentation
+import tqdm
+from .imagefunc import *
+from comfy.utils import ProgressBar
+sys.path.append(os.path.join(os.path.dirname(__file__), 'BiRefNet_v2'))
+
+
+def get_models():
+ model_path = os.path.join(folder_paths.models_dir, 'BiRefNet', 'pth')
+ model_ext = [".pth"]
+ model_dict = get_files(model_path, model_ext)
+ return model_dict
+
+class LS_LoadBiRefNetModel:
+
+ def __init__(self):
+ self.birefnet = None
+ self.state_dict = None
+
+
+ @classmethod
+ def INPUT_TYPES(s):
+ tmp_list = list(get_models().keys())
+ model_list = []
+ if 'BiRefNet-general-epoch_244.pth' in tmp_list:
+ model_list.append('BiRefNet-general-epoch_244.pth')
+ tmp_list.remove('BiRefNet-general-epoch_244.pth')
+ model_list.extend(tmp_list)
+
+ return {
+ "required": {
+ "model": (model_list,),
+ },
+ }
+
+ RETURN_TYPES = ("BIREFNET_MODEL",)
+ RETURN_NAMES = ("birefnet_model",)
+ FUNCTION = "load_birefnet_model"
+ CATEGORY = '😺dzNodes/LayerMask'
+
+ def load_birefnet_model(self, model):
+ from .BiRefNet_v2.models.birefnet import BiRefNet
+ from .BiRefNet_v2.utils import check_state_dict
+ model_dict = get_models()
+ self.birefnet = BiRefNet(bb_pretrained=False)
+ self.state_dict = torch.load(model_dict[model], map_location='cpu', weights_only=True)
+ self.state_dict = check_state_dict(self.state_dict)
+ self.birefnet.load_state_dict(self.state_dict)
+ return (self.birefnet,)
+
+class LS_LoadBiRefNetModelV2:
+ def __init__(self):
+ self.model = None
+
+ @classmethod
+ def INPUT_TYPES(s):
+ model_list = list(s.birefnet_model_repos.keys())
+ return {
+ "required": {
+ "version": (model_list,{"default": model_list[0]}),
+ },
+ }
+
+ RETURN_TYPES = ("BIREFNET_MODEL",)
+ RETURN_NAMES = ("birefnet_model",)
+ FUNCTION = "load_birefnet_model"
+ CATEGORY = '😺dzNodes/LayerMask'
+
+ birefnet_model_repos = {
+ "BiRefNet-General": "ZhengPeng7/BiRefNet",
+ "RMBG-2.0": "briaai/RMBG-2.0"
+ }
+
+ def load_birefnet_model(self, version):
+ birefnet_path = os.path.join(folder_paths.models_dir, 'BiRefNet')
+ os.makedirs(birefnet_path, exist_ok=True)
+
+ model_path = os.path.join(birefnet_path, version)
+
+ if version == "BiRefNet-General":
+ old_birefnet_path = os.path.join(birefnet_path, 'pth')
+ old_model = "BiRefNet-general-epoch_244.pth"
+ old_model_path = os.path.join(old_birefnet_path, old_model)
+ if os.path.exists(old_model_path):
+ from .BiRefNet_v2.models.birefnet import BiRefNet
+ from .BiRefNet_v2.utils import check_state_dict
+ self.birefnet = BiRefNet(bb_pretrained=False)
+ self.state_dict = torch.load(old_model_path, map_location='cpu', weights_only=True)
+ self.state_dict = check_state_dict(self.state_dict)
+ self.birefnet.load_state_dict(self.state_dict)
+ return (self.birefnet,)
+ elif not os.path.exists(model_path):
+ log(f"Downloading {version} model...")
+ repo_id = self.birefnet_model_repos[version]
+ from huggingface_hub import snapshot_download
+ snapshot_download(repo_id=repo_id, local_dir=model_path, ignore_patterns=["*.md", "*.txt"])
+
+ self.model = AutoModelForImageSegmentation.from_pretrained(model_path, trust_remote_code=True)
+ return (self.model,)
+
+class LS_BiRefNetUltraV2:
+
+ def __init__(self):
+ self.NODE_NAME = 'BiRefNetUltraV2'
+
+ @classmethod
+ def INPUT_TYPES(cls):
+
+ method_list = ['VITMatte', 'VITMatte(local)', 'PyMatting', 'GuidedFilter', ]
+ device_list = ['cuda', 'cpu']
+ return {
+ "required": {
+ "image": ("IMAGE",),
+ "birefnet_model": ("BIREFNET_MODEL",),
+ "detail_method": (method_list,),
+ "detail_erode": ("INT", {"default": 4, "min": 1, "max": 255, "step": 1}),
+ "detail_dilate": ("INT", {"default": 2, "min": 1, "max": 255, "step": 1}),
+ "black_point": ("FLOAT", {"default": 0.01, "min": 0.01, "max": 0.98, "step": 0.01, "display": "slider"}),
+ "white_point": ("FLOAT", {"default": 0.99, "min": 0.02, "max": 0.99, "step": 0.01, "display": "slider"}),
+ "process_detail": ("BOOLEAN", {"default": False}),
+ "device": (device_list,),
+ "max_megapixels": ("FLOAT", {"default": 2.0, "min": 1, "max": 999, "step": 0.1}),
+ },
+ "optional": {
+ }
+ }
+
+ RETURN_TYPES = ("IMAGE", "MASK", )
+ RETURN_NAMES = ("image", "mask", )
+ FUNCTION = "birefnet_ultra_v2"
+ CATEGORY = '😺dzNodes/LayerMask'
+
+ def birefnet_ultra_v2(self, image, birefnet_model, detail_method, detail_erode, detail_dilate,
+ black_point, white_point, process_detail, device, max_megapixels):
+ ret_images = []
+ ret_masks = []
+ inference_image_size = (1024, 1024)
+ if detail_method == 'VITMatte(local)':
+ local_files_only = True
+ else:
+ local_files_only = False
+
+ torch.set_float32_matmul_precision(['high', 'highest'][0])
+ birefnet_model.to(device)
+ birefnet_model.eval()
+
+ comfy_pbar = ProgressBar(len(image))
+ tqdm_pbar = tqdm(total=len(image), desc="Processing BiRefNet")
+ for i in image:
+ i = torch.unsqueeze(i, 0)
+ orig_image = tensor2pil(i).convert('RGB')
+
+ transform_image = transforms.Compose([
+ transforms.Resize(inference_image_size),
+ transforms.ToTensor(),
+ transforms.Normalize([0.485, 0.456, 0.406], [0.229, 0.224, 0.225])
+ ])
+
+ inference_image = transform_image(orig_image).unsqueeze(0).to(device)
+
+ # Prediction
+ with torch.no_grad():
+ preds = birefnet_model(inference_image)[-1].sigmoid().cpu()
+ pred = preds[0].squeeze()
+ pred_pil = transforms.ToPILImage()(pred)
+ _mask = pred_pil.resize(inference_image_size)
+
+ resize_sampler = Image.BILINEAR
+ _mask = _mask.resize(orig_image.size, resize_sampler)
+ brightness_image = ImageEnhance.Brightness(_mask)
+ _mask = brightness_image.enhance(factor=1.08)
+ _mask = image2mask(_mask)
+
+ detail_range = detail_erode + detail_dilate
+
+ if process_detail:
+ if detail_method == 'GuidedFilter':
+ _mask = guided_filter_alpha(i, _mask, detail_range // 6 + 1)
+ _mask = tensor2pil(histogram_remap(_mask, black_point, white_point))
+ elif detail_method == 'PyMatting':
+ _mask = tensor2pil(mask_edge_detail(i, _mask, detail_range // 8 + 1, black_point, white_point))
+ else:
+ _trimap = generate_VITMatte_trimap(_mask, detail_erode, detail_dilate)
+ _mask = generate_VITMatte(orig_image, _trimap, local_files_only=local_files_only, device=device, max_megapixels=max_megapixels)
+ _mask = tensor2pil(histogram_remap(pil2tensor(_mask), black_point, white_point))
+ else:
+ _mask = tensor2pil(_mask)
+
+ ret_image = RGB2RGBA(orig_image, _mask.convert('L'))
+ ret_images.append(pil2tensor(ret_image))
+ ret_masks.append(image2mask(_mask))
+
+ comfy_pbar.update(1)
+ tqdm_pbar.update(1)
+
+ log(f"{self.NODE_NAME} Processed {len(ret_masks)} image(s).", message_type='finish')
+ return (torch.cat(ret_images, dim=0), torch.cat(ret_masks, dim=0),)
+
+NODE_CLASS_MAPPINGS = {
+ "LayerMask: BiRefNetUltraV2": LS_BiRefNetUltraV2,
+ "LayerMask: LoadBiRefNetModel": LS_LoadBiRefNetModel,
+ "LayerMask: LoadBiRefNetModelV2": LS_LoadBiRefNetModelV2
+}
+
+NODE_DISPLAY_NAME_MAPPINGS = {
+ "LayerMask: BiRefNetUltraV2": "LayerMask: BiRefNet Ultra V2(Advance)",
+ "LayerMask: LoadBiRefNetModel": "LayerMask: Load BiRefNet Model(Advance)",
+ "LayerMask: LoadBiRefNetModelV2": "LayerMask: Load BiRefNet Model V2(Advance)"
+}
diff --git a/py/blendmodes.py b/py/blendmodes.py
new file mode 100644
index 0000000..fac31af
--- /dev/null
+++ b/py/blendmodes.py
@@ -0,0 +1,324 @@
+"""
+author: Chris Freilich
+description: This extension provides a blend modes node with 30 blend modes.
+"""
+from PIL import Image
+import numpy as np
+import torch
+import torch.nn.functional as F
+from colorsys import rgb_to_hsv
+from blend_modes import difference, normal, screen, soft_light, lighten_only, dodge, \
+ addition, darken_only, multiply, hard_light, \
+ grain_extract, grain_merge, divide, overlay
+
+def dissolve(backdrop, source, opacity):
+ # Normalize the RGB and alpha values to 0-1
+ backdrop_norm = backdrop[:, :, :3] / 255
+ source_norm = source[:, :, :3] / 255
+ source_alpha_norm = source[:, :, 3] / 255
+
+ # Calculate the transparency of each pixel in the source image
+ transparency = opacity * source_alpha_norm
+
+ # Generate a random matrix with the same shape as the source image
+ random_matrix = np.random.random(source.shape[:2])
+
+ # Create a mask where the random values are less than the transparency
+ mask = random_matrix < transparency
+
+ # Use the mask to select pixels from the source or backdrop
+ blend = np.where(mask[..., None], source_norm, backdrop_norm)
+
+ # Apply the alpha channel of the source image to the blended image
+ new_rgb = (1 - source_alpha_norm[..., None]) * backdrop_norm + source_alpha_norm[..., None] * blend
+
+ # Ensure the RGB values are within the valid range
+ new_rgb = np.clip(new_rgb, 0, 1)
+
+ # Convert the RGB values back to 0-255
+ new_rgb = new_rgb * 255
+
+ # Calculate the new alpha value by taking the maximum of the backdrop and source alpha channels
+ new_alpha = np.maximum(backdrop[:, :, 3], source[:, :, 3])
+
+ # Create a new RGBA image with the calculated RGB and alpha values
+ result = np.dstack((new_rgb, new_alpha))
+
+ return result
+
+def rgb_to_hsv_via_torch(rgb_numpy: np.ndarray, device=None) -> torch.Tensor:
+ """
+ Convert an RGB image to HSV.
+
+ :param rgb: A tensor of shape (3, H, W) where the three channels correspond to R, G, B.
+ The values should be in the range [0, 1].
+ :return: A tensor of shape (3, H, W) where the three channels correspond to H, S, V.
+ The hue (H) will be in the range [0, 1], while S and V will be in the range [0, 1].
+ """
+ if device is None:
+ device = torch.device("cuda" if torch.cuda.is_available() else "cpu")
+
+ rgb = torch.from_numpy(rgb_numpy).float().permute(2, 0, 1).to(device)
+ r, g, b = rgb[0], rgb[1], rgb[2]
+
+ max_val, _ = torch.max(rgb, dim=0)
+ min_val, _ = torch.min(rgb, dim=0)
+ delta = max_val - min_val
+
+ h = torch.zeros_like(max_val)
+ s = torch.zeros_like(max_val)
+ v = max_val
+
+ # calc hue... avoid div by zero (by masking the delta)
+ mask = delta != 0
+ r_eq_max = (r == max_val) & mask
+ g_eq_max = (g == max_val) & mask
+ b_eq_max = (b == max_val) & mask
+
+ h[r_eq_max] = (g[r_eq_max] - b[r_eq_max]) / delta[r_eq_max] % 6
+ h[g_eq_max] = (b[g_eq_max] - r[g_eq_max]) / delta[g_eq_max] + 2.0
+ h[b_eq_max] = (r[b_eq_max] - g[b_eq_max]) / delta[b_eq_max] + 4.0
+
+ h = (h / 6.0) % 1.0
+
+ # calc saturation
+ s[max_val != 0] = delta[max_val != 0] / max_val[max_val != 0]
+
+ hsv = torch.stack([h, s, v], dim=0)
+
+ hsv_numpy = hsv.permute(1, 2, 0).cpu().numpy()
+ return hsv_numpy
+
+def hsv_to_rgb_via_torch(hsv_numpy: np.ndarray, device=None) -> torch.Tensor:
+ """
+ Convert an HSV image to RGB.
+
+ :param hsv: A tensor of shape (3, H, W) where the three channels correspond to H, S, V.
+ The H channel values should be in the range [0, 1], while S and V will be in the range [0, 1].
+ :return: A tensor of shape (3, H, W) where the three channels correspond to R, G, B.
+ The RGB values will be in the range [0, 1].
+ """
+ if device is None:
+ device = torch.device("cuda" if torch.cuda.is_available() else "cpu")
+
+ hsv = torch.from_numpy(hsv_numpy).float().permute(2, 0, 1).to(device)
+ h, s, v = hsv[0], hsv[1], hsv[2]
+
+ c = v * s # chroma
+ x = c * (1 - torch.abs((h * 6) % 2 - 1))
+ m = v - c # match value
+
+ z = torch.zeros_like(h)
+ rgb = torch.zeros_like(hsv)
+
+ # define conditions for different hue ranges
+ h_cond = [
+ (h < 1/6, torch.stack([c, x, z], dim=0)),
+ ((1/6 <= h) & (h < 2/6), torch.stack([x, c, z], dim=0)),
+ ((2/6 <= h) & (h < 3/6), torch.stack([z, c, x], dim=0)),
+ ((3/6 <= h) & (h < 4/6), torch.stack([z, x, c], dim=0)),
+ ((4/6 <= h) & (h < 5/6), torch.stack([x, z, c], dim=0)),
+ (h >= 5/6, torch.stack([c, z, x], dim=0)),
+ ]
+
+ # conditionally set RGB values based on the hue range
+ for cond, result in h_cond:
+ rgb[:, cond] = result[:, cond]
+
+ # add match value to convert to final RGB values
+ rgb = rgb + m
+
+ rgb_numpy = rgb.permute(1, 2, 0).cpu().numpy()
+ return rgb_numpy
+
+def hsv(backdrop, source, opacity, channel):
+
+ # Convert RGBA to RGB, normalized
+ backdrop_rgb = backdrop[:, :, :3] / 255.0
+ source_rgb = source[:, :, :3] / 255.0
+ source_alpha = source[:, :, 3] / 255.0
+
+ # Convert RGB to HSV
+ backdrop_hsv = rgb_to_hsv_via_torch(backdrop_rgb)
+ source_hsv = rgb_to_hsv_via_torch(source_rgb)
+
+ # Combine HSV values
+ new_hsv = backdrop_hsv.copy()
+
+ # Determine which channel to operate on
+ if channel == "saturation":
+ new_hsv[:, :, 1] = (1 - opacity * source_alpha) * backdrop_hsv[:, :, 1] + opacity * source_alpha * source_hsv[:, :, 1]
+ elif channel == "luminance":
+ new_hsv[:, :, 2] = (1 - opacity * source_alpha) * backdrop_hsv[:, :, 2] + opacity * source_alpha * source_hsv[:, :, 2]
+ elif channel == "hue":
+ new_hsv[:, :, 0] = (1 - opacity * source_alpha) * backdrop_hsv[:, :, 0] + opacity * source_alpha * source_hsv[:, :, 0]
+ elif channel == "color":
+ new_hsv[:, :, :2] = (1 - opacity * source_alpha[..., None]) * backdrop_hsv[:, :, :2] + opacity * source_alpha[..., None] * source_hsv[:, :, :2]
+
+ # Convert HSV back to RGB
+ new_rgb = hsv_to_rgb_via_torch(new_hsv)
+
+ # Apply the alpha channel of the source image to the new RGB image
+ new_rgb = (1 - source_alpha[..., None]) * backdrop_rgb + source_alpha[..., None] * new_rgb
+
+ # Ensure the RGB values are within the valid range
+ new_rgb = np.clip(new_rgb, 0, 1)
+
+ # Convert RGB back to RGBA and scale to 0-255 range
+ new_rgba = np.dstack((new_rgb * 255, backdrop[:, :, 3]))
+
+ return new_rgba.astype(np.uint8)
+
+def saturation(backdrop, source, opacity):
+ return hsv(backdrop, source, opacity, "saturation")
+
+def luminance(backdrop, source, opacity):
+ return hsv(backdrop, source, opacity, "luminance")
+
+def hue(backdrop, source, opacity):
+ return hsv(backdrop, source, opacity, "hue")
+
+def color(backdrop, source, opacity):
+ return hsv(backdrop, source, opacity, "color")
+
+def darker_lighter_color(backdrop, source, opacity, type):
+
+ # Normalize the RGB and alpha values to 0-1
+ backdrop_norm = backdrop[:, :, :3] / 255
+ source_norm = source[:, :, :3] / 255
+ source_alpha_norm = source[:, :, 3] / 255
+
+ # Convert RGB to HSV
+ backdrop_hsv = np.array([rgb_to_hsv(*rgb) for row in backdrop_norm for rgb in row]).reshape(backdrop.shape[:2] + (3,))
+ source_hsv = np.array([rgb_to_hsv(*rgb) for row in source_norm for rgb in row]).reshape(source.shape[:2] + (3,))
+
+ # Create a mask where the value (brightness) of the source image is less than the value of the backdrop image
+ if type == "dark":
+ mask = source_hsv[:, :, 2] < backdrop_hsv[:, :, 2]
+ else:
+ mask = source_hsv[:, :, 2] > backdrop_hsv[:, :, 2]
+
+ # Use the mask to select pixels from the source or backdrop
+ blend = np.where(mask[..., None], source_norm, backdrop_norm)
+
+ # Apply the alpha channel of the source image to the blended image
+ new_rgb = (1 - source_alpha_norm[..., None] * opacity) * backdrop_norm + source_alpha_norm[..., None] * opacity * blend
+
+ # Ensure the RGB values are within the valid range
+ new_rgb = np.clip(new_rgb, 0, 1)
+
+ # Convert the RGB values back to 0-255
+ new_rgb = new_rgb * 255
+
+ # Calculate the new alpha value by taking the maximum of the backdrop and source alpha channels
+ new_alpha = np.maximum(backdrop[:, :, 3], source[:, :, 3])
+
+ # Create a new RGBA image with the calculated RGB and alpha values
+ result = np.dstack((new_rgb, new_alpha))
+
+ return result
+
+def darker_color(backdrop, source, opacity):
+ return darker_lighter_color(backdrop, source, opacity, "dark")
+
+def lighter_color(backdrop, source, opacity):
+ return darker_lighter_color(backdrop, source, opacity, "light")
+
+def simple_mode(backdrop, source, opacity, mode):
+ # Normalize the RGB and alpha values to 0-1
+ backdrop_norm = backdrop[:, :, :3] / 255
+ source_norm = source[:, :, :3] / 255
+ source_alpha_norm = source[:, :, 3:4] / 255
+
+ # Calculate the blend without any transparency considerations
+ if mode == "linear_burn":
+ blend = backdrop_norm + source_norm - 1
+ elif mode == "linear_light":
+ blend = backdrop_norm + (2 * source_norm) - 1
+ elif mode == "color_dodge":
+ blend = backdrop_norm / (1 - source_norm)
+ blend = np.clip(blend, 0, 1)
+ elif mode == "color_burn":
+ blend = 1 - ((1 - backdrop_norm) / source_norm)
+ blend = np.clip(blend, 0, 1)
+ elif mode == "exclusion":
+ blend = backdrop_norm + source_norm - (2 * backdrop_norm * source_norm)
+ elif mode == "subtract":
+ blend = backdrop_norm - source_norm
+ elif mode == "vivid_light":
+ blend = np.where(source_norm <= 0.5, backdrop_norm / (1 - 2 * source_norm), 1 - (1 -backdrop_norm) / (2 * source_norm - 0.5) )
+ blend = np.clip(blend, 0, 1)
+ elif mode == "pin_light":
+ blend = np.where(source_norm <= 0.5, np.minimum(backdrop_norm, 2 * source_norm), np.maximum(backdrop_norm, 2 * (source_norm - 0.5)))
+ elif mode == "hard_mix":
+ blend = simple_mode(backdrop, source, opacity, "linear_light")
+ blend = np.round(blend[:, :, :3] / 255)
+
+ # Apply the blended layer back onto the backdrop layer while utilizing the alpha channel and opacity information
+ new_rgb = (1 - source_alpha_norm * opacity) * backdrop_norm + source_alpha_norm * opacity * blend
+
+ # Ensure the RGB values are within the valid range
+ new_rgb = np.clip(new_rgb, 0, 1)
+
+ # Convert the RGB values back to 0-255
+ new_rgb = new_rgb * 255
+
+ # Calculate the new alpha value by taking the maximum of the backdrop and source alpha channels
+ new_alpha = np.maximum(backdrop[:, :, 3], source[:, :, 3])
+
+ # Create a new RGBA image with the calculated RGB and alpha values
+ result = np.dstack((new_rgb, new_alpha))
+
+ return result
+
+def linear_light(backdrop, source, opacity):
+ return simple_mode(backdrop, source, opacity, "linear_light")
+def vivid_light(backdrop, source, opacity):
+ return simple_mode(backdrop, source, opacity, "vivid_light")
+def pin_light(backdrop, source, opacity):
+ return simple_mode(backdrop, source, opacity, "pin_light")
+def hard_mix(backdrop, source, opacity):
+ return simple_mode(backdrop, source, opacity, "hard_mix")
+def linear_burn(backdrop, source, opacity):
+ return simple_mode(backdrop, source, opacity, "linear_burn")
+def color_dodge(backdrop, source, opacity):
+ return simple_mode(backdrop, source, opacity, "color_dodge")
+def color_burn(backdrop, source, opacity):
+ return simple_mode(backdrop, source, opacity, "color_burn")
+def exclusion(backdrop, source, opacity):
+ return simple_mode(backdrop, source, opacity, "exclusion")
+def subtract(backdrop, source, opacity):
+ return simple_mode(backdrop, source, opacity, "subtract")
+
+BLEND_MODES = {
+ "normal": normal,
+ "dissolve": dissolve,
+ "darken": darken_only,
+ "multiply": multiply,
+ "color burn": color_burn,
+ "linear burn": linear_burn,
+ "darker color": darker_color,
+ "lighten": lighten_only,
+ "screen": screen,
+ "color dodge": color_dodge,
+ "linear dodge(add)": addition,
+ "lighter color": lighter_color,
+ "dodge": dodge,
+ "overlay": overlay,
+ "soft light": soft_light,
+ "hard light": hard_light,
+ "vivid light": vivid_light,
+ "linear light": linear_light,
+ "pin light": pin_light,
+ "hard mix": hard_mix,
+ "difference": difference,
+ "exclusion": exclusion,
+ "subtract": subtract,
+ "divide": divide,
+ "hue": hue,
+ "saturation": saturation,
+ "color": color,
+ "luminosity": luminance,
+ "grain extract": grain_extract,
+ "grain merge": grain_merge
+}
diff --git a/py/briarmbg.py b/py/briarmbg.py
new file mode 100644
index 0000000..647bdc0
--- /dev/null
+++ b/py/briarmbg.py
@@ -0,0 +1,455 @@
+import torch
+import torch.nn as nn
+import torch.nn.functional as F
+
+class REBNCONV(nn.Module):
+ def __init__(self,in_ch=3,out_ch=3,dirate=1,stride=1):
+ super(REBNCONV,self).__init__()
+
+ self.conv_s1 = nn.Conv2d(in_ch,out_ch,3,padding=1*dirate,dilation=1*dirate,stride=stride)
+ self.bn_s1 = nn.BatchNorm2d(out_ch)
+ self.relu_s1 = nn.ReLU(inplace=True)
+
+ def forward(self,x):
+
+ hx = x
+ xout = self.relu_s1(self.bn_s1(self.conv_s1(hx)))
+
+ return xout
+
+## upsample tensor 'src' to have the same spatial size with tensor 'tar'
+def _upsample_like(src,tar):
+
+ src = F.interpolate(src,size=tar.shape[2:],mode='bilinear')
+
+ return src
+
+
+### RSU-7 ###
+class RSU7(nn.Module):
+
+ def __init__(self, in_ch=3, mid_ch=12, out_ch=3, img_size=512):
+ super(RSU7,self).__init__()
+
+ self.in_ch = in_ch
+ self.mid_ch = mid_ch
+ self.out_ch = out_ch
+
+ self.rebnconvin = REBNCONV(in_ch,out_ch,dirate=1) ## 1 -> 1/2
+
+ self.rebnconv1 = REBNCONV(out_ch,mid_ch,dirate=1)
+ self.pool1 = nn.MaxPool2d(2,stride=2,ceil_mode=True)
+
+ self.rebnconv2 = REBNCONV(mid_ch,mid_ch,dirate=1)
+ self.pool2 = nn.MaxPool2d(2,stride=2,ceil_mode=True)
+
+ self.rebnconv3 = REBNCONV(mid_ch,mid_ch,dirate=1)
+ self.pool3 = nn.MaxPool2d(2,stride=2,ceil_mode=True)
+
+ self.rebnconv4 = REBNCONV(mid_ch,mid_ch,dirate=1)
+ self.pool4 = nn.MaxPool2d(2,stride=2,ceil_mode=True)
+
+ self.rebnconv5 = REBNCONV(mid_ch,mid_ch,dirate=1)
+ self.pool5 = nn.MaxPool2d(2,stride=2,ceil_mode=True)
+
+ self.rebnconv6 = REBNCONV(mid_ch,mid_ch,dirate=1)
+
+ self.rebnconv7 = REBNCONV(mid_ch,mid_ch,dirate=2)
+
+ self.rebnconv6d = REBNCONV(mid_ch*2,mid_ch,dirate=1)
+ self.rebnconv5d = REBNCONV(mid_ch*2,mid_ch,dirate=1)
+ self.rebnconv4d = REBNCONV(mid_ch*2,mid_ch,dirate=1)
+ self.rebnconv3d = REBNCONV(mid_ch*2,mid_ch,dirate=1)
+ self.rebnconv2d = REBNCONV(mid_ch*2,mid_ch,dirate=1)
+ self.rebnconv1d = REBNCONV(mid_ch*2,out_ch,dirate=1)
+
+ def forward(self,x):
+ b, c, h, w = x.shape
+
+ hx = x
+ hxin = self.rebnconvin(hx)
+
+ hx1 = self.rebnconv1(hxin)
+ hx = self.pool1(hx1)
+
+ hx2 = self.rebnconv2(hx)
+ hx = self.pool2(hx2)
+
+ hx3 = self.rebnconv3(hx)
+ hx = self.pool3(hx3)
+
+ hx4 = self.rebnconv4(hx)
+ hx = self.pool4(hx4)
+
+ hx5 = self.rebnconv5(hx)
+ hx = self.pool5(hx5)
+
+ hx6 = self.rebnconv6(hx)
+
+ hx7 = self.rebnconv7(hx6)
+
+ hx6d = self.rebnconv6d(torch.cat((hx7,hx6),1))
+ hx6dup = _upsample_like(hx6d,hx5)
+
+ hx5d = self.rebnconv5d(torch.cat((hx6dup,hx5),1))
+ hx5dup = _upsample_like(hx5d,hx4)
+
+ hx4d = self.rebnconv4d(torch.cat((hx5dup,hx4),1))
+ hx4dup = _upsample_like(hx4d,hx3)
+
+ hx3d = self.rebnconv3d(torch.cat((hx4dup,hx3),1))
+ hx3dup = _upsample_like(hx3d,hx2)
+
+ hx2d = self.rebnconv2d(torch.cat((hx3dup,hx2),1))
+ hx2dup = _upsample_like(hx2d,hx1)
+
+ hx1d = self.rebnconv1d(torch.cat((hx2dup,hx1),1))
+
+ return hx1d + hxin
+
+
+### RSU-6 ###
+class RSU6(nn.Module):
+
+ def __init__(self, in_ch=3, mid_ch=12, out_ch=3):
+ super(RSU6,self).__init__()
+
+ self.rebnconvin = REBNCONV(in_ch,out_ch,dirate=1)
+
+ self.rebnconv1 = REBNCONV(out_ch,mid_ch,dirate=1)
+ self.pool1 = nn.MaxPool2d(2,stride=2,ceil_mode=True)
+
+ self.rebnconv2 = REBNCONV(mid_ch,mid_ch,dirate=1)
+ self.pool2 = nn.MaxPool2d(2,stride=2,ceil_mode=True)
+
+ self.rebnconv3 = REBNCONV(mid_ch,mid_ch,dirate=1)
+ self.pool3 = nn.MaxPool2d(2,stride=2,ceil_mode=True)
+
+ self.rebnconv4 = REBNCONV(mid_ch,mid_ch,dirate=1)
+ self.pool4 = nn.MaxPool2d(2,stride=2,ceil_mode=True)
+
+ self.rebnconv5 = REBNCONV(mid_ch,mid_ch,dirate=1)
+
+ self.rebnconv6 = REBNCONV(mid_ch,mid_ch,dirate=2)
+
+ self.rebnconv5d = REBNCONV(mid_ch*2,mid_ch,dirate=1)
+ self.rebnconv4d = REBNCONV(mid_ch*2,mid_ch,dirate=1)
+ self.rebnconv3d = REBNCONV(mid_ch*2,mid_ch,dirate=1)
+ self.rebnconv2d = REBNCONV(mid_ch*2,mid_ch,dirate=1)
+ self.rebnconv1d = REBNCONV(mid_ch*2,out_ch,dirate=1)
+
+ def forward(self,x):
+
+ hx = x
+
+ hxin = self.rebnconvin(hx)
+
+ hx1 = self.rebnconv1(hxin)
+ hx = self.pool1(hx1)
+
+ hx2 = self.rebnconv2(hx)
+ hx = self.pool2(hx2)
+
+ hx3 = self.rebnconv3(hx)
+ hx = self.pool3(hx3)
+
+ hx4 = self.rebnconv4(hx)
+ hx = self.pool4(hx4)
+
+ hx5 = self.rebnconv5(hx)
+
+ hx6 = self.rebnconv6(hx5)
+
+
+ hx5d = self.rebnconv5d(torch.cat((hx6,hx5),1))
+ hx5dup = _upsample_like(hx5d,hx4)
+
+ hx4d = self.rebnconv4d(torch.cat((hx5dup,hx4),1))
+ hx4dup = _upsample_like(hx4d,hx3)
+
+ hx3d = self.rebnconv3d(torch.cat((hx4dup,hx3),1))
+ hx3dup = _upsample_like(hx3d,hx2)
+
+ hx2d = self.rebnconv2d(torch.cat((hx3dup,hx2),1))
+ hx2dup = _upsample_like(hx2d,hx1)
+
+ hx1d = self.rebnconv1d(torch.cat((hx2dup,hx1),1))
+
+ return hx1d + hxin
+
+### RSU-5 ###
+class RSU5(nn.Module):
+
+ def __init__(self, in_ch=3, mid_ch=12, out_ch=3):
+ super(RSU5,self).__init__()
+
+ self.rebnconvin = REBNCONV(in_ch,out_ch,dirate=1)
+
+ self.rebnconv1 = REBNCONV(out_ch,mid_ch,dirate=1)
+ self.pool1 = nn.MaxPool2d(2,stride=2,ceil_mode=True)
+
+ self.rebnconv2 = REBNCONV(mid_ch,mid_ch,dirate=1)
+ self.pool2 = nn.MaxPool2d(2,stride=2,ceil_mode=True)
+
+ self.rebnconv3 = REBNCONV(mid_ch,mid_ch,dirate=1)
+ self.pool3 = nn.MaxPool2d(2,stride=2,ceil_mode=True)
+
+ self.rebnconv4 = REBNCONV(mid_ch,mid_ch,dirate=1)
+
+ self.rebnconv5 = REBNCONV(mid_ch,mid_ch,dirate=2)
+
+ self.rebnconv4d = REBNCONV(mid_ch*2,mid_ch,dirate=1)
+ self.rebnconv3d = REBNCONV(mid_ch*2,mid_ch,dirate=1)
+ self.rebnconv2d = REBNCONV(mid_ch*2,mid_ch,dirate=1)
+ self.rebnconv1d = REBNCONV(mid_ch*2,out_ch,dirate=1)
+
+ def forward(self,x):
+
+ hx = x
+
+ hxin = self.rebnconvin(hx)
+
+ hx1 = self.rebnconv1(hxin)
+ hx = self.pool1(hx1)
+
+ hx2 = self.rebnconv2(hx)
+ hx = self.pool2(hx2)
+
+ hx3 = self.rebnconv3(hx)
+ hx = self.pool3(hx3)
+
+ hx4 = self.rebnconv4(hx)
+
+ hx5 = self.rebnconv5(hx4)
+
+ hx4d = self.rebnconv4d(torch.cat((hx5,hx4),1))
+ hx4dup = _upsample_like(hx4d,hx3)
+
+ hx3d = self.rebnconv3d(torch.cat((hx4dup,hx3),1))
+ hx3dup = _upsample_like(hx3d,hx2)
+
+ hx2d = self.rebnconv2d(torch.cat((hx3dup,hx2),1))
+ hx2dup = _upsample_like(hx2d,hx1)
+
+ hx1d = self.rebnconv1d(torch.cat((hx2dup,hx1),1))
+
+ return hx1d + hxin
+
+### RSU-4 ###
+class RSU4(nn.Module):
+
+ def __init__(self, in_ch=3, mid_ch=12, out_ch=3):
+ super(RSU4,self).__init__()
+
+ self.rebnconvin = REBNCONV(in_ch,out_ch,dirate=1)
+
+ self.rebnconv1 = REBNCONV(out_ch,mid_ch,dirate=1)
+ self.pool1 = nn.MaxPool2d(2,stride=2,ceil_mode=True)
+
+ self.rebnconv2 = REBNCONV(mid_ch,mid_ch,dirate=1)
+ self.pool2 = nn.MaxPool2d(2,stride=2,ceil_mode=True)
+
+ self.rebnconv3 = REBNCONV(mid_ch,mid_ch,dirate=1)
+
+ self.rebnconv4 = REBNCONV(mid_ch,mid_ch,dirate=2)
+
+ self.rebnconv3d = REBNCONV(mid_ch*2,mid_ch,dirate=1)
+ self.rebnconv2d = REBNCONV(mid_ch*2,mid_ch,dirate=1)
+ self.rebnconv1d = REBNCONV(mid_ch*2,out_ch,dirate=1)
+
+ def forward(self,x):
+
+ hx = x
+
+ hxin = self.rebnconvin(hx)
+
+ hx1 = self.rebnconv1(hxin)
+ hx = self.pool1(hx1)
+
+ hx2 = self.rebnconv2(hx)
+ hx = self.pool2(hx2)
+
+ hx3 = self.rebnconv3(hx)
+
+ hx4 = self.rebnconv4(hx3)
+
+ hx3d = self.rebnconv3d(torch.cat((hx4,hx3),1))
+ hx3dup = _upsample_like(hx3d,hx2)
+
+ hx2d = self.rebnconv2d(torch.cat((hx3dup,hx2),1))
+ hx2dup = _upsample_like(hx2d,hx1)
+
+ hx1d = self.rebnconv1d(torch.cat((hx2dup,hx1),1))
+
+ return hx1d + hxin
+
+### RSU-4F ###
+class RSU4F(nn.Module):
+
+ def __init__(self, in_ch=3, mid_ch=12, out_ch=3):
+ super(RSU4F,self).__init__()
+
+ self.rebnconvin = REBNCONV(in_ch,out_ch,dirate=1)
+
+ self.rebnconv1 = REBNCONV(out_ch,mid_ch,dirate=1)
+ self.rebnconv2 = REBNCONV(mid_ch,mid_ch,dirate=2)
+ self.rebnconv3 = REBNCONV(mid_ch,mid_ch,dirate=4)
+
+ self.rebnconv4 = REBNCONV(mid_ch,mid_ch,dirate=8)
+
+ self.rebnconv3d = REBNCONV(mid_ch*2,mid_ch,dirate=4)
+ self.rebnconv2d = REBNCONV(mid_ch*2,mid_ch,dirate=2)
+ self.rebnconv1d = REBNCONV(mid_ch*2,out_ch,dirate=1)
+
+ def forward(self,x):
+
+ hx = x
+
+ hxin = self.rebnconvin(hx)
+
+ hx1 = self.rebnconv1(hxin)
+ hx2 = self.rebnconv2(hx1)
+ hx3 = self.rebnconv3(hx2)
+
+ hx4 = self.rebnconv4(hx3)
+
+ hx3d = self.rebnconv3d(torch.cat((hx4,hx3),1))
+ hx2d = self.rebnconv2d(torch.cat((hx3d,hx2),1))
+ hx1d = self.rebnconv1d(torch.cat((hx2d,hx1),1))
+
+ return hx1d + hxin
+
+
+class myrebnconv(nn.Module):
+ def __init__(self, in_ch=3,
+ out_ch=1,
+ kernel_size=3,
+ stride=1,
+ padding=1,
+ dilation=1,
+ groups=1):
+ super(myrebnconv,self).__init__()
+
+ self.conv = nn.Conv2d(in_ch,
+ out_ch,
+ kernel_size=kernel_size,
+ stride=stride,
+ padding=padding,
+ dilation=dilation,
+ groups=groups)
+ self.bn = nn.BatchNorm2d(out_ch)
+ self.rl = nn.ReLU(inplace=True)
+
+ def forward(self,x):
+ return self.rl(self.bn(self.conv(x)))
+
+
+class BriaRMBG(nn.Module):
+
+ def __init__(self,in_ch=3,out_ch=1):
+ super(BriaRMBG,self).__init__()
+
+ self.conv_in = nn.Conv2d(in_ch,64,3,stride=2,padding=1)
+ self.pool_in = nn.MaxPool2d(2,stride=2,ceil_mode=True)
+
+ self.stage1 = RSU7(64,32,64)
+ self.pool12 = nn.MaxPool2d(2,stride=2,ceil_mode=True)
+
+ self.stage2 = RSU6(64,32,128)
+ self.pool23 = nn.MaxPool2d(2,stride=2,ceil_mode=True)
+
+ self.stage3 = RSU5(128,64,256)
+ self.pool34 = nn.MaxPool2d(2,stride=2,ceil_mode=True)
+
+ self.stage4 = RSU4(256,128,512)
+ self.pool45 = nn.MaxPool2d(2,stride=2,ceil_mode=True)
+
+ self.stage5 = RSU4F(512,256,512)
+ self.pool56 = nn.MaxPool2d(2,stride=2,ceil_mode=True)
+
+ self.stage6 = RSU4F(512,256,512)
+
+ # decoder
+ self.stage5d = RSU4F(1024,256,512)
+ self.stage4d = RSU4(1024,128,256)
+ self.stage3d = RSU5(512,64,128)
+ self.stage2d = RSU6(256,32,64)
+ self.stage1d = RSU7(128,16,64)
+
+ self.side1 = nn.Conv2d(64,out_ch,3,padding=1)
+ self.side2 = nn.Conv2d(64,out_ch,3,padding=1)
+ self.side3 = nn.Conv2d(128,out_ch,3,padding=1)
+ self.side4 = nn.Conv2d(256,out_ch,3,padding=1)
+ self.side5 = nn.Conv2d(512,out_ch,3,padding=1)
+ self.side6 = nn.Conv2d(512,out_ch,3,padding=1)
+
+ # self.outconv = nn.Conv2d(6*out_ch,out_ch,1)
+
+ def forward(self,x):
+
+ hx = x
+
+ hxin = self.conv_in(hx)
+ #hx = self.pool_in(hxin)
+
+ #stage 1
+ hx1 = self.stage1(hxin)
+ hx = self.pool12(hx1)
+
+ #stage 2
+ hx2 = self.stage2(hx)
+ hx = self.pool23(hx2)
+
+ #stage 3
+ hx3 = self.stage3(hx)
+ hx = self.pool34(hx3)
+
+ #stage 4
+ hx4 = self.stage4(hx)
+ hx = self.pool45(hx4)
+
+ #stage 5
+ hx5 = self.stage5(hx)
+ hx = self.pool56(hx5)
+
+ #stage 6
+ hx6 = self.stage6(hx)
+ hx6up = _upsample_like(hx6,hx5)
+
+ #-------------------- decoder --------------------
+ hx5d = self.stage5d(torch.cat((hx6up,hx5),1))
+ hx5dup = _upsample_like(hx5d,hx4)
+
+ hx4d = self.stage4d(torch.cat((hx5dup,hx4),1))
+ hx4dup = _upsample_like(hx4d,hx3)
+
+ hx3d = self.stage3d(torch.cat((hx4dup,hx3),1))
+ hx3dup = _upsample_like(hx3d,hx2)
+
+ hx2d = self.stage2d(torch.cat((hx3dup,hx2),1))
+ hx2dup = _upsample_like(hx2d,hx1)
+
+ hx1d = self.stage1d(torch.cat((hx2dup,hx1),1))
+
+
+ #side output
+ d1 = self.side1(hx1d)
+ d1 = _upsample_like(d1,x)
+
+ d2 = self.side2(hx2d)
+ d2 = _upsample_like(d2,x)
+
+ d3 = self.side3(hx3d)
+ d3 = _upsample_like(d3,x)
+
+ d4 = self.side4(hx4d)
+ d4 = _upsample_like(d4,x)
+
+ d5 = self.side5(hx5d)
+ d5 = _upsample_like(d5,x)
+
+ d6 = self.side6(hx6)
+ d6 = _upsample_like(d6,x)
+
+ return [F.sigmoid(d1), F.sigmoid(d2), F.sigmoid(d3), F.sigmoid(d4), F.sigmoid(d5), F.sigmoid(d6)],[hx1d,hx2d,hx3d,hx4d,hx5d,hx6]
+
diff --git a/py/evf_sam/__init__.py b/py/evf_sam/__init__.py
new file mode 100644
index 0000000..e69de29
diff --git a/py/evf_sam/evf_sam_inference.py b/py/evf_sam/evf_sam_inference.py
new file mode 100644
index 0000000..9571fa4
--- /dev/null
+++ b/py/evf_sam/evf_sam_inference.py
@@ -0,0 +1,146 @@
+import os
+import sys
+from PIL import Image
+import cv2
+import numpy as np
+import torch
+import torch.nn.functional as F
+from torchvision import transforms
+from torchvision.transforms.functional import InterpolationMode
+from transformers import AutoTokenizer, BitsAndBytesConfig
+sys.path.insert(0, os.path.abspath(os.path.dirname(__file__)))
+from .model.segment_anything.utils.transforms import ResizeLongestSide
+
+def sam_preprocess(
+ x: np.ndarray,
+ pixel_mean=torch.Tensor([123.675, 116.28, 103.53]).view(-1, 1, 1),
+ pixel_std=torch.Tensor([58.395, 57.12, 57.375]).view(-1, 1, 1),
+ img_size=1024,
+ model_type="ori") -> torch.Tensor:
+ '''
+ preprocess of Segment Anything Model, including scaling, normalization and padding.
+ preprocess differs between SAM and Effi-SAM, where Effi-SAM use no padding.
+ input: ndarray
+ output: torch.Tensor
+ '''
+ assert img_size==1024, \
+ "both SAM and Effi-SAM receive images of size 1024^2, don't change this setting unless you're sure that your employed model works well with another size."
+ x = ResizeLongestSide(img_size).apply_image(x)
+ resize_shape = x.shape[:2]
+ x = torch.from_numpy(x).permute(2,0,1).contiguous()
+
+ # Normalize colors
+ x = (x - pixel_mean) / pixel_std
+ if model_type=="effi" or model_type=="sam2":
+ x = F.interpolate(x.unsqueeze(0), (img_size, img_size), mode="bilinear").squeeze(0)
+ else:
+ # Pad
+ h, w = x.shape[-2:]
+ padh = img_size - h
+ padw = img_size - w
+ x = F.pad(x, (0, padw, 0, padh))
+ return x, resize_shape
+
+def beit3_preprocess(x: np.ndarray, img_size=224) -> torch.Tensor:
+ '''
+ preprocess for BEIT-3 model.
+ input: ndarray
+ output: torch.Tensor
+ '''
+ beit_preprocess = transforms.Compose([
+ transforms.ToTensor(),
+ transforms.Resize((img_size, img_size), interpolation=InterpolationMode.BICUBIC),
+ transforms.Normalize(mean=(0.5, 0.5, 0.5), std=(0.5, 0.5, 0.5))
+ ])
+ return beit_preprocess(x)
+
+def init_models(model_path:str, model_type:str, precision:str, load_in_bit:int=16):
+ tokenizer = AutoTokenizer.from_pretrained(
+ model_path,
+ padding_side="right",
+ use_fast=False,
+ )
+
+ torch_dtype = torch.float32
+ if precision == "bf16":
+ torch_dtype = torch.bfloat16
+ elif precision == "fp16":
+ torch_dtype = torch.half
+
+ kwargs = {"torch_dtype": torch_dtype}
+
+ if load_in_bit==4:
+ kwargs.update(
+ {
+ "torch_dtype": torch.half,
+ "quantization_config": BitsAndBytesConfig(
+ llm_int8_skip_modules=["visual_model"],
+ load_in_4bit=True,
+ bnb_4bit_compute_dtype=torch.float16,
+ bnb_4bit_use_double_quant=True,
+ bnb_4bit_quant_type="nf4",
+ ),
+ }
+ )
+ elif load_in_bit==8:
+ kwargs.update(
+ {
+ "torch_dtype": torch.half,
+ "quantization_config": BitsAndBytesConfig(
+ llm_int8_skip_modules=["visual_model"],
+ load_in_8bit=True,
+ ),
+ }
+ )
+
+ if model_type=="ori":
+ from model.evf_sam import EvfSamModel
+ model = EvfSamModel.from_pretrained(
+ model_path, low_cpu_mem_usage=True, **kwargs
+ )
+ elif model_type=="effi":
+ from model.evf_effisam import EvfEffiSamModel
+ model = EvfEffiSamModel.from_pretrained(
+ model_path, low_cpu_mem_usage=True, **kwargs
+ )
+ elif model_type=="sam2":
+ from model.evf_sam2 import EvfSam2Model
+ model = EvfSam2Model.from_pretrained(
+ model_path, low_cpu_mem_usage=True, **kwargs
+ )
+
+ if load_in_bit > 8 and torch.cuda.is_available():
+ model = model.cuda()
+ model.eval()
+
+ return tokenizer, model
+
+def evf_sam_main(model_path:str, model_type:str, precision:str, load_in_bit:int, image:Image, prompt:str, ):
+
+ image_size = 224
+ # initialize model and tokenizer
+ tokenizer, model = init_models(model_path, model_type, precision, load_in_bit)
+
+ # preprocess
+ image_np = np.asarray(image)
+ image_np = cv2.cvtColor(image_np, cv2.COLOR_BGR2RGB)
+ original_size_list = [image_np.shape[:2]]
+ image_beit = beit3_preprocess(image_np, image_size).to(dtype=model.dtype, device=model.device)
+ image_sam, resize_shape = sam_preprocess(image_np, model_type=model_type)
+ image_sam = image_sam.to(dtype=model.dtype, device=model.device)
+ input_ids = tokenizer(prompt, return_tensors="pt")["input_ids"].to(device=model.device)
+
+ # infer
+ pred_mask = model.inference(
+ image_sam.unsqueeze(0),
+ image_beit.unsqueeze(0),
+ input_ids,
+ resize_list=[resize_shape],
+ original_size_list=original_size_list,
+ )
+
+ pred_mask = pred_mask.detach().cpu().numpy()[0]
+ pred_mask = (pred_mask > 0).astype(np.uint8) * 255
+ out_put_image = Image.fromarray(pred_mask.squeeze(), mode="L")
+
+ return out_put_image
\ No newline at end of file
diff --git a/py/evf_sam/model/EfficientSAM/efficient_sam/__init__.py b/py/evf_sam/model/EfficientSAM/efficient_sam/__init__.py
new file mode 100644
index 0000000..22a2d29
--- /dev/null
+++ b/py/evf_sam/model/EfficientSAM/efficient_sam/__init__.py
@@ -0,0 +1,7 @@
+# Copyright (c) Meta Platforms, Inc. and affiliates.
+# All rights reserved.
+
+from .build_efficient_sam import (
+ build_efficient_sam_vitt,
+ build_efficient_sam_vits,
+)
diff --git a/py/evf_sam/model/EfficientSAM/efficient_sam/build_efficient_sam.py b/py/evf_sam/model/EfficientSAM/efficient_sam/build_efficient_sam.py
new file mode 100644
index 0000000..b5a030d
--- /dev/null
+++ b/py/evf_sam/model/EfficientSAM/efficient_sam/build_efficient_sam.py
@@ -0,0 +1,22 @@
+# Copyright (c) Meta Platforms, Inc. and affiliates.
+# All rights reserved.
+
+# This source code is licensed under the license found in the
+# LICENSE file in the root directory of this source tree.
+
+from .efficient_sam import build_efficient_sam
+
+def build_efficient_sam_vitt(checkpoint=None):
+ return build_efficient_sam(
+ encoder_patch_embed_dim=192,
+ encoder_num_heads=3,
+ checkpoint=checkpoint,
+ ).eval()
+
+
+def build_efficient_sam_vits(checkpoint=None):
+ return build_efficient_sam(
+ encoder_patch_embed_dim=384,
+ encoder_num_heads=6,
+ checkpoint=checkpoint,
+ ).eval()
diff --git a/py/evf_sam/model/EfficientSAM/efficient_sam/efficient_sam.py b/py/evf_sam/model/EfficientSAM/efficient_sam/efficient_sam.py
new file mode 100644
index 0000000..a4ad17d
--- /dev/null
+++ b/py/evf_sam/model/EfficientSAM/efficient_sam/efficient_sam.py
@@ -0,0 +1,306 @@
+# Copyright (c) Meta Platforms, Inc. and affiliates.
+# All rights reserved.
+
+# This source code is licensed under the license found in the
+# LICENSE file in the root directory of this source tree.
+
+import math
+from typing import Any, List, Tuple, Type
+
+import torch
+import torch.nn.functional as F
+
+from torch import nn, Tensor
+
+from .efficient_sam_decoder import MaskDecoder, PromptEncoder
+from .efficient_sam_encoder import ImageEncoderViT
+from .two_way_transformer import TwoWayAttentionBlock, TwoWayTransformer
+
+class EfficientSam(nn.Module):
+ mask_threshold: float = 0.0
+ image_format: str = "RGB"
+
+ def __init__(
+ self,
+ image_encoder: ImageEncoderViT,
+ prompt_encoder: PromptEncoder,
+ decoder_max_num_input_points: int,
+ mask_decoder: MaskDecoder,
+ pixel_mean: List[float] = [0.485, 0.456, 0.406],
+ pixel_std: List[float] = [0.229, 0.224, 0.225],
+ ) -> None:
+ """
+ SAM predicts object masks from an image and input prompts.
+
+ Arguments:
+ image_encoder (ImageEncoderViT): The backbone used to encode the
+ image into image embeddings that allow for efficient mask prediction.
+ prompt_encoder (PromptEncoder): Encodes various types of input prompts.
+ mask_decoder (MaskDecoder): Predicts masks from the image embeddings
+ and encoded prompts.
+ pixel_mean (list(float)): Mean values for normalizing pixels in the input image.
+ pixel_std (list(float)): Std values for normalizing pixels in the input image.
+ """
+ super().__init__()
+ self.image_encoder = image_encoder
+ self.prompt_encoder = prompt_encoder
+ self.decoder_max_num_input_points = decoder_max_num_input_points
+ self.mask_decoder = mask_decoder
+ self.register_buffer(
+ "pixel_mean", torch.Tensor(pixel_mean).view(1, 3, 1, 1), False
+ )
+ self.register_buffer(
+ "pixel_std", torch.Tensor(pixel_std).view(1, 3, 1, 1), False
+ )
+
+ @torch.jit.export
+ def predict_masks(
+ self,
+ image_embeddings: torch.Tensor,
+ batched_points: torch.Tensor,
+ batched_point_labels: torch.Tensor,
+ multimask_output: bool,
+ input_h: int,
+ input_w: int,
+ output_h: int = -1,
+ output_w: int = -1,
+ ) -> Tuple[torch.Tensor, torch.Tensor]:
+ """
+ Predicts masks given image embeddings and prompts. This only runs the decoder.
+
+ Arguments:
+ image_embeddings: A tensor of shape [B, C, H, W] or [B*max_num_queries, C, H, W]
+ batched_points: A tensor of shape [B, max_num_queries, num_pts, 2]
+ batched_point_labels: A tensor of shape [B, max_num_queries, num_pts]
+ Returns:
+ A tuple of two tensors:
+ low_res_mask: A tensor of shape [B, max_num_queries, 256, 256] of predicted masks
+ iou_predictions: A tensor of shape [B, max_num_queries] of estimated IOU scores
+ """
+
+ batch_size, max_num_queries, num_pts, _ = batched_points.shape
+ num_pts = batched_points.shape[2]
+ rescaled_batched_points = self.get_rescaled_pts(batched_points, input_h, input_w)
+
+ if num_pts > self.decoder_max_num_input_points:
+ rescaled_batched_points = rescaled_batched_points[
+ :, :, : self.decoder_max_num_input_points, :
+ ]
+ batched_point_labels = batched_point_labels[
+ :, :, : self.decoder_max_num_input_points
+ ]
+ elif num_pts < self.decoder_max_num_input_points:
+ rescaled_batched_points = F.pad(
+ rescaled_batched_points,
+ (0, 0, 0, self.decoder_max_num_input_points - num_pts),
+ value=-1.0,
+ )
+ batched_point_labels = F.pad(
+ batched_point_labels,
+ (0, self.decoder_max_num_input_points - num_pts),
+ value=-1.0,
+ )
+
+ sparse_embeddings = self.prompt_encoder(
+ rescaled_batched_points.reshape(
+ batch_size * max_num_queries, self.decoder_max_num_input_points, 2
+ ),
+ batched_point_labels.reshape(
+ batch_size * max_num_queries, self.decoder_max_num_input_points
+ ),
+ )
+
+ sparse_embeddings = sparse_embeddings.view(
+ batch_size,
+ max_num_queries,
+ sparse_embeddings.shape[1],
+ sparse_embeddings.shape[2],
+ )
+ low_res_masks, iou_predictions = self.mask_decoder(
+ image_embeddings,
+ self.prompt_encoder.get_dense_pe(),
+ sparse_prompt_embeddings=sparse_embeddings,
+ multimask_output=multimask_output,
+ )
+ _, num_predictions, low_res_size, _ = low_res_masks.shape
+
+ if output_w > 0 and output_h > 0:
+ output_masks = F.interpolate(
+ low_res_masks, (output_h, output_w), mode="bicubic"
+ )
+ output_masks = torch.reshape(
+ output_masks,
+ (batch_size, max_num_queries, num_predictions, output_h, output_w),
+ )
+ else:
+ output_masks = torch.reshape(
+ low_res_masks,
+ (
+ batch_size,
+ max_num_queries,
+ num_predictions,
+ low_res_size,
+ low_res_size,
+ ),
+ )
+ iou_predictions = torch.reshape(
+ iou_predictions, (batch_size, max_num_queries, num_predictions)
+ )
+ return output_masks, iou_predictions
+
+ def get_rescaled_pts(self, batched_points: torch.Tensor, input_h: int, input_w: int):
+ return torch.stack(
+ [
+ torch.where(
+ batched_points[..., 0] >= 0,
+ batched_points[..., 0] * self.image_encoder.img_size / input_w,
+ -1.0,
+ ),
+ torch.where(
+ batched_points[..., 1] >= 0,
+ batched_points[..., 1] * self.image_encoder.img_size / input_h,
+ -1.0,
+ ),
+ ],
+ dim=-1,
+ )
+
+ @torch.jit.export
+ def get_image_embeddings(self, batched_images) -> torch.Tensor:
+ """
+ Predicts masks end-to-end from provided images and prompts.
+ If prompts are not known in advance, using SamPredictor is
+ recommended over calling the model directly.
+
+ Arguments:
+ batched_images: A tensor of shape [B, 3, H, W]
+ Returns:
+ List of image embeddings each of of shape [B, C(i), H(i), W(i)].
+ The last embedding corresponds to the final layer.
+ """
+ batched_images = self.preprocess(batched_images)
+ return self.image_encoder(batched_images)
+
+ def forward(
+ self,
+ batched_images: torch.Tensor,
+ batched_points: torch.Tensor,
+ batched_point_labels: torch.Tensor,
+ scale_to_original_image_size: bool = True,
+ ) -> Tuple[torch.Tensor, torch.Tensor]:
+ """
+ Predicts masks end-to-end from provided images and prompts.
+ If prompts are not known in advance, using SamPredictor is
+ recommended over calling the model directly.
+
+ Arguments:
+ batched_images: A tensor of shape [B, 3, H, W]
+ batched_points: A tensor of shape [B, num_queries, max_num_pts, 2]
+ batched_point_labels: A tensor of shape [B, num_queries, max_num_pts]
+
+ Returns:
+ A list tuples of two tensors where the ith element is by considering the first i+1 points.
+ low_res_mask: A tensor of shape [B, 256, 256] of predicted masks
+ iou_predictions: A tensor of shape [B, max_num_queries] of estimated IOU scores
+ """
+ batch_size, _, input_h, input_w = batched_images.shape
+ image_embeddings = self.get_image_embeddings(batched_images)
+ return self.predict_masks(
+ image_embeddings,
+ batched_points,
+ batched_point_labels,
+ multimask_output=True,
+ input_h=input_h,
+ input_w=input_w,
+ output_h=input_h if scale_to_original_image_size else -1,
+ output_w=input_w if scale_to_original_image_size else -1,
+ )
+
+ def preprocess(self, x: torch.Tensor) -> torch.Tensor:
+ """Normalize pixel values and pad to a square input."""
+ if (
+ x.shape[2] != self.image_encoder.img_size
+ or x.shape[3] != self.image_encoder.img_size
+ ):
+ x = F.interpolate(
+ x,
+ (self.image_encoder.img_size, self.image_encoder.img_size),
+ mode="bilinear",
+ )
+ return (x - self.pixel_mean) / self.pixel_std
+
+
+def build_efficient_sam(encoder_patch_embed_dim, encoder_num_heads, checkpoint=None):
+ img_size = 1024
+ encoder_patch_size = 16
+ encoder_depth = 12
+ encoder_mlp_ratio = 4.0
+ encoder_neck_dims = [256, 256]
+ decoder_max_num_input_points = 6
+ decoder_transformer_depth = 2
+ decoder_transformer_mlp_dim = 2048
+ decoder_num_heads = 8
+ decoder_upscaling_layer_dims = [64, 32]
+ num_multimask_outputs = 3
+ iou_head_depth = 3
+ iou_head_hidden_dim = 256
+ activation = "gelu"
+ normalization_type = "layer_norm"
+ normalize_before_activation = False
+
+ assert activation == "relu" or activation == "gelu"
+ if activation == "relu":
+ activation_fn = nn.ReLU
+ else:
+ activation_fn = nn.GELU
+
+ image_encoder = ImageEncoderViT(
+ img_size=img_size,
+ patch_size=encoder_patch_size,
+ in_chans=3,
+ patch_embed_dim=encoder_patch_embed_dim,
+ normalization_type=normalization_type,
+ depth=encoder_depth,
+ num_heads=encoder_num_heads,
+ mlp_ratio=encoder_mlp_ratio,
+ neck_dims=encoder_neck_dims,
+ act_layer=activation_fn,
+ )
+
+ image_embedding_size = image_encoder.image_embedding_size
+ encoder_transformer_output_dim = image_encoder.transformer_output_dim
+
+ sam = EfficientSam(
+ image_encoder=image_encoder,
+ prompt_encoder=PromptEncoder(
+ embed_dim=encoder_transformer_output_dim,
+ image_embedding_size=(image_embedding_size, image_embedding_size),
+ input_image_size=(img_size, img_size),
+ ),
+ decoder_max_num_input_points=decoder_max_num_input_points,
+ mask_decoder=MaskDecoder(
+ transformer_dim=encoder_transformer_output_dim,
+ transformer=TwoWayTransformer(
+ depth=decoder_transformer_depth,
+ embedding_dim=encoder_transformer_output_dim,
+ num_heads=decoder_num_heads,
+ mlp_dim=decoder_transformer_mlp_dim,
+ activation=activation_fn,
+ normalize_before_activation=normalize_before_activation,
+ ),
+ num_multimask_outputs=num_multimask_outputs,
+ activation=activation_fn,
+ normalization_type=normalization_type,
+ normalize_before_activation=normalize_before_activation,
+ iou_head_depth=iou_head_depth - 1,
+ iou_head_hidden_dim=iou_head_hidden_dim,
+ upscaling_layer_dims=decoder_upscaling_layer_dims,
+ ),
+ pixel_mean=[0.485, 0.456, 0.406],
+ pixel_std=[0.229, 0.224, 0.225],
+ )
+ if checkpoint is not None:
+ with open(checkpoint, "rb") as f:
+ state_dict = torch.load(f, map_location="cpu")
+ sam.load_state_dict(state_dict["model"])
+ return sam
diff --git a/py/evf_sam/model/EfficientSAM/efficient_sam/efficient_sam_decoder.py b/py/evf_sam/model/EfficientSAM/efficient_sam/efficient_sam_decoder.py
new file mode 100644
index 0000000..909605d
--- /dev/null
+++ b/py/evf_sam/model/EfficientSAM/efficient_sam/efficient_sam_decoder.py
@@ -0,0 +1,318 @@
+# Copyright (c) Meta Platforms, Inc. and affiliates.
+# All rights reserved.
+
+# This source code is licensed under the license found in the
+# LICENSE file in the root directory of this source tree.
+
+from typing import List, Tuple, Type
+
+import numpy as np
+import torch
+import torch.nn as nn
+import torch.nn.functional as F
+
+from .mlp import MLPBlock
+
+
+class PromptEncoder(nn.Module):
+ def __init__(
+ self,
+ embed_dim: int,
+ image_embedding_size: Tuple[int, int],
+ input_image_size: Tuple[int, int],
+ ) -> None:
+ """
+ Encodes prompts for input to SAM's mask decoder.
+
+ Arguments:
+ embed_dim (int): The prompts' embedding dimension
+ image_embedding_size (tuple(int, int)): The spatial size of the
+ image embedding, as (H, W).
+ input_image_size (int): The padded size of the image as input
+ to the image encoder, as (H, W).
+ """
+ super().__init__()
+ self.embed_dim = embed_dim
+ self.input_image_size = input_image_size
+ self.image_embedding_size = image_embedding_size
+ self.pe_layer = PositionEmbeddingRandom(embed_dim // 2)
+ self.invalid_points = nn.Embedding(1, embed_dim)
+ self.point_embeddings = nn.Embedding(1, embed_dim)
+ self.bbox_top_left_embeddings = nn.Embedding(1, embed_dim)
+ self.bbox_bottom_right_embeddings = nn.Embedding(1, embed_dim)
+
+ def get_dense_pe(self) -> torch.Tensor:
+ """
+ Returns the positional encoding used to encode point prompts,
+ applied to a dense set of points the shape of the image encoding.
+
+ Returns:
+ torch.Tensor: Positional encoding with shape
+ 1x(embed_dim)x(embedding_h)x(embedding_w)
+ """
+ return self.pe_layer(self.image_embedding_size).unsqueeze(0)
+
+ def _embed_points(
+ self,
+ points: torch.Tensor,
+ labels: torch.Tensor,
+ ) -> torch.Tensor:
+ """Embeds point prompts."""
+
+ points = points + 0.5 # Shift to center of pixel
+ point_embedding = self.pe_layer.forward_with_coords(
+ points, self.input_image_size
+ )
+ invalid_label_ids = torch.eq(labels, -1)[:,:,None]
+ point_label_ids = torch.eq(labels, 1)[:,:,None]
+ topleft_label_ids = torch.eq(labels, 2)[:,:,None]
+ bottomright_label_ids = torch.eq(labels, 3)[:,:,None]
+ point_embedding = point_embedding + self.invalid_points.weight[:,None,:] * invalid_label_ids
+ point_embedding = point_embedding + self.point_embeddings.weight[:,None,:] * point_label_ids
+ point_embedding = point_embedding + self.bbox_top_left_embeddings.weight[:,None,:] * topleft_label_ids
+ point_embedding = point_embedding + self.bbox_bottom_right_embeddings.weight[:,None,:] * bottomright_label_ids
+ return point_embedding
+
+ def forward(
+ self,
+ coords,
+ labels,
+ ) -> torch.Tensor:
+ """
+ Embeds different types of prompts, returning both sparse and dense
+ embeddings.
+
+ Arguments:
+ points: A tensor of shape [B, 2]
+ labels: An integer tensor of shape [B] where each element is 1,2 or 3.
+
+ Returns:
+ torch.Tensor: sparse embeddings for the points and boxes, with shape
+ BxNx(embed_dim), where N is determined by the number of input points
+ and boxes.
+ """
+ return self._embed_points(coords, labels)
+
+
+class PositionEmbeddingRandom(nn.Module):
+ """
+ Positional encoding using random spatial frequencies.
+ """
+
+ def __init__(self, num_pos_feats: int) -> None:
+ super().__init__()
+ self.register_buffer(
+ "positional_encoding_gaussian_matrix", torch.randn((2, num_pos_feats))
+ )
+
+ def _pe_encoding(self, coords: torch.Tensor) -> torch.Tensor:
+ """Positionally encode points that are normalized to [0,1]."""
+ # assuming coords are in [0, 1]^2 square and have d_1 x ... x d_n x 2 shape
+ coords = 2 * coords - 1
+ coords = coords @ self.positional_encoding_gaussian_matrix
+ coords = 2 * np.pi * coords
+ # outputs d_1 x ... x d_n x C shape
+ return torch.cat([torch.sin(coords), torch.cos(coords)], dim=-1)
+
+ def forward(self, size: Tuple[int, int]) -> torch.Tensor:
+ """Generate positional encoding for a grid of the specified size."""
+ h, w = size
+ device = self.positional_encoding_gaussian_matrix.device
+ grid = torch.ones([h, w], device=device, dtype=self.positional_encoding_gaussian_matrix.dtype)
+ y_embed = grid.cumsum(dim=0) - 0.5
+ x_embed = grid.cumsum(dim=1) - 0.5
+ y_embed = y_embed / h
+ x_embed = x_embed / w
+
+ pe = self._pe_encoding(torch.stack([x_embed, y_embed], dim=-1))
+ return pe.permute(2, 0, 1) # C x H x W
+
+ def forward_with_coords(
+ self, coords_input: torch.Tensor, image_size: Tuple[int, int]
+ ) -> torch.Tensor:
+ """Positionally encode points that are not normalized to [0,1]."""
+ coords = coords_input.clone()
+ coords[:, :, 0] = coords[:, :, 0] / image_size[1]
+ coords[:, :, 1] = coords[:, :, 1] / image_size[0]
+ # remove to(float) here, don't know why original implementation add this
+ return self._pe_encoding(coords) # B x N x C
+
+
+class MaskDecoder(nn.Module):
+ def __init__(
+ self,
+ *,
+ transformer_dim: int,
+ transformer: nn.Module,
+ num_multimask_outputs: int,
+ activation: Type[nn.Module],
+ normalization_type: str,
+ normalize_before_activation: bool,
+ iou_head_depth: int,
+ iou_head_hidden_dim: int,
+ upscaling_layer_dims: List[int],
+ ) -> None:
+ """
+ Predicts masks given an image and prompt embeddings, using a
+ transformer architecture.
+
+ Arguments:
+ transformer_dim (int): the channel dimension of the transformer
+ transformer (nn.Module): the transformer used to predict masks
+ num_multimask_outputs (int): the number of masks to predict
+ when disambiguating masks
+ activation (nn.Module): the type of activation to use when
+ upscaling masks
+ iou_head_depth (int): the depth of the MLP used to predict
+ mask quality
+ iou_head_hidden_dim (int): the hidden dimension of the MLP
+ used to predict mask quality
+ """
+ super().__init__()
+ self.transformer_dim = transformer_dim
+ self.transformer = transformer
+
+ self.num_multimask_outputs = num_multimask_outputs
+
+ self.iou_token = nn.Embedding(1, transformer_dim)
+ if num_multimask_outputs > 1:
+ self.num_mask_tokens = num_multimask_outputs + 1
+ else:
+ self.num_mask_tokens = 1
+ self.mask_tokens = nn.Embedding(self.num_mask_tokens, transformer_dim)
+ output_dim_after_upscaling = transformer_dim
+
+ self.final_output_upscaling_layers = nn.ModuleList([])
+ for idx, layer_dims in enumerate(upscaling_layer_dims):
+ self.final_output_upscaling_layers.append(
+ nn.Sequential(
+ nn.ConvTranspose2d(
+ output_dim_after_upscaling,
+ layer_dims,
+ kernel_size=2,
+ stride=2,
+ ),
+ nn.GroupNorm(1, layer_dims)
+ if idx < len(upscaling_layer_dims) - 1
+ else nn.Identity(),
+ activation(),
+ )
+ )
+ output_dim_after_upscaling = layer_dims
+
+ self.output_hypernetworks_mlps = nn.ModuleList(
+ [
+ MLPBlock(
+ input_dim=transformer_dim,
+ hidden_dim=transformer_dim,
+ output_dim=output_dim_after_upscaling,
+ num_layers=2,
+ act=activation,
+ )
+ for i in range(self.num_mask_tokens)
+ ]
+ )
+
+ self.iou_prediction_head = MLPBlock(
+ input_dim=transformer_dim,
+ hidden_dim=iou_head_hidden_dim,
+ output_dim=self.num_mask_tokens,
+ num_layers=iou_head_depth,
+ act=activation,
+ )
+
+ def forward(
+ self,
+ image_embeddings: torch.Tensor,
+ image_pe: torch.Tensor,
+ sparse_prompt_embeddings: torch.Tensor,
+ multimask_output: bool,
+ ) -> Tuple[torch.Tensor, torch.Tensor]:
+ """
+ Predict masks given image and prompt embeddings.
+
+ Arguments:
+ image_embeddings: A tensor of shape [B, C, H, W] or [B*max_num_queries, C, H, W]
+ image_pe (torch.Tensor): positional encoding with the shape of image_embeddings (the batch dimension is broadcastable).
+ sparse_prompt_embeddings (torch.Tensor): the embeddings of the points and boxes
+ multimask_output (bool): Whether to return multiple masks or a single
+ mask.
+
+ Returns:
+ torch.Tensor: batched predicted masks
+ torch.Tensor: batched predictions of mask quality
+ """
+
+ (
+ batch_size,
+ max_num_queries,
+ sparse_embed_dim_1,
+ sparse_embed_dim_2,
+ ) = sparse_prompt_embeddings.shape
+
+ (
+ _,
+ image_embed_dim_c,
+ image_embed_dim_h,
+ image_embed_dim_w,
+ ) = image_embeddings.shape
+
+ # Tile the image embedding for all queries.
+ image_embeddings_tiled = torch.tile(
+ image_embeddings[:, None, :, :, :], [1, max_num_queries, 1, 1, 1]
+ ).view(
+ batch_size * max_num_queries,
+ image_embed_dim_c,
+ image_embed_dim_h,
+ image_embed_dim_w,
+ )
+ sparse_prompt_embeddings = sparse_prompt_embeddings.reshape(
+ batch_size * max_num_queries, sparse_embed_dim_1, sparse_embed_dim_2
+ )
+ masks, iou_pred = self.predict_masks(
+ image_embeddings=image_embeddings_tiled,
+ image_pe=image_pe,
+ sparse_prompt_embeddings=sparse_prompt_embeddings,
+ )
+
+ if multimask_output and self.num_multimask_outputs > 1:
+ return masks[:, 1:, :], iou_pred[:, 1:]
+ else:
+ return masks[:, :1, :], iou_pred[:, :1]
+
+ def predict_masks(
+ self,
+ image_embeddings: torch.Tensor,
+ image_pe: torch.Tensor,
+ sparse_prompt_embeddings: torch.Tensor,
+ ) -> Tuple[torch.Tensor, torch.Tensor]:
+ """Predicts masks. See 'forward' for more details."""
+ # Concatenate output tokens
+ output_tokens = torch.cat(
+ [self.iou_token.weight, self.mask_tokens.weight], dim=0
+ )
+ output_tokens = output_tokens.unsqueeze(0).expand(
+ sparse_prompt_embeddings.size(0), -1, -1
+ )
+ tokens = torch.cat((output_tokens, sparse_prompt_embeddings), dim=1)
+ # Expand per-image data in batch direction to be per-mask
+ pos_src = torch.repeat_interleave(image_pe, tokens.shape[0], dim=0)
+ b, c, h, w = image_embeddings.shape
+ hs, src = self.transformer(image_embeddings, pos_src, tokens)
+ iou_token_out = hs[:, 0, :]
+ mask_tokens_out = hs[:, 1 : (1 + self.num_mask_tokens), :]
+
+ # Upscale mask embeddings and predict masks using the mask tokens
+ upscaled_embedding = src.transpose(1, 2).view(b, c, h, w)
+
+ for upscaling_layer in self.final_output_upscaling_layers:
+ upscaled_embedding = upscaling_layer(upscaled_embedding)
+ hyper_in_list: List[torch.Tensor] = []
+ for i, output_hypernetworks_mlp in enumerate(self.output_hypernetworks_mlps):
+ hyper_in_list.append(output_hypernetworks_mlp(mask_tokens_out[:, i, :]))
+ hyper_in = torch.stack(hyper_in_list, dim=1)
+ b, c, h, w = upscaled_embedding.shape
+ masks = (hyper_in @ upscaled_embedding.view(b, c, h * w)).view(b, -1, h, w)
+ # Generate mask quality predictions
+ iou_pred = self.iou_prediction_head(iou_token_out)
+ return masks, iou_pred
diff --git a/py/evf_sam/model/EfficientSAM/efficient_sam/efficient_sam_encoder.py b/py/evf_sam/model/EfficientSAM/efficient_sam/efficient_sam_encoder.py
new file mode 100644
index 0000000..73fd7ac
--- /dev/null
+++ b/py/evf_sam/model/EfficientSAM/efficient_sam/efficient_sam_encoder.py
@@ -0,0 +1,257 @@
+# Copyright (c) Meta Platforms, Inc. and affiliates.
+# All rights reserved.
+
+# This source code is licensed under the license found in the
+# LICENSE file in the root directory of this source tree.
+
+import math
+from typing import List, Optional, Tuple, Type
+
+import torch
+import torch.nn as nn
+import torch.nn.functional as F
+
+
+class LayerNorm2d(nn.Module):
+ def __init__(self, num_channels: int, eps: float = 1e-6) -> None:
+ super().__init__()
+ self.weight = nn.Parameter(torch.ones(num_channels))
+ self.bias = nn.Parameter(torch.zeros(num_channels))
+ self.eps = eps
+
+ def forward(self, x: torch.Tensor) -> torch.Tensor:
+ u = x.mean(1, keepdim=True)
+ s = (x - u).pow(2).mean(1, keepdim=True)
+ x = (x - u) / torch.sqrt(s + self.eps)
+ x = self.weight[:, None, None] * x + self.bias[:, None, None]
+ return x
+
+
+class PatchEmbed(nn.Module):
+ """2D Image to Patch Embedding"""
+
+ def __init__(
+ self,
+ img_size,
+ patch_size,
+ in_chans,
+ embed_dim,
+ ):
+ super().__init__()
+ self.proj = nn.Conv2d(
+ in_chans,
+ embed_dim,
+ kernel_size=(patch_size, patch_size),
+ stride=(patch_size, patch_size),
+ bias=True,
+ )
+
+ def forward(self, x):
+ B, C, H, W = x.shape
+ x = self.proj(x)
+ return x
+
+
+class Attention(nn.Module):
+ def __init__(
+ self,
+ dim,
+ num_heads,
+ qkv_bias,
+ qk_scale=None,
+ ):
+ super().__init__()
+ self.num_heads = num_heads
+ head_dim = dim // num_heads
+ self.scale = qk_scale or head_dim**-0.5
+ self.qkv = nn.Linear(dim, dim * 3, bias=qkv_bias)
+ self.proj = nn.Linear(dim, dim)
+
+ def forward(self, x):
+ B, N, C = x.shape
+ qkv = (
+ self.qkv(x)
+ .reshape(B, N, 3, self.num_heads, C // self.num_heads)
+ .permute(2, 0, 3, 1, 4)
+ )
+ q, k, v = (
+ qkv[0],
+ qkv[1],
+ qkv[2],
+ )
+ attn = (q @ k.transpose(-2, -1)) * self.scale
+ attn = attn.softmax(dim=-1)
+ x = (attn @ v).transpose(1, 2).reshape(B, N, C)
+ x = self.proj(x)
+ return x
+
+
+class Mlp(nn.Module):
+ def __init__(
+ self,
+ in_features,
+ hidden_features=None,
+ out_features=None,
+ act_layer=nn.GELU,
+ ):
+ super().__init__()
+ out_features = out_features or in_features
+ hidden_features = hidden_features or in_features
+ self.fc1 = nn.Linear(in_features, hidden_features)
+ self.act = act_layer()
+ self.fc2 = nn.Linear(hidden_features, out_features)
+
+ def forward(self, x):
+ x = self.fc1(x)
+ x = self.act(x)
+ x = self.fc2(x)
+ return x
+
+
+class Block(nn.Module):
+ def __init__(
+ self,
+ dim,
+ num_heads,
+ mlp_ratio=4.0,
+ qkv_bias=False,
+ qk_scale=None,
+ act_layer=nn.GELU,
+ ):
+ super().__init__()
+ self.norm1 = nn.LayerNorm(dim, eps=1e-6)
+ self.attn = Attention(
+ dim,
+ num_heads=num_heads,
+ qkv_bias=qkv_bias,
+ qk_scale=qk_scale,
+ )
+ self.norm2 = nn.LayerNorm(dim, eps=1e-6)
+ mlp_hidden_dim = int(dim * mlp_ratio)
+ self.mlp = Mlp(
+ in_features=dim,
+ hidden_features=mlp_hidden_dim,
+ act_layer=act_layer,
+ )
+
+ def forward(self, x):
+ x = x + self.attn(self.norm1(x))
+ x = x + self.mlp(self.norm2(x))
+ return x
+
+
+@torch.jit.export
+def get_abs_pos(
+ abs_pos: torch.Tensor, has_cls_token: bool, hw: List[int]
+) -> torch.Tensor:
+ """
+ Calculate absolute positional embeddings. If needed, resize embeddings and remove cls_token
+ dimension for the original embeddings.
+ Args:
+ abs_pos (Tensor): absolute positional embeddings with (1, num_position, C).
+ has_cls_token (bool): If true, has 1 embedding in abs_pos for cls token.
+ hw (Tuple): size of input image tokens.
+
+ Returns:
+ Absolute positional embeddings after processing with shape (1, H, W, C)
+ """
+ h = hw[0]
+ w = hw[1]
+ if has_cls_token:
+ abs_pos = abs_pos[:, 1:]
+ xy_num = abs_pos.shape[1]
+ size = int(math.sqrt(xy_num))
+ assert size * size == xy_num
+
+ if size != h or size != w:
+ new_abs_pos = F.interpolate(
+ abs_pos.reshape(1, size, size, -1).permute(0, 3, 1, 2),
+ size=(h, w),
+ mode="bicubic",
+ align_corners=False,
+ )
+ return new_abs_pos.permute(0, 2, 3, 1)
+ else:
+ return abs_pos.reshape(1, h, w, -1)
+
+
+# Image encoder for efficient SAM.
+class ImageEncoderViT(nn.Module):
+ def __init__(
+ self,
+ img_size: int,
+ patch_size: int,
+ in_chans: int,
+ patch_embed_dim: int,
+ normalization_type: str,
+ depth: int,
+ num_heads: int,
+ mlp_ratio: float,
+ neck_dims: List[int],
+ act_layer: Type[nn.Module],
+ ) -> None:
+ """
+ Args:
+ img_size (int): Input image size.
+ patch_size (int): Patch size.
+ in_chans (int): Number of input image channels.
+ patch_embed_dim (int): Patch embedding dimension.
+ depth (int): Depth of ViT.
+ num_heads (int): Number of attention heads in each ViT block.
+ mlp_ratio (float): Ratio of mlp hidden dim to embedding dim.
+ act_layer (nn.Module): Activation layer.
+ """
+ super().__init__()
+
+ self.img_size = img_size
+ self.image_embedding_size = img_size // ((patch_size if patch_size > 0 else 1))
+ self.transformer_output_dim = ([patch_embed_dim] + neck_dims)[-1]
+ self.pretrain_use_cls_token = True
+ pretrain_img_size = 224
+ self.patch_embed = PatchEmbed(img_size, patch_size, in_chans, patch_embed_dim)
+ # Initialize absolute positional embedding with pretrain image size.
+ num_patches = (pretrain_img_size // patch_size) * (
+ pretrain_img_size // patch_size
+ )
+ num_positions = num_patches + 1
+ self.pos_embed = nn.Parameter(torch.zeros(1, num_positions, patch_embed_dim))
+ self.blocks = nn.ModuleList()
+ for i in range(depth):
+ vit_block = Block(patch_embed_dim, num_heads, mlp_ratio, True)
+ self.blocks.append(vit_block)
+ self.neck = nn.Sequential(
+ nn.Conv2d(
+ patch_embed_dim,
+ neck_dims[0],
+ kernel_size=1,
+ bias=False,
+ ),
+ LayerNorm2d(neck_dims[0]),
+ nn.Conv2d(
+ neck_dims[0],
+ neck_dims[0],
+ kernel_size=3,
+ padding=1,
+ bias=False,
+ ),
+ LayerNorm2d(neck_dims[0]),
+ )
+
+ def forward(self, x: torch.Tensor) -> torch.Tensor:
+ assert (
+ x.shape[2] == self.img_size and x.shape[3] == self.img_size
+ ), "input image size must match self.img_size"
+ x = self.patch_embed(x)
+ # B C H W -> B H W C
+ x = x.permute(0, 2, 3, 1)
+ x = x + get_abs_pos(
+ self.pos_embed, self.pretrain_use_cls_token, [x.shape[1], x.shape[2]]
+ )
+ num_patches = x.shape[1]
+ assert x.shape[2] == num_patches
+ x = x.reshape(x.shape[0], num_patches * num_patches, x.shape[3])
+ for blk in self.blocks:
+ x = blk(x)
+ x = x.reshape(x.shape[0], num_patches, num_patches, x.shape[2])
+ x = self.neck(x.permute(0, 3, 1, 2))
+ return x
diff --git a/py/evf_sam/model/EfficientSAM/efficient_sam/mlp.py b/py/evf_sam/model/EfficientSAM/efficient_sam/mlp.py
new file mode 100644
index 0000000..b3be8db
--- /dev/null
+++ b/py/evf_sam/model/EfficientSAM/efficient_sam/mlp.py
@@ -0,0 +1,29 @@
+from typing import Type
+
+from torch import nn
+
+
+# Lightly adapted from
+# https://github.com/facebookresearch/MaskFormer/blob/main/mask_former/modeling/transformer/transformer_predictor.py # noqa
+class MLPBlock(nn.Module):
+ def __init__(
+ self,
+ input_dim: int,
+ hidden_dim: int,
+ output_dim: int,
+ num_layers: int,
+ act: Type[nn.Module],
+ ) -> None:
+ super().__init__()
+ self.num_layers = num_layers
+ h = [hidden_dim] * (num_layers - 1)
+ self.layers = nn.ModuleList(
+ nn.Sequential(nn.Linear(n, k), act())
+ for n, k in zip([input_dim] + h, [hidden_dim] * num_layers)
+ )
+ self.fc = nn.Linear(hidden_dim, output_dim)
+
+ def forward(self, x):
+ for layer in self.layers:
+ x = layer(x)
+ return self.fc(x)
diff --git a/py/evf_sam/model/EfficientSAM/efficient_sam/two_way_transformer.py b/py/evf_sam/model/EfficientSAM/efficient_sam/two_way_transformer.py
new file mode 100644
index 0000000..b06e528
--- /dev/null
+++ b/py/evf_sam/model/EfficientSAM/efficient_sam/two_way_transformer.py
@@ -0,0 +1,266 @@
+import math
+from typing import Tuple, Type
+import torch
+from torch import nn, Tensor
+from .mlp import MLPBlock
+
+
+
+
+class TwoWayTransformer(nn.Module):
+ def __init__(
+ self,
+ depth: int,
+ embedding_dim: int,
+ num_heads: int,
+ mlp_dim: int,
+ activation: Type[nn.Module],
+ normalize_before_activation: bool,
+ attention_downsample_rate: int = 2,
+ ) -> None:
+ """
+ A transformer decoder that attends to an input image using
+ queries whose positional embedding is supplied.
+
+ Args:
+ depth (int): number of layers in the transformer
+ embedding_dim (int): the channel dimension for the input embeddings
+ num_heads (int): the number of heads for multihead attention. Must
+ divide embedding_dim
+ mlp_dim (int): the channel dimension internal to the MLP block
+ activation (nn.Module): the activation to use in the MLP block
+ """
+ super().__init__()
+ self.depth = depth
+ self.embedding_dim = embedding_dim
+ self.num_heads = num_heads
+ self.mlp_dim = mlp_dim
+ self.layers = nn.ModuleList()
+
+ for i in range(depth):
+ curr_layer = TwoWayAttentionBlock(
+ embedding_dim=embedding_dim,
+ num_heads=num_heads,
+ mlp_dim=mlp_dim,
+ activation=activation,
+ normalize_before_activation=normalize_before_activation,
+ attention_downsample_rate=attention_downsample_rate,
+ skip_first_layer_pe=(i == 0),
+ )
+ self.layers.append(curr_layer)
+
+ self.final_attn_token_to_image = AttentionForTwoWayAttentionBlock(
+ embedding_dim,
+ num_heads,
+ downsample_rate=attention_downsample_rate,
+ )
+ self.norm_final_attn = nn.LayerNorm(embedding_dim)
+
+ def forward(
+ self,
+ image_embedding: Tensor,
+ image_pe: Tensor,
+ point_embedding: Tensor,
+ ) -> Tuple[Tensor, Tensor]:
+ """
+ Args:
+ image_embedding (torch.Tensor): image to attend to. Should be shape
+ B x embedding_dim x h x w for any h and w.
+ image_pe (torch.Tensor): the positional encoding to add to the image. Must
+ have the same shape as image_embedding.
+ point_embedding (torch.Tensor): the embedding to add to the query points.
+ Must have shape B x N_points x embedding_dim for any N_points.
+
+ Returns:
+ torch.Tensor: the processed point_embedding
+ torch.Tensor: the processed image_embedding
+ """
+
+ # BxCxHxW -> BxHWxC == B x N_image_tokens x C
+ bs, c, h, w = image_embedding.shape
+ image_embedding = image_embedding.flatten(2).permute(0, 2, 1)
+ image_pe = image_pe.flatten(2).permute(0, 2, 1)
+
+ # Prepare queries
+ queries = point_embedding
+ keys = image_embedding
+
+ # Apply transformer blocks and final layernorm
+ for idx, layer in enumerate(self.layers):
+ queries, keys = layer(
+ queries=queries,
+ keys=keys,
+ query_pe=point_embedding,
+ key_pe=image_pe,
+ )
+
+ # Apply the final attention layer from the points to the image
+ q = queries + point_embedding
+ k = keys + image_pe
+ attn_out = self.final_attn_token_to_image(q=q, k=k, v=keys)
+ queries = queries + attn_out
+ queries = self.norm_final_attn(queries)
+ return queries, keys
+
+
+class TwoWayAttentionBlock(nn.Module):
+ def __init__(
+ self,
+ embedding_dim: int,
+ num_heads: int,
+ mlp_dim: int,
+ activation: Type[nn.Module],
+ normalize_before_activation: bool,
+ attention_downsample_rate: int = 2,
+ skip_first_layer_pe: bool = False,
+ ) -> None:
+ """
+ A transformer block with four layers: (1) self-attention of sparse
+ inputs, (2) cross attention of sparse inputs to dense inputs, (3) mlp
+ block on sparse inputs, and (4) cross attention of dense inputs to sparse
+ inputs.
+
+ Arguments:
+ embedding_dim (int): the channel dimension of the embeddings
+ num_heads (int): the number of heads in the attention layers
+ mlp_dim (int): the hidden dimension of the mlp block
+ activation (nn.Module): the activation of the mlp block
+ skip_first_layer_pe (bool): skip the PE on the first layer
+ """
+ super().__init__()
+ self.self_attn = AttentionForTwoWayAttentionBlock(embedding_dim, num_heads)
+ self.norm1 = nn.LayerNorm(embedding_dim)
+
+ self.cross_attn_token_to_image = AttentionForTwoWayAttentionBlock(
+ embedding_dim,
+ num_heads,
+ downsample_rate=attention_downsample_rate,
+ )
+ self.norm2 = nn.LayerNorm(embedding_dim)
+
+ self.mlp = MLPBlock(
+ embedding_dim,
+ mlp_dim,
+ embedding_dim,
+ 1,
+ activation,
+ )
+
+ self.norm3 = nn.LayerNorm(embedding_dim)
+
+ self.norm4 = nn.LayerNorm(embedding_dim)
+ self.cross_attn_image_to_token = AttentionForTwoWayAttentionBlock(
+ embedding_dim,
+ num_heads,
+ downsample_rate=attention_downsample_rate,
+ )
+
+ self.skip_first_layer_pe = skip_first_layer_pe
+
+ def forward(
+ self, queries: Tensor, keys: Tensor, query_pe: Tensor, key_pe: Tensor
+ ) -> Tuple[Tensor, Tensor]:
+ # Self attention block
+ if not self.skip_first_layer_pe:
+ queries = queries + query_pe
+ attn_out = self.self_attn(q=queries, k=queries, v=queries)
+ queries = queries + attn_out
+ queries = self.norm1(queries)
+
+ # Cross attention block, tokens attending to image embedding
+ q = queries + query_pe
+ k = keys + key_pe
+ attn_out = self.cross_attn_token_to_image(q=q, k=k, v=keys)
+ queries = queries + attn_out
+ queries = self.norm2(queries)
+
+ # MLP block
+ mlp_out = self.mlp(queries)
+ queries = queries + mlp_out
+ queries = self.norm3(queries)
+
+ # Cross attention block, image embedding attending to tokens
+ q = queries + query_pe
+ k = keys + key_pe
+ attn_out = self.cross_attn_image_to_token(q=k, k=q, v=queries)
+ keys = keys + attn_out
+ keys = self.norm4(keys)
+
+ return queries, keys
+
+
+class AttentionForTwoWayAttentionBlock(nn.Module):
+ """
+ An attention layer that allows for downscaling the size of the embedding
+ after projection to queries, keys, and values.
+ """
+
+ def __init__(
+ self,
+ embedding_dim: int,
+ num_heads: int,
+ downsample_rate: int = 1,
+ ) -> None:
+ super().__init__()
+ self.embedding_dim = embedding_dim
+ self.internal_dim = embedding_dim // downsample_rate
+ self.num_heads = num_heads
+ assert (
+ self.internal_dim % num_heads == 0
+ ), "num_heads must divide embedding_dim."
+ self.c_per_head = self.internal_dim / num_heads
+ self.inv_sqrt_c_per_head = 1.0 / math.sqrt(self.c_per_head)
+
+ self.q_proj = nn.Linear(embedding_dim, self.internal_dim)
+ self.k_proj = nn.Linear(embedding_dim, self.internal_dim)
+ self.v_proj = nn.Linear(embedding_dim, self.internal_dim)
+ self.out_proj = nn.Linear(self.internal_dim, embedding_dim)
+ self._reset_parameters()
+
+ def _reset_parameters(self) -> None:
+ # The fan_out is incorrect, but matches pytorch's initialization
+ # for which qkv is a single 3*embedding_dim x embedding_dim matrix
+ fan_in = self.embedding_dim
+ fan_out = 3 * self.internal_dim
+ # Xavier uniform with our custom fan_out
+ bnd = math.sqrt(6 / (fan_in + fan_out))
+ nn.init.uniform_(self.q_proj.weight, -bnd, bnd)
+ nn.init.uniform_(self.k_proj.weight, -bnd, bnd)
+ nn.init.uniform_(self.v_proj.weight, -bnd, bnd)
+ # out_proj.weight is left with default initialization, like pytorch attention
+ nn.init.zeros_(self.q_proj.bias)
+ nn.init.zeros_(self.k_proj.bias)
+ nn.init.zeros_(self.v_proj.bias)
+ nn.init.zeros_(self.out_proj.bias)
+
+ def _separate_heads(self, x: Tensor, num_heads: int) -> Tensor:
+ b, n, c = x.shape
+ x = x.reshape(b, n, num_heads, c // num_heads)
+ return x.transpose(1, 2) # B x N_heads x N_tokens x C_per_head
+
+ def _recombine_heads(self, x: Tensor) -> Tensor:
+ b, n_heads, n_tokens, c_per_head = x.shape
+ x = x.transpose(1, 2)
+ return x.reshape(b, n_tokens, n_heads * c_per_head) # B x N_tokens x C
+
+ def forward(self, q: Tensor, k: Tensor, v: Tensor) -> Tensor:
+ # Input projections
+ q = self.q_proj(q)
+ k = self.k_proj(k)
+ v = self.v_proj(v)
+
+ # Separate into heads
+ q = self._separate_heads(q, self.num_heads)
+ k = self._separate_heads(k, self.num_heads)
+ v = self._separate_heads(v, self.num_heads)
+
+ # Attention
+ _, _, _, c_per_head = q.shape
+ attn = q @ k.permute(0, 1, 3, 2) # B x N_heads x N_tokens x N_tokens
+ attn = attn * self.inv_sqrt_c_per_head
+ attn = torch.softmax(attn, dim=-1)
+ # Get output
+ out = attn @ v
+ out = self._recombine_heads(out)
+ out = self.out_proj(out)
+ return out
diff --git a/py/evf_sam/model/__init__.py b/py/evf_sam/model/__init__.py
new file mode 100644
index 0000000..e69de29
diff --git a/py/evf_sam/model/configuration_evf.py b/py/evf_sam/model/configuration_evf.py
new file mode 100644
index 0000000..fc1383f
--- /dev/null
+++ b/py/evf_sam/model/configuration_evf.py
@@ -0,0 +1,113 @@
+# coding=utf-8
+# Copyright 2022 EleutherAI and the HuggingFace Inc. team. All rights reserved.
+#
+# This code is based on EleutherAI's GPT-NeoX library and the GPT-NeoX
+# and OPT implementations in this library. It has been modified from its
+# original forms to accommodate minor architectural differences compared
+# to GPT-NeoX and OPT used by the Meta AI team that trained the model.
+#
+# Licensed under the Apache License, Version 2.0 (the "License");
+# you may not use this file except in compliance with the License.
+# You may obtain a copy of the License at
+#
+# http://www.apache.org/licenses/LICENSE-2.0
+#
+# Unless required by applicable law or agreed to in writing, software
+# distributed under the License is distributed on an "AS IS" BASIS,
+# WITHOUT WARRANTIES OR CONDITIONS OF ANY KIND, either express or implied.
+# See the License for the specific language governing permissions and
+# limitations under the License.
+""" Evf model configuration"""
+
+from transformers.configuration_utils import PretrainedConfig
+from transformers.utils import logging
+
+logger = logging.get_logger(__name__)
+
+EVF_PRETRAINED_CONFIG_ARCHIVE_MAP = {}
+
+
+class EvfConfig(PretrainedConfig):
+ r"""
+ This is the configuration class to store the configuration of a [`EvfSam`].
+
+ Configuration objects inherit from [`PretrainedConfig`] and can be used to control the model outputs. Read the
+ documentation from [`PretrainedConfig`] for more information.
+
+ Args:
+ hidden_size (`int`, *optional*, defaults to 4096):
+ Dimension of the hidden representations.
+ pretraining_tp (`int`, *optional*, defaults to `1`):
+ Experimental feature. Tensor parallelism rank used during pretraining. Please refer to [this
+ document](https://huggingface.co/docs/transformers/parallelism) to understand more about it. This value is
+ necessary to ensure exact reproducibility of the pretraining results. Please refer to [this
+ issue](https://github.com/pytorch/pytorch/issues/76232).
+ rope_scaling (`Dict`, *optional*):
+ Dictionary containing the scaling configuration for the RoPE embeddings. Currently supports three scaling
+ strategies: linear and dynamic. Their scaling factor must be an float greater than 1. The expected format
+ is `{"type": strategy name, "factor": scaling factor}`. When using this flag, don't update
+ `max_position_embeddings` to the expected new maximum. See the following thread for more information on how
+ these scaling strategies behave:
+ https://www.reddit.com/r/LocalLLaMA/comments/14mrgpr/dynamically_scaled_rope_further_increases/. This is an
+ experimental feature, subject to breaking API changes in future versions.
+
+ Example:
+
+ ```python
+
+ >>> configuration = EvfConfig()
+ >>> model = EvfSam(configuration)
+
+ >>> # Accessing the model configuration
+ >>> configuration = model.config
+ ```"""
+ model_type = "evf"
+ keys_to_ignore_at_inference = ["past_key_values"]
+
+ def __init__(
+ self,
+ hidden_size=768,
+ pad_token_id=1,
+ bos_token_id=0,
+ eos_token_id=2,
+ pretraining_tp=1,
+ tie_word_embeddings=False,
+ rope_scaling=None,
+ out_dim=256,
+ **kwargs,
+ ):
+ self.hidden_size = hidden_size
+ self.out_dim = out_dim
+
+ # self.pretraining_tp = pretraining_tp
+ # self.rope_scaling = rope_scaling
+ # self._rope_scaling_validation()
+
+ super().__init__(
+ pad_token_id=pad_token_id,
+ bos_token_id=bos_token_id,
+ eos_token_id=eos_token_id,
+ tie_word_embeddings=tie_word_embeddings,
+ **kwargs,
+ )
+
+ def _rope_scaling_validation(self):
+ """
+ Validate the `rope_scaling` configuration.
+ """
+ if self.rope_scaling is None:
+ return
+
+ if not isinstance(self.rope_scaling, dict) or len(self.rope_scaling) != 2:
+ raise ValueError(
+ "`rope_scaling` must be a dictionary with with two fields, `name` and `factor`, "
+ f"got {self.rope_scaling}"
+ )
+ rope_scaling_type = self.rope_scaling.get("type", None)
+ rope_scaling_factor = self.rope_scaling.get("factor", None)
+ if rope_scaling_type is None or rope_scaling_type not in ["linear", "dynamic"]:
+ raise ValueError(
+ f"`rope_scaling`'s name field must be one of ['linear', 'dynamic'], got {rope_scaling_type}"
+ )
+ if rope_scaling_factor is None or not isinstance(rope_scaling_factor, float) or rope_scaling_factor <= 1.0:
+ raise ValueError(f"`rope_scaling`'s factor field must be an float > 1, got {rope_scaling_factor}")
diff --git a/py/evf_sam/model/evf_effisam.py b/py/evf_sam/model/evf_effisam.py
new file mode 100644
index 0000000..9820624
--- /dev/null
+++ b/py/evf_sam/model/evf_effisam.py
@@ -0,0 +1,313 @@
+from typing import List, Tuple
+
+import torch
+import torch.nn as nn
+import torch.nn.functional as F
+from transformers import PreTrainedModel, AutoConfig, AutoModelForCausalLM
+from .EfficientSAM.efficient_sam.build_efficient_sam import build_efficient_sam_vits, build_efficient_sam_vitt
+from .unilm.beit3.modeling_utils import BEiT3Wrapper, _get_base_config, _get_large_config
+from .configuration_evf import EvfConfig
+
+
+def dice_loss(
+ inputs: torch.Tensor,
+ targets: torch.Tensor,
+ num_masks: float,
+ scale=1000, # 100000.0,
+ eps=1e-6,
+):
+ """
+ Compute the DICE loss, similar to generalized IOU for masks
+ Args:
+ inputs: A float tensor of arbitrary shape.
+ The predictions for each example.
+ targets: A float tensor with the same shape as inputs. Stores the binary
+ classification label for each element in inputs
+ (0 for the negative class and 1 for the positive class).
+ """
+ inputs = inputs.sigmoid()
+ inputs = inputs.flatten(1, 2)
+ targets = targets.flatten(1, 2)
+ numerator = 2 * (inputs / scale * targets).sum(-1)
+ denominator = (inputs / scale).sum(-1) + (targets / scale).sum(-1)
+ loss = 1 - (numerator + eps) / (denominator + eps)
+ loss = loss.sum() / (num_masks + 1e-8)
+ return loss
+
+
+def sigmoid_ce_loss(
+ inputs: torch.Tensor,
+ targets: torch.Tensor,
+ num_masks: float,
+):
+ """
+ Args:
+ inputs: A float tensor of arbitrary shape.
+ The predictions for each example.
+ targets: A float tensor with the same shape as inputs. Stores the binary
+ classification label for each element in inputs
+ (0 for the negative class and 1 for the positive class).
+ Returns:
+ Loss tensor
+ """
+ loss = F.binary_cross_entropy_with_logits(inputs, targets, reduction="none")
+ loss = loss.flatten(1, 2).mean(1).sum() / (num_masks + 1e-8)
+ return loss
+
+
+
+class EvfEffiSamModel(PreTrainedModel):
+ config_class = EvfConfig
+ def __init__(
+ self,
+ config,
+ **kwargs
+ ):
+ super(EvfEffiSamModel, self).__init__(config)
+
+ self.config = config
+ self.vision_pretrained = kwargs.get("vision_pretrained", None)
+ self.encoder_pretrained = kwargs.get("encoder_pretrained", None)
+ self.dice_loss_weight = kwargs.get("dice_loss_weight", None)
+ self.bce_loss_weight = kwargs.get("bce_loss_weight", None)
+ self.train_mask_decoder = kwargs.get("train_mask_decoder", False)
+ self.initialize_evf_modules(config)
+
+
+ def initialize_evf_modules(self, config):
+ # EffiSAM
+ if config.sam_scale=="tiny":
+ self.visual_model = build_efficient_sam_vitt(self.vision_pretrained)
+ elif config.sam_scale=="small":
+ # vits scale, or without pretrained weight (self.vision_pretrained=None)
+ self.visual_model = build_efficient_sam_vits(self.vision_pretrained)
+ else:
+ raise NotImplementedError
+
+ for param in self.visual_model.parameters():
+ param.requires_grad = False
+ if self.train_mask_decoder:
+ self.visual_model.mask_decoder.train()
+ for param in self.visual_model.mask_decoder.parameters():
+ param.requires_grad = True
+
+ # beit-3
+ if self.config.mm_extractor_scale == "base":
+ beit_config = _get_base_config()
+ elif self.config.mm_extractor_scale == "large":
+ beit_config = _get_large_config()
+ else:
+ raise AttributeError(f"model config should contain key 'mm_extractor_scale', with value 'base' or 'large'.")
+
+ self.mm_extractor = BEiT3Wrapper(beit_config)
+ if self.encoder_pretrained is not None:
+ beit_state_dict = torch.load(self.encoder_pretrained)["model"]
+ self.mm_extractor.load_state_dict(
+ beit_state_dict,
+ strict=False
+ )
+
+ for param in self.mm_extractor.parameters():
+ param.requires_grad = True
+
+ # Projection layer
+ in_dim = config.hidden_size
+ assert in_dim==beit_config.encoder_embed_dim, \
+ f"projection layer dim {in_dim} mismatch with mm_extractor dim {beit_config.encoder_embed_dim}"
+ out_dim = config.out_dim
+ text_fc = [
+ nn.Linear(in_dim, in_dim),
+ nn.ReLU(),
+ nn.Linear(in_dim, out_dim)
+ ]
+ self.text_hidden_fcs = nn.ModuleList([nn.Sequential(*text_fc)])
+ self.text_hidden_fcs.train()
+ for param in self.text_hidden_fcs.parameters():
+ param.requires_grad = True
+
+ def get_visual_embs(self, pixel_values: torch.Tensor):
+ with torch.no_grad():
+ image_embeddings_list = []
+ for i in range(pixel_values.shape[0]):
+ torch.cuda.empty_cache()
+ image_embeddings = self.visual_model.image_encoder(
+ pixel_values[i].unsqueeze(0)
+ )
+ image_embeddings_list.append(image_embeddings)
+ torch.cuda.empty_cache()
+ image_embeddings = torch.cat(image_embeddings_list, 0)
+ return image_embeddings
+
+ def forward(
+ self,
+ images: torch.Tensor,
+ images_evf: torch.Tensor,
+ input_ids: torch.Tensor,
+ attention_masks: torch.Tensor,
+ offset: torch.Tensor,
+ masks_list: List[torch.Tensor],
+ label_list: List[torch.Tensor],
+ resize_list: List[tuple],
+ inference: bool = False,
+ **kwargs,
+ ):
+ image_embeddings = self.get_visual_embs(images)
+ batch_size = image_embeddings.shape[0]
+ assert batch_size == len(offset) - 1
+
+ images_evf_list = []
+ for i in range(len(offset) - 1):
+ start_i, end_i = offset[i], offset[i + 1]
+ images_evf_i = (
+ images_evf[i]
+ .unsqueeze(0)
+ .expand(end_i - start_i, -1, -1, -1)
+ .contiguous()
+ )
+ images_evf_list.append(images_evf_i)
+ images_evf = torch.cat(images_evf_list, dim=0)
+
+ multimask_output = False
+ output = self.mm_extractor.beit3(
+ visual_tokens=images_evf,
+ textual_tokens=input_ids,
+ text_padding_position=~attention_masks
+ )
+
+ feat = output["encoder_out"][:, :1, ...]
+
+ feat = self.text_hidden_fcs[0](feat)
+ feat = torch.split(feat, [offset[i+1] - offset[i] for i in range(len(offset)-1)])
+
+ pred_masks = []
+ for i in range(len(feat)):
+ sparse_embeddings = feat[i].unsqueeze(0)
+ sparse_embeddings = sparse_embeddings.to(feat[i].dtype)
+ low_res_masks, iou_predictions = self.visual_model.mask_decoder(
+ image_embeddings=image_embeddings[i].unsqueeze(0),
+ image_pe=self.visual_model.prompt_encoder.get_dense_pe(),
+ sparse_prompt_embeddings=sparse_embeddings,
+ multimask_output=multimask_output,
+ )
+
+ if multimask_output:
+ sorted_ids = torch.argsort(iou_predictions, dim=-1, descending=True)
+ low_res_masks = torch.take_along_dim(low_res_masks, sorted_ids[..., None, None], dim=1)
+
+ pred_mask = self.postprocess_masks(
+ low_res_masks[:, :1],
+ input_size=resize_list[i],
+ original_size=label_list[i].shape,
+ )
+ pred_masks.append(pred_mask[:, 0])
+
+ gt_masks = masks_list
+
+ if inference:
+ return {
+ "pred_masks": pred_masks,
+ "gt_masks": gt_masks,
+ }
+
+ mask_bce_loss = 0
+ mask_dice_loss = 0
+ num_masks = 0
+ for batch_idx in range(len(pred_masks)):
+ gt_mask = gt_masks[batch_idx]
+ pred_mask = pred_masks[batch_idx]
+
+ assert (
+ gt_mask.shape[0] == pred_mask.shape[0]
+ ), "gt_mask.shape: {}, pred_mask.shape: {}".format(
+ gt_mask.shape, pred_mask.shape
+ )
+ mask_bce_loss += (
+ sigmoid_ce_loss(pred_mask, gt_mask, num_masks=gt_mask.shape[0])
+ * gt_mask.shape[0]
+ )
+ mask_dice_loss += (
+ dice_loss(pred_mask, gt_mask, num_masks=gt_mask.shape[0])
+ * gt_mask.shape[0]
+ )
+ num_masks += gt_mask.shape[0]
+
+ mask_bce_loss = self.bce_loss_weight * mask_bce_loss / (num_masks + 1e-8)
+ mask_dice_loss = self.dice_loss_weight * mask_dice_loss / (num_masks + 1e-8)
+ mask_loss = mask_bce_loss + mask_dice_loss
+
+ loss = mask_loss
+
+ return {
+ "loss": loss,
+ "mask_bce_loss": mask_bce_loss,
+ "mask_dice_loss": mask_dice_loss,
+ "mask_loss": mask_loss,
+ }
+
+ def postprocess_masks(
+ self,
+ masks: torch.Tensor,
+ input_size: Tuple[int, ...],
+ original_size: Tuple[int, ...],
+ ) -> torch.Tensor:
+ """
+ pre-process of Effi-SAM is different from SAM, where there is no padding,
+ so cropping is not needed in post-process.
+ """
+
+ dtype = masks.dtype
+
+ # masks = F.interpolate(
+ # masks.float(),
+ # (1024, 1024),
+ # mode="bilinear",
+ # align_corners=False,
+ # )
+ # masks = masks.to(dtype)
+ # masks = masks[..., : input_size[0], : input_size[1]]
+
+ masks = F.interpolate(
+ masks, original_size, mode="bilinear", align_corners=False
+ )
+ masks = masks.to(dtype)
+ return masks
+
+ def inference(
+ self,
+ images,
+ images_evf,
+ input_ids,
+ resize_list,
+ original_size_list,
+ multimask_output=False,
+ ):
+ with torch.no_grad():
+ image_embeddings = self.visual_model.image_encoder(images)
+
+ output = self.mm_extractor.beit3(visual_tokens=images_evf, textual_tokens=input_ids, text_padding_position=torch.zeros_like(input_ids))
+
+ feat = output["encoder_out"][:, :1, ...]
+ feat = self.text_hidden_fcs[0](feat)
+ sparse_embeddings = feat.unsqueeze(0)
+ sparse_embeddings = sparse_embeddings.to(feat.dtype)
+ low_res_masks, iou_predictions = self.visual_model.mask_decoder(
+ image_embeddings=image_embeddings,
+ image_pe=self.visual_model.prompt_encoder.get_dense_pe(),
+ sparse_prompt_embeddings=sparse_embeddings,
+ multimask_output=multimask_output,
+ )
+ if multimask_output:
+ sorted_ids = torch.argsort(iou_predictions, dim=-1, descending=True)
+ low_res_masks = torch.take_along_dim(low_res_masks, sorted_ids[..., None, None], dim=1)
+
+ pred_mask = self.postprocess_masks(
+ low_res_masks[:, :1],
+ input_size=resize_list[0],
+ original_size=original_size_list[0],
+ )
+
+ return pred_mask[:, 0]
+
+
+AutoConfig.register("evf", EvfConfig)
+AutoModelForCausalLM.register(EvfConfig, EvfEffiSamModel)
diff --git a/py/evf_sam/model/evf_sam.py b/py/evf_sam/model/evf_sam.py
new file mode 100644
index 0000000..a0ec88e
--- /dev/null
+++ b/py/evf_sam/model/evf_sam.py
@@ -0,0 +1,303 @@
+from typing import List
+
+import torch
+import torch.nn as nn
+import torch.nn.functional as F
+from transformers import PreTrainedModel, AutoConfig, AutoModelForCausalLM
+from .segment_anything import build_sam_vit_h
+from .unilm.beit3.modeling_utils import BEiT3Wrapper, _get_base_config, _get_large_config
+from .configuration_evf import EvfConfig
+
+def dice_loss(
+ inputs: torch.Tensor,
+ targets: torch.Tensor,
+ num_masks: float,
+ scale=1000, # 100000.0,
+ eps=1e-6,
+):
+ """
+ Compute the DICE loss, similar to generalized IOU for masks
+ Args:
+ inputs: A float tensor of arbitrary shape.
+ The predictions for each example.
+ targets: A float tensor with the same shape as inputs. Stores the binary
+ classification label for each element in inputs
+ (0 for the negative class and 1 for the positive class).
+ """
+ inputs = inputs.sigmoid()
+ inputs = inputs.flatten(1, 2)
+ targets = targets.flatten(1, 2)
+ numerator = 2 * (inputs / scale * targets).sum(-1)
+ denominator = (inputs / scale).sum(-1) + (targets / scale).sum(-1)
+ loss = 1 - (numerator + eps) / (denominator + eps)
+ loss = loss.sum() / (num_masks + 1e-8)
+ return loss
+
+
+def sigmoid_ce_loss(
+ inputs: torch.Tensor,
+ targets: torch.Tensor,
+ num_masks: float,
+):
+ """
+ Args:
+ inputs: A float tensor of arbitrary shape.
+ The predictions for each example.
+ targets: A float tensor with the same shape as inputs. Stores the binary
+ classification label for each element in inputs
+ (0 for the negative class and 1 for the positive class).
+ Returns:
+ Loss tensor
+ """
+ loss = F.binary_cross_entropy_with_logits(inputs, targets, reduction="none")
+ loss = loss.flatten(1, 2).mean(1).sum() / (num_masks + 1e-8)
+ return loss
+
+
+
+class EvfSamModel(PreTrainedModel):
+ config_class = EvfConfig
+ def __init__(
+ self,
+ config,
+ **kwargs
+ ):
+ super(EvfSamModel, self).__init__(config)
+
+ self.config = config
+ self.vision_pretrained = kwargs.get("vision_pretrained", None)
+ self.encoder_pretrained = kwargs.get("encoder_pretrained", None)
+ self.dice_loss_weight = kwargs.get("dice_loss_weight", None)
+ self.bce_loss_weight = kwargs.get("bce_loss_weight", None)
+ self.train_mask_decoder = kwargs.get("train_mask_decoder", False)
+ self.train_prompt_encoder = kwargs.get("train_prompt_encoder", False)
+ self.initialize_evf_modules(config)
+
+
+ def initialize_evf_modules(self, config):
+ # SAM
+ if config.sam_scale=="huge":
+ self.visual_model = build_sam_vit_h(self.vision_pretrained)
+ else:
+ raise NotImplementedError
+
+ for param in self.visual_model.parameters():
+ param.requires_grad = False
+ if self.train_mask_decoder:
+ self.visual_model.mask_decoder.train()
+ for param in self.visual_model.mask_decoder.parameters():
+ param.requires_grad = True
+ if self.train_prompt_encoder:
+ self.visual_model.prompt_encoder.no_mask_embed.requires_grad_(True)
+
+ # beit-3
+ if self.config.mm_extractor_scale == "base":
+ beit_config = _get_base_config()
+ elif self.config.mm_extractor_scale == "large":
+ beit_config = _get_large_config()
+ else:
+ raise AttributeError(f"model config should contain key 'mm_extractor_scale', with value 'base' or 'large'.")
+
+ self.mm_extractor = BEiT3Wrapper(beit_config)
+ if self.encoder_pretrained is not None:
+ beit_state_dict = torch.load(self.encoder_pretrained)["model"]
+ self.mm_extractor.load_state_dict(
+ beit_state_dict,
+ strict=False
+ )
+
+ for param in self.mm_extractor.parameters():
+ param.requires_grad = True
+
+ # Projection layer
+ in_dim = config.hidden_size
+ assert in_dim==beit_config.encoder_embed_dim, \
+ f"projection layer dim {in_dim} mismatch with mm_extractor dim {beit_config.encoder_embed_dim}"
+ out_dim = config.out_dim
+ text_fc = [
+ nn.Linear(in_dim, in_dim),
+ nn.ReLU(),
+ nn.Linear(in_dim, out_dim)
+ ]
+ self.text_hidden_fcs = nn.ModuleList([nn.Sequential(*text_fc)])
+ self.text_hidden_fcs.train()
+ for param in self.text_hidden_fcs.parameters():
+ param.requires_grad = True
+
+ def get_visual_embs(self, pixel_values: torch.FloatTensor):
+ with torch.no_grad():
+ image_embeddings_list = []
+ for i in range(pixel_values.shape[0]):
+ torch.cuda.empty_cache()
+ image_embeddings = self.visual_model.image_encoder(
+ pixel_values[i].unsqueeze(0)
+ )
+ image_embeddings_list.append(image_embeddings)
+ torch.cuda.empty_cache()
+ image_embeddings = torch.cat(image_embeddings_list, 0)
+ return image_embeddings
+
+ def forward(
+ self,
+ images: torch.FloatTensor,
+ images_evf: torch.FloatTensor,
+ input_ids: torch.LongTensor,
+ attention_masks: torch.LongTensor,
+ offset: torch.LongTensor,
+ masks_list: List[torch.FloatTensor],
+ label_list: List[torch.Tensor],
+ resize_list: List[tuple],
+ inference: bool = False,
+ **kwargs,
+ ):
+ image_embeddings = self.get_visual_embs(images)
+ batch_size = image_embeddings.shape[0]
+ assert batch_size == len(offset) - 1
+
+ images_evf_list = []
+ for i in range(len(offset) - 1):
+ start_i, end_i = offset[i], offset[i + 1]
+ images_evf_i = (
+ images_evf[i]
+ .unsqueeze(0)
+ .expand(end_i - start_i, -1, -1, -1)
+ .contiguous()
+ )
+ images_evf_list.append(images_evf_i)
+ images_evf = torch.cat(images_evf_list, dim=0)
+
+ multimask_output = False
+ output = self.mm_extractor.beit3(
+ visual_tokens=images_evf,
+ textual_tokens=input_ids,
+ text_padding_position=~attention_masks
+ )
+
+ feat = output["encoder_out"][:, :1, ...]
+
+ feat = self.text_hidden_fcs[0](feat)
+ feat = torch.split(feat, [offset[i+1] - offset[i] for i in range(len(offset)-1)])
+
+ pred_masks = []
+ for i in range(len(feat)):
+ (
+ sparse_embeddings,
+ dense_embeddings,
+ ) = self.visual_model.prompt_encoder(
+ points=None,
+ boxes=None,
+ masks=None,
+ text_embeds=feat[i],
+ )
+ sparse_embeddings = sparse_embeddings.to(feat[i].dtype)
+ low_res_masks, iou_predictions = self.visual_model.mask_decoder(
+ image_embeddings=image_embeddings[i].unsqueeze(0),
+ image_pe=self.visual_model.prompt_encoder.get_dense_pe(),
+ sparse_prompt_embeddings=sparse_embeddings,
+ dense_prompt_embeddings=dense_embeddings,
+ multimask_output=multimask_output,
+ )
+
+ if multimask_output:
+ sorted_ids = torch.argsort(iou_predictions, dim=-1, descending=True)
+ low_res_masks = torch.take_along_dim(low_res_masks, sorted_ids[..., None, None], dim=1)[:, :1]
+
+ pred_mask = self.visual_model.postprocess_masks(
+ low_res_masks,
+ input_size=resize_list[i],
+ original_size=label_list[i].shape,
+ )
+ pred_masks.append(pred_mask[:, 0])
+
+ gt_masks = masks_list
+
+ if inference:
+ return {
+ "pred_masks": pred_masks,
+ "gt_masks": gt_masks,
+ }
+
+ mask_bce_loss = 0
+ mask_dice_loss = 0
+ num_masks = 0
+ for batch_idx in range(len(pred_masks)):
+ gt_mask = gt_masks[batch_idx]
+ pred_mask = pred_masks[batch_idx]
+
+ assert (
+ gt_mask.shape[0] == pred_mask.shape[0]
+ ), "gt_mask.shape: {}, pred_mask.shape: {}".format(
+ gt_mask.shape, pred_mask.shape
+ )
+ mask_bce_loss += (
+ sigmoid_ce_loss(pred_mask, gt_mask, num_masks=gt_mask.shape[0])
+ * gt_mask.shape[0]
+ )
+ mask_dice_loss += (
+ dice_loss(pred_mask, gt_mask, num_masks=gt_mask.shape[0])
+ * gt_mask.shape[0]
+ )
+ num_masks += gt_mask.shape[0]
+
+ mask_bce_loss = self.bce_loss_weight * mask_bce_loss / (num_masks + 1e-8)
+ mask_dice_loss = self.dice_loss_weight * mask_dice_loss / (num_masks + 1e-8)
+ mask_loss = mask_bce_loss + mask_dice_loss
+
+ loss = mask_loss
+
+ return {
+ "loss": loss,
+ "mask_bce_loss": mask_bce_loss,
+ "mask_dice_loss": mask_dice_loss,
+ "mask_loss": mask_loss,
+ }
+
+ def inference(
+ self,
+ images,
+ images_evf,
+ input_ids,
+ resize_list,
+ original_size_list,
+ multimask_output=False,
+ ):
+ with torch.no_grad():
+ image_embeddings = self.visual_model.image_encoder(images)
+ multimask_output = multimask_output
+
+ output = self.mm_extractor.beit3(visual_tokens=images_evf, textual_tokens=input_ids, text_padding_position=torch.zeros_like(input_ids))
+
+ feat = output["encoder_out"][:, :1, ...]
+ feat = self.text_hidden_fcs[0](feat)
+ (
+ sparse_embeddings,
+ dense_embeddings,
+ ) = self.visual_model.prompt_encoder(
+ points=None,
+ boxes=None,
+ masks=None,
+ text_embeds=feat,
+ )
+ sparse_embeddings = sparse_embeddings.to(feat.dtype)
+ low_res_masks, iou_predictions = self.visual_model.mask_decoder(
+ image_embeddings=image_embeddings,
+ image_pe=self.visual_model.prompt_encoder.get_dense_pe(),
+ sparse_prompt_embeddings=sparse_embeddings,
+ dense_prompt_embeddings=dense_embeddings,
+ multimask_output=multimask_output,
+ )
+ if multimask_output:
+ sorted_ids = torch.argsort(iou_predictions, dim=-1, descending=True)
+ low_res_masks = torch.take_along_dim(low_res_masks, sorted_ids[..., None, None], dim=1)[:, :1]
+
+ pred_mask = self.visual_model.postprocess_masks(
+ low_res_masks,
+ input_size=resize_list[0],
+ original_size=original_size_list[0],
+ )
+
+ return pred_mask[:, 0]
+
+
+AutoConfig.register("evf", EvfConfig)
+AutoModelForCausalLM.register(EvfConfig, EvfSamModel)
\ No newline at end of file
diff --git a/py/evf_sam/model/evf_sam2.py b/py/evf_sam/model/evf_sam2.py
new file mode 100644
index 0000000..93aa295
--- /dev/null
+++ b/py/evf_sam/model/evf_sam2.py
@@ -0,0 +1,341 @@
+from typing import List
+import os
+import torch
+import torch.nn as nn
+import torch.nn.functional as F
+from transformers import PreTrainedModel, AutoConfig, AutoModelForCausalLM
+from .segment_anything_2.sam2.build_sam import build_sam2
+from .unilm.beit3.modeling_utils import BEiT3Wrapper, _get_base_config, _get_large_config
+from .configuration_evf import EvfConfig
+from .segment_anything_2.sam2.utils.misc import load_video_frames
+from collections import OrderedDict
+
+
+def dice_loss(
+ inputs: torch.Tensor,
+ targets: torch.Tensor,
+ num_masks: float,
+ scale=1000, # 100000.0,
+ eps=1e-6,
+):
+ """
+ Compute the DICE loss, similar to generalized IOU for masks
+ Args:
+ inputs: A float tensor of arbitrary shape.
+ The predictions for each example.
+ targets: A float tensor with the same shape as inputs. Stores the binary
+ classification label for each element in inputs
+ (0 for the negative class and 1 for the positive class).
+ """
+ inputs = inputs.sigmoid()
+ inputs = inputs.flatten(1, 2)
+ targets = targets.flatten(1, 2)
+ numerator = 2 * (inputs / scale * targets).sum(-1)
+ denominator = (inputs / scale).sum(-1) + (targets / scale).sum(-1)
+ loss = 1 - (numerator + eps) / (denominator + eps)
+ loss = loss.sum() / (num_masks + 1e-8)
+ return loss
+
+
+def sigmoid_ce_loss(
+ inputs: torch.Tensor,
+ targets: torch.Tensor,
+ num_masks: float,
+):
+ """
+ Args:
+ inputs: A float tensor of arbitrary shape.
+ The predictions for each example.
+ targets: A float tensor with the same shape as inputs. Stores the binary
+ classification label for each element in inputs
+ (0 for the negative class and 1 for the positive class).
+ Returns:
+ Loss tensor
+ """
+ loss = F.binary_cross_entropy_with_logits(inputs, targets, reduction="none")
+ loss = loss.flatten(1, 2).mean(1).sum() / (num_masks + 1e-8)
+ return loss
+
+class EvfSam2Model(PreTrainedModel):
+ config_class = EvfConfig
+ def __init__(
+ self,
+ config,
+ **kwargs
+ ):
+ super(EvfSam2Model, self).__init__(config)
+
+ self.config = config
+ self.vision_pretrained = kwargs.get("vision_pretrained", None)
+ self.encoder_pretrained = kwargs.get("encoder_pretrained", None)
+ self.dice_loss_weight = kwargs.get("dice_loss_weight", None)
+ self.bce_loss_weight = kwargs.get("bce_loss_weight", None)
+ self.train_mask_decoder = kwargs.get("train_mask_decoder", False)
+ self.train_prompt_encoder = kwargs.get("train_prompt_encoder", False)
+ self.initialize_evf_modules(config)
+ self._bb_feat_sizes = [
+ (256, 256),
+ (128, 128),
+ (64, 64),
+ ]
+
+ def initialize_evf_modules(self, config):
+ # SAM
+ if config.sam_scale=="large":
+ self.visual_model = build_sam2("sam2_hiera_l.yaml", self.vision_pretrained, device=None)
+ elif config.sam_scale=="tiny":
+ self.visual_model = build_sam2("sam2_hiera_t.yaml", self.vision_pretrained, device=None)
+ else:
+ raise NotImplementedError
+
+ for param in self.visual_model.parameters():
+ param.requires_grad = False
+ if self.train_mask_decoder:
+ self.visual_model.sam_mask_decoder.train()
+ for param in self.visual_model.sam_mask_decoder.parameters():
+ param.requires_grad = True
+ if self.train_prompt_encoder:
+ self.visual_model.sam_prompt_encoder.no_mask_embed.requires_grad_(True)
+
+ # beit-3
+ if self.config.mm_extractor_scale == "base":
+ beit_config = _get_base_config()
+ elif self.config.mm_extractor_scale == "large":
+ beit_config = _get_large_config()
+ else:
+ raise AttributeError(f"model config should contain key 'mm_extractor_scale', with value 'base' or 'large'.")
+
+ self.mm_extractor = BEiT3Wrapper(beit_config)
+ if self.encoder_pretrained is not None:
+ beit_state_dict = torch.load(self.encoder_pretrained)["model"]
+ self.mm_extractor.load_state_dict(
+ beit_state_dict,
+ strict=False
+ )
+
+ for param in self.mm_extractor.parameters():
+ param.requires_grad = True
+
+ # Projection layer
+ in_dim = config.hidden_size
+ assert in_dim==beit_config.encoder_embed_dim, \
+ f"projection layer dim {in_dim} mismatch with mm_extractor dim {beit_config.encoder_embed_dim}"
+ out_dim = config.out_dim
+ text_fc = [
+ nn.Linear(in_dim, in_dim),
+ nn.ReLU(),
+ nn.Linear(in_dim, out_dim)
+ ]
+ self.text_hidden_fcs = nn.ModuleList([nn.Sequential(*text_fc)])
+ self.text_hidden_fcs.train()
+ for param in self.text_hidden_fcs.parameters():
+ param.requires_grad = True
+
+ def postprocess_masks(self, masks: torch.Tensor, orig_hw) -> torch.Tensor:
+ """
+ Perform PostProcessing on output masks.
+ """
+ masks = masks.float()
+ masks = F.interpolate(masks, orig_hw, mode="bilinear", align_corners=False)
+ return masks
+
+ def forward(
+ self,
+ images: torch.FloatTensor,
+ images_evf: torch.FloatTensor,
+ input_ids: torch.LongTensor,
+ attention_masks: torch.LongTensor,
+ offset: torch.LongTensor,
+ masks_list: List[torch.FloatTensor],
+ label_list: List[torch.Tensor],
+ resize_list: List[tuple],
+ inference: bool = False,
+ **kwargs,
+ ):
+ # image_embeddings = self.get_visual_embs(images)
+ backbone_out = self.visual_model.forward_image(images)
+ # dict_keys(['vision_features', 'vision_pos_enc', 'backbone_fpn'])
+ _, image_embeddings, _, _ = self.visual_model._prepare_backbone_features(backbone_out)
+ image_embeddings = [_.to(images.dtype) for _ in image_embeddings]
+ batch_size = images.shape[0]
+ if self.visual_model.directly_add_no_mem_embed:
+ image_embeddings[-1] = image_embeddings[-1] + self.visual_model.no_mem_embed
+
+ feats = [
+ feat.permute(1, 2, 0).view(batch_size, -1, *feat_size)
+ for feat, feat_size in zip(image_embeddings[::-1], self._bb_feat_sizes[::-1])
+ ][::-1]
+ _features = {"image_embed": feats[-1], "high_res_feats": feats[:-1]}
+
+
+ assert batch_size == len(offset) - 1
+
+ images_evf_list = []
+ for i in range(len(offset) - 1):
+ start_i, end_i = offset[i], offset[i + 1]
+ images_evf_i = (
+ images_evf[i]
+ .unsqueeze(0)
+ .expand(end_i - start_i, -1, -1, -1)
+ .contiguous()
+ )
+ images_evf_list.append(images_evf_i)
+ images_evf = torch.cat(images_evf_list, dim=0)
+
+ multimask_output = False
+ output = self.mm_extractor.beit3(
+ visual_tokens=images_evf,
+ textual_tokens=input_ids,
+ text_padding_position=~attention_masks
+ )
+
+ feat = output["encoder_out"][:, :1, ...]
+
+ feat = self.text_hidden_fcs[0](feat)
+ feat = torch.split(feat, [offset[i+1] - offset[i] for i in range(len(offset)-1)])
+
+ pred_masks = []
+
+ for i in range(len(feat)):
+ (
+ sparse_embeddings,
+ dense_embeddings,
+ ) = self.visual_model.sam_prompt_encoder(
+ points=None,
+ boxes=None,
+ masks=None,
+ text_embeds=feat[i],
+ )
+ sparse_embeddings = sparse_embeddings.to(feat[i].dtype)
+ high_res_features = [
+ feat_level[i].unsqueeze(0)
+ for feat_level in _features["high_res_feats"]
+ ]
+ low_res_masks, iou_predictions, _, _ = self.visual_model.sam_mask_decoder(
+ image_embeddings=_features["image_embed"][i].unsqueeze(0),
+ image_pe=self.visual_model.sam_prompt_encoder.get_dense_pe(),
+ sparse_prompt_embeddings=sparse_embeddings,
+ dense_prompt_embeddings=dense_embeddings,
+ multimask_output=multimask_output,
+ repeat_image = True,
+ high_res_features=high_res_features,
+ )
+
+ if multimask_output:
+ sorted_ids = torch.argsort(iou_predictions, dim=-1, descending=True)
+ low_res_masks = torch.take_along_dim(low_res_masks, sorted_ids[..., None, None], dim=1)[:, :1]
+
+ pred_mask = self.postprocess_masks(
+ low_res_masks,
+ orig_hw=label_list[i].shape,
+ )
+ pred_masks.append(pred_mask[:, 0])
+
+ gt_masks = masks_list
+
+ if inference:
+ return {
+ "pred_masks": pred_masks,
+ "gt_masks": gt_masks,
+ }
+
+ mask_bce_loss = 0
+ mask_dice_loss = 0
+ num_masks = 0
+ for batch_idx in range(len(pred_masks)):
+ gt_mask = gt_masks[batch_idx]
+ pred_mask = pred_masks[batch_idx]
+
+ assert (
+ gt_mask.shape[0] == pred_mask.shape[0]
+ ), "gt_mask.shape: {}, pred_mask.shape: {}".format(
+ gt_mask.shape, pred_mask.shape
+ )
+ mask_bce_loss += (
+ sigmoid_ce_loss(pred_mask, gt_mask, num_masks=gt_mask.shape[0])
+ * gt_mask.shape[0]
+ )
+ mask_dice_loss += (
+ dice_loss(pred_mask, gt_mask, num_masks=gt_mask.shape[0])
+ * gt_mask.shape[0]
+ )
+ num_masks += gt_mask.shape[0]
+
+ mask_bce_loss = self.bce_loss_weight * mask_bce_loss / (num_masks + 1e-8)
+ mask_dice_loss = self.dice_loss_weight * mask_dice_loss / (num_masks + 1e-8)
+ mask_loss = mask_bce_loss + mask_dice_loss
+
+ loss = mask_loss
+
+ return {
+ "loss": loss,
+ "mask_bce_loss": mask_bce_loss,
+ "mask_dice_loss": mask_dice_loss,
+ "mask_loss": mask_loss,
+ }
+
+ def inference(
+ self,
+ images,
+ images_evf,
+ input_ids,
+ resize_list,
+ original_size_list,
+ multimask_output=False,
+ ):
+ with torch.no_grad():
+ backbone_out = self.visual_model.forward_image(images)
+ # dict_keys(['vision_features', 'vision_pos_enc', 'backbone_fpn'])
+ _, image_embeddings, _, _ = self.visual_model._prepare_backbone_features(backbone_out)
+ image_embeddings = [_.to(images.dtype) for _ in image_embeddings]
+ batch_size = images.shape[0]
+ if self.visual_model.directly_add_no_mem_embed:
+ image_embeddings[-1] = image_embeddings[-1] + self.visual_model.no_mem_embed
+
+ feats = [
+ feat.permute(1, 2, 0).view(batch_size, -1, *feat_size)
+ for feat, feat_size in zip(image_embeddings[::-1], self._bb_feat_sizes[::-1])
+ ][::-1]
+ _features = {"image_embed": feats[-1], "high_res_feats": feats[:-1]}
+
+
+ multimask_output = multimask_output
+
+ output = self.mm_extractor.beit3(visual_tokens=images_evf, textual_tokens=input_ids, text_padding_position=torch.zeros_like(input_ids))
+
+ feat = output["encoder_out"][:, :1, ...]
+ feat = self.text_hidden_fcs[0](feat)
+ (
+ sparse_embeddings,
+ dense_embeddings,
+ ) = self.visual_model.sam_prompt_encoder(
+ points=None,
+ boxes=None,
+ masks=None,
+ text_embeds=feat,
+ )
+ high_res_features = _features["high_res_feats"]
+ sparse_embeddings = sparse_embeddings.to(feat.dtype)
+ low_res_masks, iou_predictions, _, _ = self.visual_model.sam_mask_decoder(
+ image_embeddings=_features["image_embed"],
+ image_pe=self.visual_model.sam_prompt_encoder.get_dense_pe(),
+ sparse_prompt_embeddings=sparse_embeddings,
+ dense_prompt_embeddings=dense_embeddings,
+ multimask_output=multimask_output,
+ repeat_image = True,
+ high_res_features=high_res_features,
+ )
+ if multimask_output:
+ sorted_ids = torch.argsort(iou_predictions, dim=-1, descending=True)
+ low_res_masks = torch.take_along_dim(low_res_masks, sorted_ids[..., None, None], dim=1)[:, :1]
+
+ pred_mask = self.postprocess_masks(
+ low_res_masks,
+ orig_hw=original_size_list[0],
+ )
+
+ return pred_mask[:, 0]
+
+
+AutoConfig.register("evf", EvfConfig)
+AutoModelForCausalLM.register(EvfConfig, EvfSam2Model)
\ No newline at end of file
diff --git a/py/evf_sam/model/evf_sam2_video.py b/py/evf_sam/model/evf_sam2_video.py
new file mode 100644
index 0000000..ef2499a
--- /dev/null
+++ b/py/evf_sam/model/evf_sam2_video.py
@@ -0,0 +1,321 @@
+from typing import List
+import os
+import torch
+import torch.nn as nn
+import torch.nn.functional as F
+from transformers import PreTrainedModel, AutoConfig, AutoModelForCausalLM
+from .segment_anything_2.sam2.build_sam import build_sam2, build_sam2_video_predictor
+from .unilm.beit3.modeling_utils import BEiT3Wrapper, _get_base_config, _get_large_config
+from .configuration_evf import EvfConfig
+from .segment_anything_2.sam2.utils.misc import load_video_frames
+from collections import OrderedDict
+
+
+
+def dice_loss(
+ inputs: torch.Tensor,
+ targets: torch.Tensor,
+ num_masks: float,
+ scale=1000, # 100000.0,
+ eps=1e-6,
+):
+ """
+ Compute the DICE loss, similar to generalized IOU for masks
+ Args:
+ inputs: A float tensor of arbitrary shape.
+ The predictions for each example.
+ targets: A float tensor with the same shape as inputs. Stores the binary
+ classification label for each element in inputs
+ (0 for the negative class and 1 for the positive class).
+ """
+ inputs = inputs.sigmoid()
+ inputs = inputs.flatten(1, 2)
+ targets = targets.flatten(1, 2)
+ numerator = 2 * (inputs / scale * targets).sum(-1)
+ denominator = (inputs / scale).sum(-1) + (targets / scale).sum(-1)
+ loss = 1 - (numerator + eps) / (denominator + eps)
+ loss = loss.sum() / (num_masks + 1e-8)
+ return loss
+
+
+def sigmoid_ce_loss(
+ inputs: torch.Tensor,
+ targets: torch.Tensor,
+ num_masks: float,
+):
+ """
+ Args:
+ inputs: A float tensor of arbitrary shape.
+ The predictions for each example.
+ targets: A float tensor with the same shape as inputs. Stores the binary
+ classification label for each element in inputs
+ (0 for the negative class and 1 for the positive class).
+ Returns:
+ Loss tensor
+ """
+ loss = F.binary_cross_entropy_with_logits(inputs, targets, reduction="none")
+ loss = loss.flatten(1, 2).mean(1).sum() / (num_masks + 1e-8)
+ return loss
+
+class EvfSam2Model(PreTrainedModel):
+ config_class = EvfConfig
+ def __init__(
+ self,
+ config,
+ **kwargs
+ ):
+ super(EvfSam2Model, self).__init__(config)
+
+ self.config = config
+ self.vision_pretrained = kwargs.get("vision_pretrained", None)
+ self.encoder_pretrained = kwargs.get("encoder_pretrained", None)
+ self.dice_loss_weight = kwargs.get("dice_loss_weight", None)
+ self.bce_loss_weight = kwargs.get("bce_loss_weight", None)
+ self.train_mask_decoder = kwargs.get("train_mask_decoder", False)
+ self.train_prompt_encoder = kwargs.get("train_prompt_encoder", False)
+ self.initialize_evf_modules(config)
+ self._bb_feat_sizes = [
+ (256, 256),
+ (128, 128),
+ (64, 64),
+ ]
+
+ def initialize_evf_modules(self, config):
+ # SAM
+ if config.sam_scale=="large":
+ self.visual_model = build_sam2_video_predictor("sam2_hiera_l.yaml", self.vision_pretrained, device=None)
+ elif config.sam_scale=="tiny":
+ self.visual_model = build_sam2_video_predictor("sam2_hiera_t.yaml", self.vision_pretrained, device=None)
+ else:
+ raise NotImplementedError
+
+ for param in self.visual_model.parameters():
+ param.requires_grad = False
+ if self.train_mask_decoder:
+ self.visual_model.sam_mask_decoder.train()
+ for param in self.visual_model.sam_mask_decoder.parameters():
+ param.requires_grad = True
+ if self.train_prompt_encoder:
+ self.visual_model.sam_prompt_encoder.no_mask_embed.requires_grad_(True)
+
+ # beit-3
+ if self.config.mm_extractor_scale == "base":
+ beit_config = _get_base_config()
+ elif self.config.mm_extractor_scale == "large":
+ beit_config = _get_large_config()
+ else:
+ raise AttributeError(f"model config should contain key 'mm_extractor_scale', with value 'base' or 'large'.")
+
+ self.mm_extractor = BEiT3Wrapper(beit_config)
+ if self.encoder_pretrained is not None:
+ beit_state_dict = torch.load(self.encoder_pretrained)["model"]
+ self.mm_extractor.load_state_dict(
+ beit_state_dict,
+ strict=False
+ )
+
+ for param in self.mm_extractor.parameters():
+ param.requires_grad = True
+
+ # Projection layer
+ in_dim = config.hidden_size
+ assert in_dim==beit_config.encoder_embed_dim, \
+ f"projection layer dim {in_dim} mismatch with mm_extractor dim {beit_config.encoder_embed_dim}"
+ out_dim = config.out_dim
+ text_fc = [
+ nn.Linear(in_dim, in_dim),
+ nn.ReLU(),
+ nn.Linear(in_dim, out_dim)
+ ]
+ self.text_hidden_fcs = nn.ModuleList([nn.Sequential(*text_fc)])
+ self.text_hidden_fcs.train()
+ for param in self.text_hidden_fcs.parameters():
+ param.requires_grad = True
+
+
+ def postprocess_masks(self, masks: torch.Tensor, orig_hw) -> torch.Tensor:
+ """
+ Perform PostProcessing on output masks.
+ """
+ masks = masks.float()
+ masks = F.interpolate(masks, orig_hw, mode="bilinear", align_corners=False)
+ return masks
+
+ # def forward(
+ # self,
+ # images: torch.FloatTensor,
+ # images_evf: torch.FloatTensor,
+ # input_ids: torch.LongTensor,
+ # attention_masks: torch.LongTensor,
+ # offset: torch.LongTensor,
+ # masks_list: List[torch.FloatTensor],
+ # label_list: List[torch.Tensor],
+ # resize_list: List[tuple],
+ # inference: bool = False,
+ # **kwargs,
+ # ):
+ # # image_embeddings = self.get_visual_embs(images)
+ # backbone_out = self.visual_model.forward_image(images)
+ # # dict_keys(['vision_features', 'vision_pos_enc', 'backbone_fpn'])
+ # _, image_embeddings, _, _ = self.visual_model._prepare_backbone_features(backbone_out)
+ # image_embeddings = [_.to(images.dtype) for _ in image_embeddings]
+ # batch_size = images.shape[0]
+ # if self.visual_model.directly_add_no_mem_embed:
+ # image_embeddings[-1] = image_embeddings[-1] + self.visual_model.no_mem_embed
+
+ # feats = [
+ # feat.permute(1, 2, 0).view(batch_size, -1, *feat_size)
+ # for feat, feat_size in zip(image_embeddings[::-1], self._bb_feat_sizes[::-1])
+ # ][::-1]
+ # _features = {"image_embed": feats[-1], "high_res_feats": feats[:-1]}
+
+
+ # assert batch_size == len(offset) - 1
+
+ # images_evf_list = []
+ # for i in range(len(offset) - 1):
+ # start_i, end_i = offset[i], offset[i + 1]
+ # images_evf_i = (
+ # images_evf[i]
+ # .unsqueeze(0)
+ # .expand(end_i - start_i, -1, -1, -1)
+ # .contiguous()
+ # )
+ # images_evf_list.append(images_evf_i)
+ # images_evf = torch.cat(images_evf_list, dim=0)
+
+ # multimask_output = False
+ # output = self.mm_extractor.beit3(
+ # visual_tokens=images_evf,
+ # textual_tokens=input_ids,
+ # text_padding_position=~attention_masks
+ # )
+
+ # feat = output["encoder_out"][:, :1, ...]
+
+ # feat = self.text_hidden_fcs[0](feat)
+ # feat = torch.split(feat, [offset[i+1] - offset[i] for i in range(len(offset)-1)])
+
+ # pred_masks = []
+
+ # for i in range(len(feat)):
+ # (
+ # sparse_embeddings,
+ # dense_embeddings,
+ # ) = self.visual_model.sam_prompt_encoder(
+ # points=None,
+ # boxes=None,
+ # masks=None,
+ # text_embeds=feat[i],
+ # )
+ # sparse_embeddings = sparse_embeddings.to(feat[i].dtype)
+ # high_res_features = [
+ # feat_level[i].unsqueeze(0)
+ # for feat_level in _features["high_res_feats"]
+ # ]
+ # low_res_masks, iou_predictions, _, _ = self.visual_model.sam_mask_decoder(
+ # image_embeddings=_features["image_embed"][i].unsqueeze(0),
+ # image_pe=self.visual_model.sam_prompt_encoder.get_dense_pe(),
+ # sparse_prompt_embeddings=sparse_embeddings,
+ # dense_prompt_embeddings=dense_embeddings,
+ # multimask_output=multimask_output,
+ # repeat_image = True,
+ # high_res_features=high_res_features,
+ # )
+
+ # if multimask_output:
+ # sorted_ids = torch.argsort(iou_predictions, dim=-1, descending=True)
+ # low_res_masks = torch.take_along_dim(low_res_masks, sorted_ids[..., None, None], dim=1)[:, :1]
+
+ # pred_mask = self.postprocess_masks(
+ # low_res_masks,
+ # orig_hw=label_list[i].shape,
+ # )
+ # pred_masks.append(pred_mask[:, 0])
+
+ # gt_masks = masks_list
+
+ # if inference:
+ # return {
+ # "pred_masks": pred_masks,
+ # "gt_masks": gt_masks,
+ # }
+
+ # mask_bce_loss = 0
+ # mask_dice_loss = 0
+ # num_masks = 0
+ # for batch_idx in range(len(pred_masks)):
+ # gt_mask = gt_masks[batch_idx]
+ # pred_mask = pred_masks[batch_idx]
+
+ # assert (
+ # gt_mask.shape[0] == pred_mask.shape[0]
+ # ), "gt_mask.shape: {}, pred_mask.shape: {}".format(
+ # gt_mask.shape, pred_mask.shape
+ # )
+ # mask_bce_loss += (
+ # sigmoid_ce_loss(pred_mask, gt_mask, num_masks=gt_mask.shape[0])
+ # * gt_mask.shape[0]
+ # )
+ # mask_dice_loss += (
+ # dice_loss(pred_mask, gt_mask, num_masks=gt_mask.shape[0])
+ # * gt_mask.shape[0]
+ # )
+ # num_masks += gt_mask.shape[0]
+
+ # mask_bce_loss = self.bce_loss_weight * mask_bce_loss / (num_masks + 1e-8)
+ # mask_dice_loss = self.dice_loss_weight * mask_dice_loss / (num_masks + 1e-8)
+ # mask_loss = mask_bce_loss + mask_dice_loss
+
+ # loss = mask_loss
+
+ # return {
+ # "loss": loss,
+ # "mask_bce_loss": mask_bce_loss,
+ # "mask_dice_loss": mask_dice_loss,
+ # "mask_loss": mask_loss,
+ # }
+
+ def inference(
+ self,
+ video_path,
+ images_evf,
+ input_ids,
+ # original_size_list,
+ multimask_output=False,
+ ):
+ predictor = self.visual_model
+ inference_state = predictor.init_state(video_path=video_path)
+ predictor.reset_state(inference_state)
+
+
+ multimask_output = multimask_output
+
+ output = self.mm_extractor.beit3(visual_tokens=images_evf, textual_tokens=input_ids, text_padding_position=torch.zeros_like(input_ids))
+
+ feat = output["encoder_out"][:, :1, ...]
+ feat = self.text_hidden_fcs[0](feat)
+
+ ann_frame_idx = 0 # the frame index we interact with
+ ann_obj_id = 1 # give a unique id to each object we interact with (it can be any integers)
+
+ _, out_obj_ids, out_mask_logits = predictor.add_new_text(
+ inference_state=inference_state,
+ frame_idx=ann_frame_idx,
+ obj_id=ann_obj_id,
+ text=feat
+ )
+
+ # run propagation throughout the video and collect the results in a dict
+ video_segments = {} # video_segments contains the per-frame segmentation results
+ for out_frame_idx, out_obj_ids, out_mask_logits in predictor.propagate_in_video(inference_state):
+ video_segments[out_frame_idx] = {
+ out_obj_id: (out_mask_logits[i] > 0.0).cpu().numpy()
+ for i, out_obj_id in enumerate(out_obj_ids)
+ }
+
+ return video_segments
+
+
+AutoConfig.register("evf", EvfConfig)
+AutoModelForCausalLM.register(EvfConfig, EvfSam2Model)
\ No newline at end of file
diff --git a/py/evf_sam/model/segment_anything/__init__.py b/py/evf_sam/model/segment_anything/__init__.py
new file mode 100644
index 0000000..e66218b
--- /dev/null
+++ b/py/evf_sam/model/segment_anything/__init__.py
@@ -0,0 +1,10 @@
+# Copyright (c) Meta Platforms, Inc. and affiliates.
+# All rights reserved.
+
+# This source code is licensed under the license found in the
+# LICENSE file in the root directory of this source tree.
+
+from .automatic_mask_generator import SamAutomaticMaskGenerator
+from .build_sam import (build_sam, build_sam_vit_b, build_sam_vit_h,
+ build_sam_vit_l, sam_model_registry)
+from .predictor import SamPredictor
diff --git a/py/evf_sam/model/segment_anything/automatic_mask_generator.py b/py/evf_sam/model/segment_anything/automatic_mask_generator.py
new file mode 100644
index 0000000..aa4bc4f
--- /dev/null
+++ b/py/evf_sam/model/segment_anything/automatic_mask_generator.py
@@ -0,0 +1,372 @@
+# Copyright (c) Meta Platforms, Inc. and affiliates.
+# All rights reserved.
+
+# This source code is licensed under the license found in the
+# LICENSE file in the root directory of this source tree.
+
+from typing import Any, Dict, List, Optional, Tuple
+
+import numpy as np
+import torch
+from torchvision.ops.boxes import batched_nms, box_area # type: ignore
+
+from .modeling import Sam
+from .predictor import SamPredictor
+from .utils.amg import (MaskData, area_from_rle, batch_iterator,
+ batched_mask_to_box, box_xyxy_to_xywh,
+ build_all_layer_point_grids, calculate_stability_score,
+ coco_encode_rle, generate_crop_boxes,
+ is_box_near_crop_edge, mask_to_rle_pytorch,
+ remove_small_regions, rle_to_mask, uncrop_boxes_xyxy,
+ uncrop_masks, uncrop_points)
+
+
+class SamAutomaticMaskGenerator:
+ def __init__(
+ self,
+ model: Sam,
+ points_per_side: Optional[int] = 32,
+ points_per_batch: int = 64,
+ pred_iou_thresh: float = 0.88,
+ stability_score_thresh: float = 0.95,
+ stability_score_offset: float = 1.0,
+ box_nms_thresh: float = 0.7,
+ crop_n_layers: int = 0,
+ crop_nms_thresh: float = 0.7,
+ crop_overlap_ratio: float = 512 / 1500,
+ crop_n_points_downscale_factor: int = 1,
+ point_grids: Optional[List[np.ndarray]] = None,
+ min_mask_region_area: int = 0,
+ output_mode: str = "binary_mask",
+ ) -> None:
+ """
+ Using a SAM model, generates masks for the entire image.
+ Generates a grid of point prompts over the image, then filters
+ low quality and duplicate masks. The default settings are chosen
+ for SAM with a ViT-H backbone.
+
+ Arguments:
+ model (Sam): The SAM model to use for mask prediction.
+ points_per_side (int or None): The number of points to be sampled
+ along one side of the image. The total number of points is
+ points_per_side**2. If None, 'point_grids' must provide explicit
+ point sampling.
+ points_per_batch (int): Sets the number of points run simultaneously
+ by the model. Higher numbers may be faster but use more GPU memory.
+ pred_iou_thresh (float): A filtering threshold in [0,1], using the
+ model's predicted mask quality.
+ stability_score_thresh (float): A filtering threshold in [0,1], using
+ the stability of the mask under changes to the cutoff used to binarize
+ the model's mask predictions.
+ stability_score_offset (float): The amount to shift the cutoff when
+ calculated the stability score.
+ box_nms_thresh (float): The box IoU cutoff used by non-maximal
+ suppression to filter duplicate masks.
+ crop_n_layers (int): If >0, mask prediction will be run again on
+ crops of the image. Sets the number of layers to run, where each
+ layer has 2**i_layer number of image crops.
+ crop_nms_thresh (float): The box IoU cutoff used by non-maximal
+ suppression to filter duplicate masks between different crops.
+ crop_overlap_ratio (float): Sets the degree to which crops overlap.
+ In the first crop layer, crops will overlap by this fraction of
+ the image length. Later layers with more crops scale down this overlap.
+ crop_n_points_downscale_factor (int): The number of points-per-side
+ sampled in layer n is scaled down by crop_n_points_downscale_factor**n.
+ point_grids (list(np.ndarray) or None): A list over explicit grids
+ of points used for sampling, normalized to [0,1]. The nth grid in the
+ list is used in the nth crop layer. Exclusive with points_per_side.
+ min_mask_region_area (int): If >0, postprocessing will be applied
+ to remove disconnected regions and holes in masks with area smaller
+ than min_mask_region_area. Requires opencv.
+ output_mode (str): The form masks are returned in. Can be 'binary_mask',
+ 'uncompressed_rle', or 'coco_rle'. 'coco_rle' requires pycocotools.
+ For large resolutions, 'binary_mask' may consume large amounts of
+ memory.
+ """
+
+ assert (points_per_side is None) != (
+ point_grids is None
+ ), "Exactly one of points_per_side or point_grid must be provided."
+ if points_per_side is not None:
+ self.point_grids = build_all_layer_point_grids(
+ points_per_side,
+ crop_n_layers,
+ crop_n_points_downscale_factor,
+ )
+ elif point_grids is not None:
+ self.point_grids = point_grids
+ else:
+ raise ValueError("Can't have both points_per_side and point_grid be None.")
+
+ assert output_mode in [
+ "binary_mask",
+ "uncompressed_rle",
+ "coco_rle",
+ ], f"Unknown output_mode {output_mode}."
+ if output_mode == "coco_rle":
+ from pycocotools import \
+ mask as mask_utils # type: ignore # noqa: F401
+
+ if min_mask_region_area > 0:
+ import cv2 # type: ignore # noqa: F401
+
+ self.predictor = SamPredictor(model)
+ self.points_per_batch = points_per_batch
+ self.pred_iou_thresh = pred_iou_thresh
+ self.stability_score_thresh = stability_score_thresh
+ self.stability_score_offset = stability_score_offset
+ self.box_nms_thresh = box_nms_thresh
+ self.crop_n_layers = crop_n_layers
+ self.crop_nms_thresh = crop_nms_thresh
+ self.crop_overlap_ratio = crop_overlap_ratio
+ self.crop_n_points_downscale_factor = crop_n_points_downscale_factor
+ self.min_mask_region_area = min_mask_region_area
+ self.output_mode = output_mode
+
+ @torch.no_grad()
+ def generate(self, image: np.ndarray) -> List[Dict[str, Any]]:
+ """
+ Generates masks for the given image.
+
+ Arguments:
+ image (np.ndarray): The image to generate masks for, in HWC uint8 format.
+
+ Returns:
+ list(dict(str, any)): A list over records for masks. Each record is
+ a dict containing the following keys:
+ segmentation (dict(str, any) or np.ndarray): The mask. If
+ output_mode='binary_mask', is an array of shape HW. Otherwise,
+ is a dictionary containing the RLE.
+ bbox (list(float)): The box around the mask, in XYWH format.
+ area (int): The area in pixels of the mask.
+ predicted_iou (float): The model's own prediction of the mask's
+ quality. This is filtered by the pred_iou_thresh parameter.
+ point_coords (list(list(float))): The point coordinates input
+ to the model to generate this mask.
+ stability_score (float): A measure of the mask's quality. This
+ is filtered on using the stability_score_thresh parameter.
+ crop_box (list(float)): The crop of the image used to generate
+ the mask, given in XYWH format.
+ """
+
+ # Generate masks
+ mask_data = self._generate_masks(image)
+
+ # Filter small disconnected regions and holes in masks
+ if self.min_mask_region_area > 0:
+ mask_data = self.postprocess_small_regions(
+ mask_data,
+ self.min_mask_region_area,
+ max(self.box_nms_thresh, self.crop_nms_thresh),
+ )
+
+ # Encode masks
+ if self.output_mode == "coco_rle":
+ mask_data["segmentations"] = [
+ coco_encode_rle(rle) for rle in mask_data["rles"]
+ ]
+ elif self.output_mode == "binary_mask":
+ mask_data["segmentations"] = [rle_to_mask(rle) for rle in mask_data["rles"]]
+ else:
+ mask_data["segmentations"] = mask_data["rles"]
+
+ # Write mask records
+ curr_anns = []
+ for idx in range(len(mask_data["segmentations"])):
+ ann = {
+ "segmentation": mask_data["segmentations"][idx],
+ "area": area_from_rle(mask_data["rles"][idx]),
+ "bbox": box_xyxy_to_xywh(mask_data["boxes"][idx]).tolist(),
+ "predicted_iou": mask_data["iou_preds"][idx].item(),
+ "point_coords": [mask_data["points"][idx].tolist()],
+ "stability_score": mask_data["stability_score"][idx].item(),
+ "crop_box": box_xyxy_to_xywh(mask_data["crop_boxes"][idx]).tolist(),
+ }
+ curr_anns.append(ann)
+
+ return curr_anns
+
+ def _generate_masks(self, image: np.ndarray) -> MaskData:
+ orig_size = image.shape[:2]
+ crop_boxes, layer_idxs = generate_crop_boxes(
+ orig_size, self.crop_n_layers, self.crop_overlap_ratio
+ )
+
+ # Iterate over image crops
+ data = MaskData()
+ for crop_box, layer_idx in zip(crop_boxes, layer_idxs):
+ crop_data = self._process_crop(image, crop_box, layer_idx, orig_size)
+ data.cat(crop_data)
+
+ # Remove duplicate masks between crops
+ if len(crop_boxes) > 1:
+ # Prefer masks from smaller crops
+ scores = 1 / box_area(data["crop_boxes"])
+ scores = scores.to(data["boxes"].device)
+ keep_by_nms = batched_nms(
+ data["boxes"].float(),
+ scores,
+ torch.zeros_like(data["boxes"][:, 0]), # categories
+ iou_threshold=self.crop_nms_thresh,
+ )
+ data.filter(keep_by_nms)
+
+ data.to_numpy()
+ return data
+
+ def _process_crop(
+ self,
+ image: np.ndarray,
+ crop_box: List[int],
+ crop_layer_idx: int,
+ orig_size: Tuple[int, ...],
+ ) -> MaskData:
+ # Crop the image and calculate embeddings
+ x0, y0, x1, y1 = crop_box
+ cropped_im = image[y0:y1, x0:x1, :]
+ cropped_im_size = cropped_im.shape[:2]
+ self.predictor.set_image(cropped_im)
+
+ # Get points for this crop
+ points_scale = np.array(cropped_im_size)[None, ::-1]
+ points_for_image = self.point_grids[crop_layer_idx] * points_scale
+
+ # Generate masks for this crop in batches
+ data = MaskData()
+ for (points,) in batch_iterator(self.points_per_batch, points_for_image):
+ batch_data = self._process_batch(
+ points, cropped_im_size, crop_box, orig_size
+ )
+ data.cat(batch_data)
+ del batch_data
+ self.predictor.reset_image()
+
+ # Remove duplicates within this crop.
+ keep_by_nms = batched_nms(
+ data["boxes"].float(),
+ data["iou_preds"],
+ torch.zeros_like(data["boxes"][:, 0]), # categories
+ iou_threshold=self.box_nms_thresh,
+ )
+ data.filter(keep_by_nms)
+
+ # Return to the original image frame
+ data["boxes"] = uncrop_boxes_xyxy(data["boxes"], crop_box)
+ data["points"] = uncrop_points(data["points"], crop_box)
+ data["crop_boxes"] = torch.tensor([crop_box for _ in range(len(data["rles"]))])
+
+ return data
+
+ def _process_batch(
+ self,
+ points: np.ndarray,
+ im_size: Tuple[int, ...],
+ crop_box: List[int],
+ orig_size: Tuple[int, ...],
+ ) -> MaskData:
+ orig_h, orig_w = orig_size
+
+ # Run model on this batch
+ transformed_points = self.predictor.transform.apply_coords(points, im_size)
+ in_points = torch.as_tensor(transformed_points, device=self.predictor.device)
+ in_labels = torch.ones(
+ in_points.shape[0], dtype=torch.int, device=in_points.device
+ )
+ masks, iou_preds, _ = self.predictor.predict_torch(
+ in_points[:, None, :],
+ in_labels[:, None],
+ multimask_output=True,
+ return_logits=True,
+ )
+
+ # Serialize predictions and store in MaskData
+ data = MaskData(
+ masks=masks.flatten(0, 1),
+ iou_preds=iou_preds.flatten(0, 1),
+ points=torch.as_tensor(points.repeat(masks.shape[1], axis=0)),
+ )
+ del masks
+
+ # Filter by predicted IoU
+ if self.pred_iou_thresh > 0.0:
+ keep_mask = data["iou_preds"] > self.pred_iou_thresh
+ data.filter(keep_mask)
+
+ # Calculate stability score
+ data["stability_score"] = calculate_stability_score(
+ data["masks"],
+ self.predictor.model.mask_threshold,
+ self.stability_score_offset,
+ )
+ if self.stability_score_thresh > 0.0:
+ keep_mask = data["stability_score"] >= self.stability_score_thresh
+ data.filter(keep_mask)
+
+ # Threshold masks and calculate boxes
+ data["masks"] = data["masks"] > self.predictor.model.mask_threshold
+ data["boxes"] = batched_mask_to_box(data["masks"])
+
+ # Filter boxes that touch crop boundaries
+ keep_mask = ~is_box_near_crop_edge(
+ data["boxes"], crop_box, [0, 0, orig_w, orig_h]
+ )
+ if not torch.all(keep_mask):
+ data.filter(keep_mask)
+
+ # Compress to RLE
+ data["masks"] = uncrop_masks(data["masks"], crop_box, orig_h, orig_w)
+ data["rles"] = mask_to_rle_pytorch(data["masks"])
+ del data["masks"]
+
+ return data
+
+ @staticmethod
+ def postprocess_small_regions(
+ mask_data: MaskData, min_area: int, nms_thresh: float
+ ) -> MaskData:
+ """
+ Removes small disconnected regions and holes in masks, then reruns
+ box NMS to remove any new duplicates.
+
+ Edits mask_data in place.
+
+ Requires open-cv as a dependency.
+ """
+ if len(mask_data["rles"]) == 0:
+ return mask_data
+
+ # Filter small disconnected regions and holes
+ new_masks = []
+ scores = []
+ for rle in mask_data["rles"]:
+ mask = rle_to_mask(rle)
+
+ mask, changed = remove_small_regions(mask, min_area, mode="holes")
+ unchanged = not changed
+ mask, changed = remove_small_regions(mask, min_area, mode="islands")
+ unchanged = unchanged and not changed
+
+ new_masks.append(torch.as_tensor(mask).unsqueeze(0))
+ # Give score=0 to changed masks and score=1 to unchanged masks
+ # so NMS will prefer ones that didn't need postprocessing
+ scores.append(float(unchanged))
+
+ # Recalculate boxes and remove any new duplicates
+ masks = torch.cat(new_masks, dim=0)
+ boxes = batched_mask_to_box(masks)
+ keep_by_nms = batched_nms(
+ boxes.float(),
+ torch.as_tensor(scores),
+ torch.zeros_like(boxes[:, 0]), # categories
+ iou_threshold=nms_thresh,
+ )
+
+ # Only recalculate RLEs for masks that have changed
+ for i_mask in keep_by_nms:
+ if scores[i_mask] == 0.0:
+ mask_torch = masks[i_mask].unsqueeze(0)
+ mask_data["rles"][i_mask] = mask_to_rle_pytorch(mask_torch)[0]
+ mask_data["boxes"][i_mask] = boxes[i_mask] # update res directly
+ mask_data.filter(keep_by_nms)
+
+ return mask_data
diff --git a/py/evf_sam/model/segment_anything/build_sam.py b/py/evf_sam/model/segment_anything/build_sam.py
new file mode 100644
index 0000000..788d25a
--- /dev/null
+++ b/py/evf_sam/model/segment_anything/build_sam.py
@@ -0,0 +1,108 @@
+# Copyright (c) Meta Platforms, Inc. and affiliates.
+# All rights reserved.
+
+# This source code is licensed under the license found in the
+# LICENSE file in the root directory of this source tree.
+
+from functools import partial
+
+import torch
+
+from .modeling import (ImageEncoderViT, MaskDecoder, PromptEncoder, Sam,
+ TwoWayTransformer)
+
+
+def build_sam_vit_h(checkpoint=None):
+ return _build_sam(
+ encoder_embed_dim=1280,
+ encoder_depth=32,
+ encoder_num_heads=16,
+ encoder_global_attn_indexes=[7, 15, 23, 31],
+ checkpoint=checkpoint,
+ )
+
+
+build_sam = build_sam_vit_h
+
+
+def build_sam_vit_l(checkpoint=None):
+ return _build_sam(
+ encoder_embed_dim=1024,
+ encoder_depth=24,
+ encoder_num_heads=16,
+ encoder_global_attn_indexes=[5, 11, 17, 23],
+ checkpoint=checkpoint,
+ )
+
+
+def build_sam_vit_b(checkpoint=None):
+ return _build_sam(
+ encoder_embed_dim=768,
+ encoder_depth=12,
+ encoder_num_heads=12,
+ encoder_global_attn_indexes=[2, 5, 8, 11],
+ checkpoint=checkpoint,
+ )
+
+
+sam_model_registry = {
+ "default": build_sam_vit_h,
+ "vit_h": build_sam_vit_h,
+ "vit_l": build_sam_vit_l,
+ "vit_b": build_sam_vit_b,
+}
+
+
+def _build_sam(
+ encoder_embed_dim,
+ encoder_depth,
+ encoder_num_heads,
+ encoder_global_attn_indexes,
+ checkpoint=None,
+):
+ prompt_embed_dim = 256
+ image_size = 1024
+ vit_patch_size = 16
+ image_embedding_size = image_size // vit_patch_size
+ sam = Sam(
+ image_encoder=ImageEncoderViT(
+ depth=encoder_depth,
+ embed_dim=encoder_embed_dim,
+ img_size=image_size,
+ mlp_ratio=4,
+ norm_layer=partial(torch.nn.LayerNorm, eps=1e-6),
+ num_heads=encoder_num_heads,
+ patch_size=vit_patch_size,
+ qkv_bias=True,
+ use_rel_pos=True,
+ global_attn_indexes=encoder_global_attn_indexes,
+ window_size=14,
+ out_chans=prompt_embed_dim,
+ ),
+ prompt_encoder=PromptEncoder(
+ embed_dim=prompt_embed_dim,
+ image_embedding_size=(image_embedding_size, image_embedding_size),
+ input_image_size=(image_size, image_size),
+ mask_in_chans=16,
+ ),
+ mask_decoder=MaskDecoder(
+ num_multimask_outputs=3,
+ transformer=TwoWayTransformer(
+ depth=2,
+ embedding_dim=prompt_embed_dim,
+ mlp_dim=2048,
+ num_heads=8,
+ ),
+ transformer_dim=prompt_embed_dim,
+ iou_head_depth=3,
+ iou_head_hidden_dim=256,
+ ),
+ pixel_mean=[123.675, 116.28, 103.53],
+ pixel_std=[58.395, 57.12, 57.375],
+ )
+ sam.eval()
+ if checkpoint is not None:
+ with open(checkpoint, "rb") as f:
+ state_dict = torch.load(f)
+ sam.load_state_dict(state_dict, strict=False)
+ return sam
diff --git a/py/evf_sam/model/segment_anything/modeling/__init__.py b/py/evf_sam/model/segment_anything/modeling/__init__.py
new file mode 100644
index 0000000..088af38
--- /dev/null
+++ b/py/evf_sam/model/segment_anything/modeling/__init__.py
@@ -0,0 +1,11 @@
+# Copyright (c) Meta Platforms, Inc. and affiliates.
+# All rights reserved.
+
+# This source code is licensed under the license found in the
+# LICENSE file in the root directory of this source tree.
+
+from .image_encoder import ImageEncoderViT
+from .mask_decoder import MaskDecoder
+from .prompt_encoder import PromptEncoder
+from .sam import Sam
+from .transformer import TwoWayTransformer
diff --git a/py/evf_sam/model/segment_anything/modeling/common.py b/py/evf_sam/model/segment_anything/modeling/common.py
new file mode 100644
index 0000000..e872781
--- /dev/null
+++ b/py/evf_sam/model/segment_anything/modeling/common.py
@@ -0,0 +1,43 @@
+# Copyright (c) Meta Platforms, Inc. and affiliates.
+# All rights reserved.
+
+# This source code is licensed under the license found in the
+# LICENSE file in the root directory of this source tree.
+
+from typing import Type
+
+import torch
+import torch.nn as nn
+
+
+class MLPBlock(nn.Module):
+ def __init__(
+ self,
+ embedding_dim: int,
+ mlp_dim: int,
+ act: Type[nn.Module] = nn.GELU,
+ ) -> None:
+ super().__init__()
+ self.lin1 = nn.Linear(embedding_dim, mlp_dim)
+ self.lin2 = nn.Linear(mlp_dim, embedding_dim)
+ self.act = act()
+
+ def forward(self, x: torch.Tensor) -> torch.Tensor:
+ return self.lin2(self.act(self.lin1(x)))
+
+
+# From https://github.com/facebookresearch/detectron2/blob/main/detectron2/layers/batch_norm.py # noqa
+# Itself from https://github.com/facebookresearch/ConvNeXt/blob/d1fa8f6fef0a165b27399986cc2bdacc92777e40/models/convnext.py#L119 # noqa
+class LayerNorm2d(nn.Module):
+ def __init__(self, num_channels: int, eps: float = 1e-6) -> None:
+ super().__init__()
+ self.weight = nn.Parameter(torch.ones(num_channels))
+ self.bias = nn.Parameter(torch.zeros(num_channels))
+ self.eps = eps
+
+ def forward(self, x: torch.Tensor) -> torch.Tensor:
+ u = x.mean(1, keepdim=True)
+ s = (x - u).pow(2).mean(1, keepdim=True)
+ x = (x - u) / torch.sqrt(s + self.eps)
+ x = self.weight[:, None, None] * x + self.bias[:, None, None]
+ return x
diff --git a/py/evf_sam/model/segment_anything/modeling/image_encoder.py b/py/evf_sam/model/segment_anything/modeling/image_encoder.py
new file mode 100644
index 0000000..b472a3d
--- /dev/null
+++ b/py/evf_sam/model/segment_anything/modeling/image_encoder.py
@@ -0,0 +1,426 @@
+# Copyright (c) Meta Platforms, Inc. and affiliates.
+# All rights reserved.
+
+# This source code is licensed under the license found in the
+# LICENSE file in the root directory of this source tree.
+
+from typing import Optional, Tuple, Type
+
+import torch
+import torch.nn as nn
+import torch.nn.functional as F
+
+from .common import LayerNorm2d, MLPBlock
+
+
+# This class and its supporting functions below lightly adapted from the ViTDet backbone available at: https://github.com/facebookresearch/detectron2/blob/main/detectron2/modeling/backbone/vit.py # noqa
+class ImageEncoderViT(nn.Module):
+ def __init__(
+ self,
+ img_size: int = 1024,
+ patch_size: int = 16,
+ in_chans: int = 3,
+ embed_dim: int = 768,
+ depth: int = 12,
+ num_heads: int = 12,
+ mlp_ratio: float = 4.0,
+ out_chans: int = 256,
+ qkv_bias: bool = True,
+ norm_layer: Type[nn.Module] = nn.LayerNorm,
+ act_layer: Type[nn.Module] = nn.GELU,
+ use_abs_pos: bool = True,
+ use_rel_pos: bool = False,
+ rel_pos_zero_init: bool = True,
+ window_size: int = 0,
+ global_attn_indexes: Tuple[int, ...] = (),
+ ) -> None:
+ """
+ Args:
+ img_size (int): Input image size.
+ patch_size (int): Patch size.
+ in_chans (int): Number of input image channels.
+ embed_dim (int): Patch embedding dimension.
+ depth (int): Depth of ViT.
+ num_heads (int): Number of attention heads in each ViT block.
+ mlp_ratio (float): Ratio of mlp hidden dim to embedding dim.
+ qkv_bias (bool): If True, add a learnable bias to query, key, value.
+ norm_layer (nn.Module): Normalization layer.
+ act_layer (nn.Module): Activation layer.
+ use_abs_pos (bool): If True, use absolute positional embeddings.
+ use_rel_pos (bool): If True, add relative positional embeddings to the attention map.
+ rel_pos_zero_init (bool): If True, zero initialize relative positional parameters.
+ window_size (int): Window size for window attention blocks.
+ global_attn_indexes (list): Indexes for blocks using global attention.
+ """
+ super().__init__()
+ self.img_size = img_size
+ self.embed_dim = embed_dim
+ self.out_chans = out_chans
+
+ self.patch_embed = PatchEmbed(
+ kernel_size=(patch_size, patch_size),
+ stride=(patch_size, patch_size),
+ in_chans=in_chans,
+ embed_dim=embed_dim,
+ )
+
+ self.pos_embed: Optional[nn.Parameter] = None
+ if use_abs_pos:
+ # Initialize absolute positional embedding with pretrain image size.
+ self.pos_embed = nn.Parameter(
+ torch.zeros(
+ 1, img_size // patch_size, img_size // patch_size, embed_dim
+ )
+ )
+
+ self.blocks = nn.ModuleList()
+ for i in range(depth):
+ block = Block(
+ dim=embed_dim,
+ num_heads=num_heads,
+ mlp_ratio=mlp_ratio,
+ qkv_bias=qkv_bias,
+ norm_layer=norm_layer,
+ act_layer=act_layer,
+ use_rel_pos=use_rel_pos,
+ rel_pos_zero_init=rel_pos_zero_init,
+ window_size=window_size if i not in global_attn_indexes else 0,
+ input_size=(img_size // patch_size, img_size // patch_size),
+ )
+ self.blocks.append(block)
+
+ self.neck = nn.Sequential(
+ nn.Conv2d(
+ embed_dim,
+ out_chans,
+ kernel_size=1,
+ bias=False,
+ ),
+ LayerNorm2d(out_chans),
+ nn.Conv2d(
+ out_chans,
+ out_chans,
+ kernel_size=3,
+ padding=1,
+ bias=False,
+ ),
+ LayerNorm2d(out_chans),
+ )
+
+ def forward(self, x: torch.Tensor) -> torch.Tensor:
+ x = self.patch_embed(x)
+ if self.pos_embed is not None:
+ x = x + self.pos_embed
+
+ for blk in self.blocks:
+ x = blk(x)
+
+ dtype = x.dtype
+ if dtype == torch.float16: # prevent overflow
+ with torch.autocast(device_type="cuda", dtype=torch.float32):
+ x = self.neck(x.permute(0, 3, 1, 2))
+ x = x.to(dtype)
+ else:
+ x = self.neck(x.permute(0, 3, 1, 2))
+ return x
+
+
+class Block(nn.Module):
+ """Transformer blocks with support of window attention and residual propagation blocks"""
+
+ def __init__(
+ self,
+ dim: int,
+ num_heads: int,
+ mlp_ratio: float = 4.0,
+ qkv_bias: bool = True,
+ norm_layer: Type[nn.Module] = nn.LayerNorm,
+ act_layer: Type[nn.Module] = nn.GELU,
+ use_rel_pos: bool = False,
+ rel_pos_zero_init: bool = True,
+ window_size: int = 0,
+ input_size: Optional[Tuple[int, int]] = None,
+ ) -> None:
+ """
+ Args:
+ dim (int): Number of input channels.
+ num_heads (int): Number of attention heads in each ViT block.
+ mlp_ratio (float): Ratio of mlp hidden dim to embedding dim.
+ qkv_bias (bool): If True, add a learnable bias to query, key, value.
+ norm_layer (nn.Module): Normalization layer.
+ act_layer (nn.Module): Activation layer.
+ use_rel_pos (bool): If True, add relative positional embeddings to the attention map.
+ rel_pos_zero_init (bool): If True, zero initialize relative positional parameters.
+ window_size (int): Window size for window attention blocks. If it equals 0, then
+ use global attention.
+ input_size (tuple(int, int) or None): Input resolution for calculating the relative
+ positional parameter size.
+ """
+ super().__init__()
+ self.norm1 = norm_layer(dim)
+ self.attn = Attention(
+ dim,
+ num_heads=num_heads,
+ qkv_bias=qkv_bias,
+ use_rel_pos=use_rel_pos,
+ rel_pos_zero_init=rel_pos_zero_init,
+ input_size=input_size if window_size == 0 else (window_size, window_size),
+ )
+
+ self.norm2 = norm_layer(dim)
+ self.mlp = MLPBlock(
+ embedding_dim=dim, mlp_dim=int(dim * mlp_ratio), act=act_layer
+ )
+
+ self.window_size = window_size
+
+ def forward(self, x: torch.Tensor) -> torch.Tensor:
+ shortcut = x
+ x = self.norm1(x)
+ # Window partition
+ if self.window_size > 0:
+ H, W = x.shape[1], x.shape[2]
+ x, pad_hw = window_partition(x, self.window_size)
+
+ x = self.attn(x)
+ # Reverse window partition
+ if self.window_size > 0:
+ x = window_unpartition(x, self.window_size, pad_hw, (H, W))
+
+ x = shortcut + x
+ x = x + self.mlp(self.norm2(x))
+
+ return x
+
+
+class Attention(nn.Module):
+ """Multi-head Attention block with relative position embeddings."""
+
+ def __init__(
+ self,
+ dim: int,
+ num_heads: int = 8,
+ qkv_bias: bool = True,
+ use_rel_pos: bool = False,
+ rel_pos_zero_init: bool = True,
+ input_size: Optional[Tuple[int, int]] = None,
+ ) -> None:
+ """
+ Args:
+ dim (int): Number of input channels.
+ num_heads (int): Number of attention heads.
+ qkv_bias (bool): If True, add a learnable bias to query, key, value.
+ rel_pos (bool): If True, add relative positional embeddings to the attention map.
+ rel_pos_zero_init (bool): If True, zero initialize relative positional parameters.
+ input_size (tuple(int, int) or None): Input resolution for calculating the relative
+ positional parameter size.
+ """
+ super().__init__()
+ self.num_heads = num_heads
+ head_dim = dim // num_heads
+ self.scale = head_dim**-0.5
+
+ self.qkv = nn.Linear(dim, dim * 3, bias=qkv_bias)
+ self.proj = nn.Linear(dim, dim)
+
+ self.use_rel_pos = use_rel_pos
+ if self.use_rel_pos:
+ assert (
+ input_size is not None
+ ), "Input size must be provided if using relative positional encoding."
+ # initialize relative positional embeddings
+ self.rel_pos_h = nn.Parameter(torch.zeros(2 * input_size[0] - 1, head_dim))
+ self.rel_pos_w = nn.Parameter(torch.zeros(2 * input_size[1] - 1, head_dim))
+
+ def forward(self, x: torch.Tensor) -> torch.Tensor:
+ B, H, W, _ = x.shape
+ # qkv with shape (3, B, nHead, H * W, C)
+ qkv = (
+ self.qkv(x).reshape(B, H * W, 3, self.num_heads, -1).permute(2, 0, 3, 1, 4)
+ )
+ # q, k, v with shape (B * nHead, H * W, C)
+ q, k, v = qkv.reshape(3, B * self.num_heads, H * W, -1).unbind(0)
+
+ attn = (q * self.scale) @ k.transpose(-2, -1)
+
+ if self.use_rel_pos:
+ attn = add_decomposed_rel_pos(
+ attn, q, self.rel_pos_h, self.rel_pos_w, (H, W), (H, W)
+ )
+
+ attn = attn.softmax(dim=-1)
+ x = (
+ (attn @ v)
+ .view(B, self.num_heads, H, W, -1)
+ .permute(0, 2, 3, 1, 4)
+ .reshape(B, H, W, -1)
+ )
+ x = self.proj(x)
+
+ return x
+
+
+def window_partition(
+ x: torch.Tensor, window_size: int
+) -> Tuple[torch.Tensor, Tuple[int, int]]:
+ """
+ Partition into non-overlapping windows with padding if needed.
+ Args:
+ x (tensor): input tokens with [B, H, W, C].
+ window_size (int): window size.
+
+ Returns:
+ windows: windows after partition with [B * num_windows, window_size, window_size, C].
+ (Hp, Wp): padded height and width before partition
+ """
+ B, H, W, C = x.shape
+
+ pad_h = (window_size - H % window_size) % window_size
+ pad_w = (window_size - W % window_size) % window_size
+ if pad_h > 0 or pad_w > 0:
+ x = F.pad(x, (0, 0, 0, pad_w, 0, pad_h))
+ Hp, Wp = H + pad_h, W + pad_w
+
+ x = x.view(B, Hp // window_size, window_size, Wp // window_size, window_size, C)
+ windows = (
+ x.permute(0, 1, 3, 2, 4, 5).contiguous().view(-1, window_size, window_size, C)
+ )
+ return windows, (Hp, Wp)
+
+
+def window_unpartition(
+ windows: torch.Tensor,
+ window_size: int,
+ pad_hw: Tuple[int, int],
+ hw: Tuple[int, int],
+) -> torch.Tensor:
+ """
+ Window unpartition into original sequences and removing padding.
+ Args:
+ windows (tensor): input tokens with [B * num_windows, window_size, window_size, C].
+ window_size (int): window size.
+ pad_hw (Tuple): padded height and width (Hp, Wp).
+ hw (Tuple): original height and width (H, W) before padding.
+
+ Returns:
+ x: unpartitioned sequences with [B, H, W, C].
+ """
+ Hp, Wp = pad_hw
+ H, W = hw
+ B = windows.shape[0] // (Hp * Wp // window_size // window_size)
+ x = windows.view(
+ B, Hp // window_size, Wp // window_size, window_size, window_size, -1
+ )
+ x = x.permute(0, 1, 3, 2, 4, 5).contiguous().view(B, Hp, Wp, -1)
+
+ if Hp > H or Wp > W:
+ x = x[:, :H, :W, :].contiguous()
+ return x
+
+
+def get_rel_pos(q_size: int, k_size: int, rel_pos: torch.Tensor) -> torch.Tensor:
+ """
+ Get relative positional embeddings according to the relative positions of
+ query and key sizes.
+ Args:
+ q_size (int): size of query q.
+ k_size (int): size of key k.
+ rel_pos (Tensor): relative position embeddings (L, C).
+
+ Returns:
+ Extracted positional embeddings according to relative positions.
+ """
+ max_rel_dist = int(2 * max(q_size, k_size) - 1)
+ # Interpolate rel pos if needed.
+ if rel_pos.shape[0] != max_rel_dist:
+ # Interpolate rel pos.
+ rel_pos_resized = F.interpolate(
+ rel_pos.reshape(1, rel_pos.shape[0], -1).permute(0, 2, 1),
+ size=max_rel_dist,
+ mode="linear",
+ )
+ rel_pos_resized = rel_pos_resized.reshape(-1, max_rel_dist).permute(1, 0)
+ else:
+ rel_pos_resized = rel_pos
+
+ # Scale the coords with short length if shapes for q and k are different.
+ q_coords = torch.arange(q_size)[:, None] * max(k_size / q_size, 1.0)
+ k_coords = torch.arange(k_size)[None, :] * max(q_size / k_size, 1.0)
+ relative_coords = (q_coords - k_coords) + (k_size - 1) * max(q_size / k_size, 1.0)
+
+ return rel_pos_resized[relative_coords.long()]
+
+
+def add_decomposed_rel_pos(
+ attn: torch.Tensor,
+ q: torch.Tensor,
+ rel_pos_h: torch.Tensor,
+ rel_pos_w: torch.Tensor,
+ q_size: Tuple[int, int],
+ k_size: Tuple[int, int],
+) -> torch.Tensor:
+ """
+ Calculate decomposed Relative Positional Embeddings from :paper:`mvitv2`.
+ https://github.com/facebookresearch/mvit/blob/19786631e330df9f3622e5402b4a419a263a2c80/mvit/models/attention.py # noqa B950
+ Args:
+ attn (Tensor): attention map.
+ q (Tensor): query q in the attention layer with shape (B, q_h * q_w, C).
+ rel_pos_h (Tensor): relative position embeddings (Lh, C) for height axis.
+ rel_pos_w (Tensor): relative position embeddings (Lw, C) for width axis.
+ q_size (Tuple): spatial sequence size of query q with (q_h, q_w).
+ k_size (Tuple): spatial sequence size of key k with (k_h, k_w).
+
+ Returns:
+ attn (Tensor): attention map with added relative positional embeddings.
+ """
+ q_h, q_w = q_size
+ k_h, k_w = k_size
+ Rh = get_rel_pos(q_h, k_h, rel_pos_h)
+ Rw = get_rel_pos(q_w, k_w, rel_pos_w)
+
+ B, _, dim = q.shape
+ r_q = q.reshape(B, q_h, q_w, dim)
+ rel_h = torch.einsum("bhwc,hkc->bhwk", r_q, Rh)
+ rel_w = torch.einsum("bhwc,wkc->bhwk", r_q, Rw)
+
+ attn = (
+ attn.view(B, q_h, q_w, k_h, k_w)
+ + rel_h[:, :, :, :, None]
+ + rel_w[:, :, :, None, :]
+ ).view(B, q_h * q_w, k_h * k_w)
+
+ return attn
+
+
+class PatchEmbed(nn.Module):
+ """
+ Image to Patch Embedding.
+ """
+
+ def __init__(
+ self,
+ kernel_size: Tuple[int, int] = (16, 16),
+ stride: Tuple[int, int] = (16, 16),
+ padding: Tuple[int, int] = (0, 0),
+ in_chans: int = 3,
+ embed_dim: int = 768,
+ ) -> None:
+ """
+ Args:
+ kernel_size (Tuple): kernel size of the projection layer.
+ stride (Tuple): stride of the projection layer.
+ padding (Tuple): padding size of the projection layer.
+ in_chans (int): Number of input image channels.
+ embed_dim (int): Patch embedding dimension.
+ """
+ super().__init__()
+
+ self.proj = nn.Conv2d(
+ in_chans, embed_dim, kernel_size=kernel_size, stride=stride, padding=padding
+ )
+
+ def forward(self, x: torch.Tensor) -> torch.Tensor:
+ x = self.proj(x)
+ # B C H W -> B H W C
+ x = x.permute(0, 2, 3, 1)
+ return x
diff --git a/py/evf_sam/model/segment_anything/modeling/mask_decoder.py b/py/evf_sam/model/segment_anything/modeling/mask_decoder.py
new file mode 100644
index 0000000..fb104ea
--- /dev/null
+++ b/py/evf_sam/model/segment_anything/modeling/mask_decoder.py
@@ -0,0 +1,191 @@
+# Copyright (c) Meta Platforms, Inc. and affiliates.
+# All rights reserved.
+
+# This source code is licensed under the license found in the
+# LICENSE file in the root directory of this source tree.
+
+from typing import List, Tuple, Type
+
+import torch
+from torch import nn
+from torch.nn import functional as F
+
+from .common import LayerNorm2d
+
+
+class MaskDecoder(nn.Module):
+ def __init__(
+ self,
+ *,
+ transformer_dim: int,
+ transformer: nn.Module,
+ num_multimask_outputs: int = 3,
+ activation: Type[nn.Module] = nn.GELU,
+ iou_head_depth: int = 3,
+ iou_head_hidden_dim: int = 256,
+ ) -> None:
+ """
+ Predicts masks given an image and prompt embeddings, using a
+ transformer architecture.
+
+ Arguments:
+ transformer_dim (int): the channel dimension of the transformer
+ transformer (nn.Module): the transformer used to predict masks
+ num_multimask_outputs (int): the number of masks to predict
+ when disambiguating masks
+ activation (nn.Module): the type of activation to use when
+ upscaling masks
+ iou_head_depth (int): the depth of the MLP used to predict
+ mask quality
+ iou_head_hidden_dim (int): the hidden dimension of the MLP
+ used to predict mask quality
+ """
+ super().__init__()
+ self.transformer_dim = transformer_dim
+ self.transformer = transformer
+
+ self.num_multimask_outputs = num_multimask_outputs
+
+ self.iou_token = nn.Embedding(1, transformer_dim)
+ self.num_mask_tokens = num_multimask_outputs + 1
+ self.mask_tokens = nn.Embedding(self.num_mask_tokens, transformer_dim)
+
+ self.output_upscaling = nn.Sequential(
+ nn.ConvTranspose2d(
+ transformer_dim, transformer_dim // 4, kernel_size=2, stride=2
+ ),
+ LayerNorm2d(transformer_dim // 4),
+ activation(),
+ nn.ConvTranspose2d(
+ transformer_dim // 4, transformer_dim // 8, kernel_size=2, stride=2
+ ),
+ activation(),
+ )
+ self.output_hypernetworks_mlps = nn.ModuleList(
+ [
+ MLP(transformer_dim, transformer_dim, transformer_dim // 8, 3)
+ for i in range(self.num_mask_tokens)
+ ]
+ )
+
+ self.iou_prediction_head = MLP(
+ transformer_dim, iou_head_hidden_dim, self.num_mask_tokens, iou_head_depth
+ )
+
+ def forward(
+ self,
+ image_embeddings: torch.Tensor,
+ image_pe: torch.Tensor,
+ sparse_prompt_embeddings: torch.Tensor,
+ dense_prompt_embeddings: torch.Tensor,
+ multimask_output: bool,
+ ) -> Tuple[torch.Tensor, torch.Tensor]:
+ """
+ Predict masks given image and prompt embeddings.
+
+ Arguments:
+ image_embeddings (torch.Tensor): the embeddings from the image encoder
+ image_pe (torch.Tensor): positional encoding with the shape of image_embeddings
+ sparse_prompt_embeddings (torch.Tensor): the embeddings of the points and boxes
+ dense_prompt_embeddings (torch.Tensor): the embeddings of the mask inputs
+ multimask_output (bool): Whether to return multiple masks or a single
+ mask.
+
+ Returns:
+ torch.Tensor: batched predicted masks
+ torch.Tensor: batched predictions of mask quality
+ """
+ masks, iou_pred = self.predict_masks(
+ image_embeddings=image_embeddings,
+ image_pe=image_pe,
+ sparse_prompt_embeddings=sparse_prompt_embeddings,
+ dense_prompt_embeddings=dense_prompt_embeddings,
+ )
+
+ # Select the correct mask or masks for output
+ if multimask_output:
+ mask_slice = slice(1, None)
+ else:
+ mask_slice = slice(0, 1)
+ masks = masks[:, mask_slice, :, :]
+ iou_pred = iou_pred[:, mask_slice]
+
+ # Prepare output
+ return masks, iou_pred
+
+ def predict_masks(
+ self,
+ image_embeddings: torch.Tensor,
+ image_pe: torch.Tensor,
+ sparse_prompt_embeddings: torch.Tensor,
+ dense_prompt_embeddings: torch.Tensor,
+ ) -> Tuple[torch.Tensor, torch.Tensor]:
+ """Predicts masks. See 'forward' for more details."""
+ # Concatenate output tokens
+ output_tokens = torch.cat(
+ [self.iou_token.weight, self.mask_tokens.weight], dim=0
+ )
+ output_tokens = output_tokens.unsqueeze(0).expand(
+ sparse_prompt_embeddings.size(0), -1, -1
+ )
+
+ tokens = torch.cat((output_tokens, sparse_prompt_embeddings), dim=1)
+
+ # image_embeddings: [1, C, H, W], tokens: [B, N, C]
+ # dense_prompt_embeddings: [B, C, H, W]
+ # Expand per-image data in batch direction to be per-mask
+ src = torch.repeat_interleave(image_embeddings, tokens.shape[0], dim=0)
+ src = src + dense_prompt_embeddings
+ pos_src = torch.repeat_interleave(image_pe, tokens.shape[0], dim=0)
+ b, c, h, w = src.shape
+
+ # Run the transformer
+ hs, src = self.transformer(src, pos_src, tokens)
+ iou_token_out = hs[:, 0, :]
+ mask_tokens_out = hs[:, 1 : (1 + self.num_mask_tokens), :]
+
+ # Upscale mask embeddings and predict masks using the mask tokens
+ src = src.transpose(1, 2).view(b, c, h, w)
+ upscaled_embedding = self.output_upscaling(src)
+ hyper_in_list: List[torch.Tensor] = []
+ for i in range(self.num_mask_tokens):
+ hyper_in_list.append(
+ self.output_hypernetworks_mlps[i](mask_tokens_out[:, i, :])
+ )
+ hyper_in = torch.stack(hyper_in_list, dim=1)
+ b, c, h, w = upscaled_embedding.shape
+ masks = (hyper_in @ upscaled_embedding.view(b, c, h * w)).view(
+ b, self.num_mask_tokens, h, w
+ )
+
+ # Generate mask quality predictions
+ iou_pred = self.iou_prediction_head(iou_token_out)
+
+ return masks, iou_pred
+
+
+# Lightly adapted from
+# https://github.com/facebookresearch/MaskFormer/blob/main/mask_former/modeling/transformer/transformer_predictor.py # noqa
+class MLP(nn.Module):
+ def __init__(
+ self,
+ input_dim: int,
+ hidden_dim: int,
+ output_dim: int,
+ num_layers: int,
+ sigmoid_output: bool = False,
+ ) -> None:
+ super().__init__()
+ self.num_layers = num_layers
+ h = [hidden_dim] * (num_layers - 1)
+ self.layers = nn.ModuleList(
+ nn.Linear(n, k) for n, k in zip([input_dim] + h, h + [output_dim])
+ )
+ self.sigmoid_output = sigmoid_output
+
+ def forward(self, x):
+ for i, layer in enumerate(self.layers):
+ x = F.relu(layer(x)) if i < self.num_layers - 1 else layer(x)
+ if self.sigmoid_output:
+ x = F.sigmoid(x)
+ return x
diff --git a/py/evf_sam/model/segment_anything/modeling/prompt_encoder.py b/py/evf_sam/model/segment_anything/modeling/prompt_encoder.py
new file mode 100644
index 0000000..16bc3a4
--- /dev/null
+++ b/py/evf_sam/model/segment_anything/modeling/prompt_encoder.py
@@ -0,0 +1,238 @@
+# Copyright (c) Meta Platforms, Inc. and affiliates.
+# All rights reserved.
+
+# This source code is licensed under the license found in the
+# LICENSE file in the root directory of this source tree.
+
+from typing import Any, Optional, Tuple, Type
+
+import numpy as np
+import torch
+from torch import nn
+
+from .common import LayerNorm2d
+
+
+class PromptEncoder(nn.Module):
+ def __init__(
+ self,
+ embed_dim: int,
+ image_embedding_size: Tuple[int, int],
+ input_image_size: Tuple[int, int],
+ mask_in_chans: int,
+ activation: Type[nn.Module] = nn.GELU,
+ ) -> None:
+ """
+ Encodes prompts for input to SAM's mask decoder.
+
+ Arguments:
+ embed_dim (int): The prompts' embedding dimension
+ image_embedding_size (tuple(int, int)): The spatial size of the
+ image embedding, as (H, W).
+ input_image_size (int): The padded size of the image as input
+ to the image encoder, as (H, W).
+ mask_in_chans (int): The number of hidden channels used for
+ encoding input masks.
+ activation (nn.Module): The activation to use when encoding
+ input masks.
+ """
+ super().__init__()
+ self.embed_dim = embed_dim
+ self.input_image_size = input_image_size
+ self.image_embedding_size = image_embedding_size
+ self.pe_layer = PositionEmbeddingRandom(embed_dim // 2)
+
+ self.num_point_embeddings: int = 4 # pos/neg point + 2 box corners
+ point_embeddings = [
+ nn.Embedding(1, embed_dim) for i in range(self.num_point_embeddings)
+ ]
+ self.point_embeddings = nn.ModuleList(point_embeddings)
+ self.not_a_point_embed = nn.Embedding(1, embed_dim)
+
+ self.mask_input_size = (
+ 4 * image_embedding_size[0],
+ 4 * image_embedding_size[1],
+ )
+ self.mask_downscaling = nn.Sequential(
+ nn.Conv2d(1, mask_in_chans // 4, kernel_size=2, stride=2),
+ LayerNorm2d(mask_in_chans // 4),
+ activation(),
+ nn.Conv2d(mask_in_chans // 4, mask_in_chans, kernel_size=2, stride=2),
+ LayerNorm2d(mask_in_chans),
+ activation(),
+ nn.Conv2d(mask_in_chans, embed_dim, kernel_size=1),
+ )
+ self.no_mask_embed = nn.Embedding(1, embed_dim)
+
+ def get_dense_pe(self) -> torch.Tensor:
+ """
+ Returns the positional encoding used to encode point prompts,
+ applied to a dense set of points the shape of the image encoding.
+
+ Returns:
+ torch.Tensor: Positional encoding with shape
+ 1x(embed_dim)x(embedding_h)x(embedding_w)
+ """
+ return self.pe_layer(self.image_embedding_size).unsqueeze(0)
+
+ def _embed_points(
+ self,
+ points: torch.Tensor,
+ labels: torch.Tensor,
+ pad: bool,
+ ) -> torch.Tensor:
+ """Embeds point prompts."""
+ points = points + 0.5 # Shift to center of pixel
+ if pad:
+ padding_point = torch.zeros((points.shape[0], 1, 2), device=points.device)
+ padding_label = -torch.ones((labels.shape[0], 1), device=labels.device)
+ points = torch.cat([points, padding_point], dim=1)
+ labels = torch.cat([labels, padding_label], dim=1)
+ point_embedding = self.pe_layer.forward_with_coords(
+ points, self.input_image_size
+ )
+ point_embedding[labels == -1] = 0.0
+ point_embedding[labels == -1] += self.not_a_point_embed.weight
+ point_embedding[labels == 0] += self.point_embeddings[0].weight
+ point_embedding[labels == 1] += self.point_embeddings[1].weight
+ return point_embedding
+
+ def _embed_boxes(self, boxes: torch.Tensor) -> torch.Tensor:
+ """Embeds box prompts."""
+ boxes = boxes + 0.5 # Shift to center of pixel
+ coords = boxes.reshape(-1, 2, 2)
+ corner_embedding = self.pe_layer.forward_with_coords(
+ coords, self.input_image_size
+ )
+ corner_embedding[:, 0, :] += self.point_embeddings[2].weight
+ corner_embedding[:, 1, :] += self.point_embeddings[3].weight
+ return corner_embedding
+
+ def _embed_masks(self, masks: torch.Tensor) -> torch.Tensor:
+ """Embeds mask inputs."""
+ mask_embedding = self.mask_downscaling(masks)
+ return mask_embedding
+
+ def _get_batch_size(
+ self,
+ points: Optional[Tuple[torch.Tensor, torch.Tensor]],
+ boxes: Optional[torch.Tensor],
+ masks: Optional[torch.Tensor],
+ text_embeds: Optional[torch.Tensor],
+ ) -> int:
+ """
+ Gets the batch size of the output given the batch size of the input prompts.
+ """
+ if points is not None:
+ return points[0].shape[0]
+ elif boxes is not None:
+ return boxes.shape[0]
+ elif masks is not None:
+ return masks.shape[0]
+ elif text_embeds is not None:
+ return text_embeds.shape[0]
+ else:
+ return 1
+
+ def _get_device(self) -> torch.device:
+ return self.point_embeddings[0].weight.device
+
+ def forward(
+ self,
+ points: Optional[Tuple[torch.Tensor, torch.Tensor]],
+ boxes: Optional[torch.Tensor],
+ masks: Optional[torch.Tensor],
+ text_embeds: Optional[torch.Tensor],
+ ) -> Tuple[torch.Tensor, torch.Tensor]:
+ """
+ Embeds different types of prompts, returning both sparse and dense
+ embeddings.
+
+ Arguments:
+ points (tuple(torch.Tensor, torch.Tensor) or none): point coordinates
+ and labels to embed.
+ boxes (torch.Tensor or none): boxes to embed
+ masks (torch.Tensor or none): masks to embed
+
+ Returns:
+ torch.Tensor: sparse embeddings for the points and boxes, with shape
+ BxNx(embed_dim), where N is determined by the number of input points
+ and boxes.
+ torch.Tensor: dense embeddings for the masks, in the shape
+ Bx(embed_dim)x(embed_H)x(embed_W)
+ """
+ bs = self._get_batch_size(points, boxes, masks, text_embeds)
+ sparse_embeddings = torch.empty(
+ (bs, 0, self.embed_dim), device=self._get_device()
+ )
+ if points is not None:
+ coords, labels = points
+ point_embeddings = self._embed_points(coords, labels, pad=(boxes is None))
+ sparse_embeddings = torch.cat([sparse_embeddings, point_embeddings], dim=1)
+ if boxes is not None:
+ box_embeddings = self._embed_boxes(boxes)
+ sparse_embeddings = torch.cat([sparse_embeddings, box_embeddings], dim=1)
+
+ if text_embeds is not None:
+ sparse_embeddings = torch.cat([sparse_embeddings, text_embeds], dim=1)
+
+ if masks is not None:
+ dense_embeddings = self._embed_masks(masks)
+ else:
+ dense_embeddings = self.no_mask_embed.weight.reshape(1, -1, 1, 1).expand(
+ bs, -1, self.image_embedding_size[0], self.image_embedding_size[1]
+ )
+
+ return sparse_embeddings, dense_embeddings
+
+
+class PositionEmbeddingRandom(nn.Module):
+ """
+ Positional encoding using random spatial frequencies.
+ """
+
+ def __init__(self, num_pos_feats: int = 64, scale: Optional[float] = None) -> None:
+ super().__init__()
+ if scale is None or scale <= 0.0:
+ scale = 1.0
+ self.register_buffer(
+ "positional_encoding_gaussian_matrix",
+ scale * torch.randn((2, num_pos_feats)),
+ )
+
+ def _pe_encoding(self, coords: torch.Tensor) -> torch.Tensor:
+ """Positionally encode points that are normalized to [0,1]."""
+ # assuming coords are in [0, 1]^2 square and have d_1 x ... x d_n x 2 shape
+ coords = 2 * coords - 1
+
+ if coords.dtype != self.positional_encoding_gaussian_matrix.dtype:
+ coords = coords.to(self.positional_encoding_gaussian_matrix.dtype)
+
+ coords = coords @ self.positional_encoding_gaussian_matrix
+ coords = 2 * np.pi * coords
+ # outputs d_1 x ... x d_n x C shape
+ return torch.cat([torch.sin(coords), torch.cos(coords)], dim=-1)
+
+ def forward(self, size: Tuple[int, int]) -> torch.Tensor:
+ """Generate positional encoding for a grid of the specified size."""
+ h, w = size
+ device: Any = self.positional_encoding_gaussian_matrix.device
+ grid = torch.ones(
+ (h, w), device=device, dtype=self.positional_encoding_gaussian_matrix.dtype
+ )
+ y_embed = grid.cumsum(dim=0) - 0.5
+ x_embed = grid.cumsum(dim=1) - 0.5
+ y_embed = y_embed / h
+ x_embed = x_embed / w
+
+ pe = self._pe_encoding(torch.stack([x_embed, y_embed], dim=-1))
+ return pe.permute(2, 0, 1) # C x H x W
+
+ def forward_with_coords(
+ self, coords_input: torch.Tensor, image_size: Tuple[int, int]
+ ) -> torch.Tensor:
+ """Positionally encode points that are not normalized to [0,1]."""
+ coords = coords_input.clone()
+ coords[:, :, 0] = coords[:, :, 0] / image_size[1]
+ coords[:, :, 1] = coords[:, :, 1] / image_size[0]
+ return self._pe_encoding(coords.to(torch.float)) # B x N x C
diff --git a/py/evf_sam/model/segment_anything/modeling/sam.py b/py/evf_sam/model/segment_anything/modeling/sam.py
new file mode 100644
index 0000000..f1d82ca
--- /dev/null
+++ b/py/evf_sam/model/segment_anything/modeling/sam.py
@@ -0,0 +1,184 @@
+# Copyright (c) Meta Platforms, Inc. and affiliates.
+# All rights reserved.
+
+# This source code is licensed under the license found in the
+# LICENSE file in the root directory of this source tree.
+
+from typing import Any, Dict, List, Tuple
+
+import torch
+from torch import nn
+from torch.nn import functional as F
+
+from .image_encoder import ImageEncoderViT
+from .mask_decoder import MaskDecoder
+from .prompt_encoder import PromptEncoder
+
+
+class Sam(nn.Module):
+ mask_threshold: float = 0.0
+ image_format: str = "RGB"
+
+ def __init__(
+ self,
+ image_encoder: ImageEncoderViT,
+ prompt_encoder: PromptEncoder,
+ mask_decoder: MaskDecoder,
+ pixel_mean: List[float] = [123.675, 116.28, 103.53],
+ pixel_std: List[float] = [58.395, 57.12, 57.375],
+ ) -> None:
+ """
+ SAM predicts object masks from an image and input prompts.
+
+ Arguments:
+ image_encoder (ImageEncoderViT): The backbone used to encode the
+ image into image embeddings that allow for efficient mask prediction.
+ prompt_encoder (PromptEncoder): Encodes various types of input prompts.
+ mask_decoder (MaskDecoder): Predicts masks from the image embeddings
+ and encoded prompts.
+ pixel_mean (list(float)): Mean values for normalizing pixels in the input image.
+ pixel_std (list(float)): Std values for normalizing pixels in the input image.
+ """
+ super().__init__()
+ self.image_encoder = image_encoder
+ self.prompt_encoder = prompt_encoder
+ self.mask_decoder = mask_decoder
+ self.register_buffer(
+ "pixel_mean", torch.Tensor(pixel_mean).view(-1, 1, 1), False
+ )
+ self.register_buffer("pixel_std", torch.Tensor(pixel_std).view(-1, 1, 1), False)
+
+ @property
+ def device(self) -> Any:
+ return self.pixel_mean.device
+
+ @torch.no_grad()
+ def forward(
+ self,
+ batched_input: List[Dict[str, Any]],
+ multimask_output: bool,
+ ) -> List[Dict[str, torch.Tensor]]:
+ """
+ Predicts masks end-to-end from provided images and prompts.
+ If prompts are not known in advance, using SamPredictor is
+ recommended over calling the model directly.
+
+ Arguments:
+ batched_input (list(dict)): A list over input images, each a
+ dictionary with the following keys. A prompt key can be
+ excluded if it is not present.
+ 'image': The image as a torch tensor in 3xHxW format,
+ already transformed for input to the model.
+ 'original_size': (tuple(int, int)) The original size of
+ the image before transformation, as (H, W).
+ 'point_coords': (torch.Tensor) Batched point prompts for
+ this image, with shape BxNx2. Already transformed to the
+ input frame of the model.
+ 'point_labels': (torch.Tensor) Batched labels for point prompts,
+ with shape BxN.
+ 'boxes': (torch.Tensor) Batched box inputs, with shape Bx4.
+ Already transformed to the input frame of the model.
+ 'mask_inputs': (torch.Tensor) Batched mask inputs to the model,
+ in the form Bx1xHxW.
+ multimask_output (bool): Whether the model should predict multiple
+ disambiguating masks, or return a single mask.
+
+ Returns:
+ (list(dict)): A list over input images, where each element is
+ as dictionary with the following keys.
+ 'masks': (torch.Tensor) Batched binary mask predictions,
+ with shape BxCxHxW, where B is the number of input prompts,
+ C is determined by multimask_output, and (H, W) is the
+ original size of the image.
+ 'iou_predictions': (torch.Tensor) The model's predictions
+ of mask quality, in shape BxC.
+ 'low_res_logits': (torch.Tensor) Low resolution logits with
+ shape BxCxHxW, where H=W=256. Can be passed as mask input
+ to subsequent iterations of prediction.
+ """
+ input_images = torch.stack(
+ [self.preprocess(x["image"]) for x in batched_input], dim=0
+ )
+ image_embeddings = self.image_encoder(input_images)
+
+ outputs = []
+ for image_record, curr_embedding in zip(batched_input, image_embeddings):
+ if "point_coords" in image_record:
+ points = (image_record["point_coords"], image_record["point_labels"])
+ else:
+ points = None
+ sparse_embeddings, dense_embeddings = self.prompt_encoder(
+ points=points,
+ boxes=image_record.get("boxes", None),
+ masks=image_record.get("mask_inputs", None),
+ )
+ low_res_masks, iou_predictions = self.mask_decoder(
+ image_embeddings=curr_embedding.unsqueeze(0),
+ image_pe=self.prompt_encoder.get_dense_pe(),
+ sparse_prompt_embeddings=sparse_embeddings,
+ dense_prompt_embeddings=dense_embeddings,
+ multimask_output=multimask_output,
+ )
+ masks = self.postprocess_masks(
+ low_res_masks,
+ input_size=image_record["image"].shape[-2:],
+ original_size=image_record["original_size"],
+ )
+ masks = masks > self.mask_threshold
+ outputs.append(
+ {
+ "masks": masks,
+ "iou_predictions": iou_predictions,
+ "low_res_logits": low_res_masks,
+ }
+ )
+ return outputs
+
+ def postprocess_masks(
+ self,
+ masks: torch.Tensor,
+ input_size: Tuple[int, ...],
+ original_size: Tuple[int, ...],
+ ) -> torch.Tensor:
+ """
+ Remove padding and upscale masks to the original image size.
+
+ Arguments:
+ masks (torch.Tensor): Batched masks from the mask_decoder,
+ in BxCxHxW format.
+ input_size (tuple(int, int)): The size of the image input to the
+ model, in (H, W) format. Used to remove padding.
+ original_size (tuple(int, int)): The original size of the image
+ before resizing for input to the model, in (H, W) format.
+
+ Returns:
+ (torch.Tensor): Batched masks in BxCxHxW format, where (H, W)
+ is given by original_size.
+ """
+
+ dtype = masks.dtype
+
+ masks = F.interpolate(
+ masks.float(),
+ (self.image_encoder.img_size, self.image_encoder.img_size),
+ mode="bilinear",
+ align_corners=False,
+ )
+ # masks = masks.to(dtype)
+ masks = masks[..., : input_size[0], : input_size[1]]
+ masks = F.interpolate(
+ masks, original_size, mode="bilinear", align_corners=False
+ )
+ return masks
+
+ def preprocess(self, x: torch.Tensor) -> torch.Tensor:
+ """Normalize pixel values and pad to a square input."""
+ # Normalize colors
+ x = (x - self.pixel_mean) / self.pixel_std
+
+ # Pad
+ h, w = x.shape[-2:]
+ padh = self.image_encoder.img_size - h
+ padw = self.image_encoder.img_size - w
+ x = F.pad(x, (0, padw, 0, padh))
+ return x
diff --git a/py/evf_sam/model/segment_anything/modeling/transformer.py b/py/evf_sam/model/segment_anything/modeling/transformer.py
new file mode 100644
index 0000000..8c511e4
--- /dev/null
+++ b/py/evf_sam/model/segment_anything/modeling/transformer.py
@@ -0,0 +1,242 @@
+# Copyright (c) Meta Platforms, Inc. and affiliates.
+# All rights reserved.
+
+# This source code is licensed under the license found in the
+# LICENSE file in the root directory of this source tree.
+
+import math
+from typing import Tuple, Type
+
+import torch
+from torch import Tensor, nn
+
+from .common import MLPBlock
+
+
+class TwoWayTransformer(nn.Module):
+ def __init__(
+ self,
+ depth: int,
+ embedding_dim: int,
+ num_heads: int,
+ mlp_dim: int,
+ activation: Type[nn.Module] = nn.ReLU,
+ attention_downsample_rate: int = 2,
+ ) -> None:
+ """
+ A transformer decoder that attends to an input image using
+ queries whose positional embedding is supplied.
+
+ Args:
+ depth (int): number of layers in the transformer
+ embedding_dim (int): the channel dimension for the input embeddings
+ num_heads (int): the number of heads for multihead attention. Must
+ divide embedding_dim
+ mlp_dim (int): the channel dimension internal to the MLP block
+ activation (nn.Module): the activation to use in the MLP block
+ """
+ super().__init__()
+ self.depth = depth
+ self.embedding_dim = embedding_dim
+ self.num_heads = num_heads
+ self.mlp_dim = mlp_dim
+ self.layers = nn.ModuleList()
+
+ for i in range(depth):
+ self.layers.append(
+ TwoWayAttentionBlock(
+ embedding_dim=embedding_dim,
+ num_heads=num_heads,
+ mlp_dim=mlp_dim,
+ activation=activation,
+ attention_downsample_rate=attention_downsample_rate,
+ skip_first_layer_pe=(i == 0),
+ )
+ )
+
+ self.final_attn_token_to_image = Attention(
+ embedding_dim, num_heads, downsample_rate=attention_downsample_rate
+ )
+ self.norm_final_attn = nn.LayerNorm(embedding_dim)
+
+ def forward(
+ self,
+ image_embedding: Tensor,
+ image_pe: Tensor,
+ point_embedding: Tensor,
+ ) -> Tuple[Tensor, Tensor]:
+ """
+ Args:
+ image_embedding (torch.Tensor): image to attend to. Should be shape
+ B x embedding_dim x h x w for any h and w.
+ image_pe (torch.Tensor): the positional encoding to add to the image. Must
+ have the same shape as image_embedding.
+ point_embedding (torch.Tensor): the embedding to add to the query points.
+ Must have shape B x N_points x embedding_dim for any N_points.
+
+ Returns:
+ torch.Tensor: the processed point_embedding
+ torch.Tensor: the processed image_embedding
+ """
+ # BxCxHxW -> BxHWxC == B x N_image_tokens x C
+ bs, c, h, w = image_embedding.shape
+ image_embedding = image_embedding.flatten(2).permute(0, 2, 1)
+ image_pe = image_pe.flatten(2).permute(0, 2, 1)
+
+ # Prepare queries
+ queries = point_embedding
+ keys = image_embedding
+
+ # Apply transformer blocks and final layernorm
+ for layer in self.layers:
+ queries, keys = layer(
+ queries=queries,
+ keys=keys,
+ query_pe=point_embedding,
+ key_pe=image_pe,
+ )
+
+ # Apply the final attention layer from the points to the image
+ q = queries + point_embedding
+ k = keys + image_pe
+ attn_out = self.final_attn_token_to_image(q=q, k=k, v=keys)
+ queries = queries + attn_out
+ queries = self.norm_final_attn(queries)
+
+ return queries, keys
+
+
+class TwoWayAttentionBlock(nn.Module):
+ def __init__(
+ self,
+ embedding_dim: int,
+ num_heads: int,
+ mlp_dim: int = 2048,
+ activation: Type[nn.Module] = nn.ReLU,
+ attention_downsample_rate: int = 2,
+ skip_first_layer_pe: bool = False,
+ ) -> None:
+ """
+ A transformer block with four layers: (1) self-attention of sparse
+ inputs, (2) cross attention of sparse inputs to dense inputs, (3) mlp
+ block on sparse inputs, and (4) cross attention of dense inputs to sparse
+ inputs.
+
+ Arguments:
+ embedding_dim (int): the channel dimension of the embeddings
+ num_heads (int): the number of heads in the attention layers
+ mlp_dim (int): the hidden dimension of the mlp block
+ activation (nn.Module): the activation of the mlp block
+ skip_first_layer_pe (bool): skip the PE on the first layer
+ """
+ super().__init__()
+ self.self_attn = Attention(embedding_dim, num_heads)
+ self.norm1 = nn.LayerNorm(embedding_dim)
+
+ self.cross_attn_token_to_image = Attention(
+ embedding_dim, num_heads, downsample_rate=attention_downsample_rate
+ )
+ self.norm2 = nn.LayerNorm(embedding_dim)
+
+ self.mlp = MLPBlock(embedding_dim, mlp_dim, activation)
+ self.norm3 = nn.LayerNorm(embedding_dim)
+
+ self.norm4 = nn.LayerNorm(embedding_dim)
+ self.cross_attn_image_to_token = Attention(
+ embedding_dim, num_heads, downsample_rate=attention_downsample_rate
+ )
+
+ self.skip_first_layer_pe = skip_first_layer_pe
+
+ def forward(
+ self, queries: Tensor, keys: Tensor, query_pe: Tensor, key_pe: Tensor
+ ) -> Tuple[Tensor, Tensor]:
+ # Self attention block
+ if self.skip_first_layer_pe:
+ queries = self.self_attn(q=queries, k=queries, v=queries)
+ else:
+ q = queries + query_pe
+ attn_out = self.self_attn(q=q, k=q, v=queries)
+ queries = queries + attn_out
+ queries = self.norm1(queries)
+
+ # Cross attention block, tokens attending to image embedding
+ q = queries + query_pe
+ k = keys + key_pe
+ attn_out = self.cross_attn_token_to_image(q=q, k=k, v=keys)
+ queries = queries + attn_out
+ queries = self.norm2(queries)
+
+ # MLP block
+ mlp_out = self.mlp(queries)
+ queries = queries + mlp_out
+ queries = self.norm3(queries)
+
+ # Cross attention block, image embedding attending to tokens
+ q = queries + query_pe
+ k = keys + key_pe
+ attn_out = self.cross_attn_image_to_token(q=k, k=q, v=queries)
+ keys = keys + attn_out
+ keys = self.norm4(keys)
+
+ return queries, keys
+
+
+class Attention(nn.Module):
+ """
+ An attention layer that allows for downscaling the size of the embedding
+ after projection to queries, keys, and values.
+ """
+
+ def __init__(
+ self,
+ embedding_dim: int,
+ num_heads: int,
+ downsample_rate: int = 1,
+ ) -> None:
+ super().__init__()
+ self.embedding_dim = embedding_dim
+ self.internal_dim = embedding_dim // downsample_rate
+ self.num_heads = num_heads
+ assert (
+ self.internal_dim % num_heads == 0
+ ), "num_heads must divide embedding_dim."
+
+ self.q_proj = nn.Linear(embedding_dim, self.internal_dim)
+ self.k_proj = nn.Linear(embedding_dim, self.internal_dim)
+ self.v_proj = nn.Linear(embedding_dim, self.internal_dim)
+ self.out_proj = nn.Linear(self.internal_dim, embedding_dim)
+
+ def _separate_heads(self, x: Tensor, num_heads: int) -> Tensor:
+ b, n, c = x.shape
+ x = x.reshape(b, n, num_heads, c // num_heads)
+ return x.transpose(1, 2) # B x N_heads x N_tokens x C_per_head
+
+ def _recombine_heads(self, x: Tensor) -> Tensor:
+ b, n_heads, n_tokens, c_per_head = x.shape
+ x = x.transpose(1, 2)
+ return x.reshape(b, n_tokens, n_heads * c_per_head) # B x N_tokens x C
+
+ def forward(self, q: Tensor, k: Tensor, v: Tensor) -> Tensor:
+ # Input projections
+ q = self.q_proj(q)
+ k = self.k_proj(k)
+ v = self.v_proj(v)
+
+ # Separate into heads
+ q = self._separate_heads(q, self.num_heads)
+ k = self._separate_heads(k, self.num_heads)
+ v = self._separate_heads(v, self.num_heads)
+
+ # Attention
+ _, _, _, c_per_head = q.shape
+ attn = q @ k.permute(0, 1, 3, 2) # B x N_heads x N_tokens x N_tokens
+ attn = attn / math.sqrt(c_per_head)
+ attn = torch.softmax(attn, dim=-1)
+
+ # Get output
+ out = attn @ v
+ out = self._recombine_heads(out)
+ out = self.out_proj(out)
+
+ return out
diff --git a/py/evf_sam/model/segment_anything/predictor.py b/py/evf_sam/model/segment_anything/predictor.py
new file mode 100644
index 0000000..bf52d81
--- /dev/null
+++ b/py/evf_sam/model/segment_anything/predictor.py
@@ -0,0 +1,284 @@
+# Copyright (c) Meta Platforms, Inc. and affiliates.
+# All rights reserved.
+
+# This source code is licensed under the license found in the
+# LICENSE file in the root directory of this source tree.
+
+from typing import Optional, Tuple
+
+import numpy as np
+import torch
+
+from .modeling import Sam
+from .utils.transforms import ResizeLongestSide
+
+
+class SamPredictor:
+ def __init__(
+ self,
+ sam_model: Sam,
+ ) -> None:
+ """
+ Uses SAM to calculate the image embedding for an image, and then
+ allow repeated, efficient mask prediction given prompts.
+
+ Arguments:
+ sam_model (Sam): The model to use for mask prediction.
+ """
+ super().__init__()
+ self.model = sam_model
+ self.transform = ResizeLongestSide(sam_model.image_encoder.img_size)
+ self.reset_image()
+
+ def set_image(
+ self,
+ image: np.ndarray,
+ image_format: str = "RGB",
+ ) -> None:
+ """
+ Calculates the image embeddings for the provided image, allowing
+ masks to be predicted with the 'predict' method.
+
+ Arguments:
+ image (np.ndarray): The image for calculating masks. Expects an
+ image in HWC uint8 format, with pixel values in [0, 255].
+ image_format (str): The color format of the image, in ['RGB', 'BGR'].
+ """
+ assert image_format in [
+ "RGB",
+ "BGR",
+ ], f"image_format must be in ['RGB', 'BGR'], is {image_format}."
+ if image_format != self.model.image_format:
+ image = image[..., ::-1]
+
+ # Transform the image to the form expected by the model
+ input_image = self.transform.apply_image(image)
+ input_image_torch = torch.as_tensor(input_image, device=self.device)
+ input_image_torch = input_image_torch.permute(2, 0, 1).contiguous()[
+ None, :, :, :
+ ]
+
+ self.set_torch_image(input_image_torch, image.shape[:2])
+
+ @torch.no_grad()
+ def set_torch_image(
+ self,
+ transformed_image: torch.Tensor,
+ original_image_size: Tuple[int, ...],
+ ) -> None:
+ """
+ Calculates the image embeddings for the provided image, allowing
+ masks to be predicted with the 'predict' method. Expects the input
+ image to be already transformed to the format expected by the model.
+
+ Arguments:
+ transformed_image (torch.Tensor): The input image, with shape
+ 1x3xHxW, which has been transformed with ResizeLongestSide.
+ original_image_size (tuple(int, int)): The size of the image
+ before transformation, in (H, W) format.
+ """
+ assert (
+ len(transformed_image.shape) == 4
+ and transformed_image.shape[1] == 3
+ and max(*transformed_image.shape[2:]) == self.model.image_encoder.img_size
+ ), f"set_torch_image input must be BCHW with long side {self.model.image_encoder.img_size}."
+ self.reset_image()
+
+ self.original_size = original_image_size
+ self.input_size = tuple(transformed_image.shape[-2:])
+ input_image = self.model.preprocess(transformed_image)
+ self.features = self.model.image_encoder(input_image)
+ self.is_image_set = True
+
+ def predict(
+ self,
+ point_coords: Optional[np.ndarray] = None,
+ point_labels: Optional[np.ndarray] = None,
+ box: Optional[np.ndarray] = None,
+ mask_input: Optional[np.ndarray] = None,
+ multimask_output: bool = True,
+ return_logits: bool = False,
+ ) -> Tuple[np.ndarray, np.ndarray, np.ndarray]:
+ """
+ Predict masks for the given input prompts, using the currently set image.
+
+ Arguments:
+ point_coords (np.ndarray or None): A Nx2 array of point prompts to the
+ model. Each point is in (X,Y) in pixels.
+ point_labels (np.ndarray or None): A length N array of labels for the
+ point prompts. 1 indicates a foreground point and 0 indicates a
+ background point.
+ box (np.ndarray or None): A length 4 array given a box prompt to the
+ model, in XYXY format.
+ mask_input (np.ndarray): A low resolution mask input to the model, typically
+ coming from a previous prediction iteration. Has form 1xHxW, where
+ for SAM, H=W=256.
+ multimask_output (bool): If true, the model will return three masks.
+ For ambiguous input prompts (such as a single click), this will often
+ produce better masks than a single prediction. If only a single
+ mask is needed, the model's predicted quality score can be used
+ to select the best mask. For non-ambiguous prompts, such as multiple
+ input prompts, multimask_output=False can give better results.
+ return_logits (bool): If true, returns un-thresholded masks logits
+ instead of a binary mask.
+
+ Returns:
+ (np.ndarray): The output masks in CxHxW format, where C is the
+ number of masks, and (H, W) is the original image size.
+ (np.ndarray): An array of length C containing the model's
+ predictions for the quality of each mask.
+ (np.ndarray): An array of shape CxHxW, where C is the number
+ of masks and H=W=256. These low resolution logits can be passed to
+ a subsequent iteration as mask input.
+ """
+ if not self.is_image_set:
+ raise RuntimeError(
+ "An image must be set with .set_image(...) before mask prediction."
+ )
+
+ # Transform input prompts
+ coords_torch, labels_torch, box_torch, mask_input_torch = None, None, None, None
+ if point_coords is not None:
+ assert (
+ point_labels is not None
+ ), "point_labels must be supplied if point_coords is supplied."
+ point_coords = self.transform.apply_coords(point_coords, self.original_size)
+ coords_torch = torch.as_tensor(
+ point_coords, dtype=torch.float, device=self.device
+ )
+ labels_torch = torch.as_tensor(
+ point_labels, dtype=torch.int, device=self.device
+ )
+ coords_torch, labels_torch = coords_torch[None, :, :], labels_torch[None, :]
+ if box is not None:
+ box = self.transform.apply_boxes(box, self.original_size)
+ box_torch = torch.as_tensor(box, dtype=torch.float, device=self.device)
+ box_torch = box_torch[None, :]
+ if mask_input is not None:
+ mask_input_torch = torch.as_tensor(
+ mask_input, dtype=torch.float, device=self.device
+ )
+ mask_input_torch = mask_input_torch[None, :, :, :]
+
+ masks, iou_predictions, low_res_masks = self.predict_torch(
+ coords_torch,
+ labels_torch,
+ box_torch,
+ mask_input_torch,
+ multimask_output,
+ return_logits=return_logits,
+ )
+
+ masks_np = masks[0].detach().cpu().numpy()
+ iou_predictions_np = iou_predictions[0].detach().cpu().numpy()
+ low_res_masks_np = low_res_masks[0].detach().cpu().numpy()
+ return masks_np, iou_predictions_np, low_res_masks_np
+
+ @torch.no_grad()
+ def predict_torch(
+ self,
+ point_coords: Optional[torch.Tensor],
+ point_labels: Optional[torch.Tensor],
+ boxes: Optional[torch.Tensor] = None,
+ mask_input: Optional[torch.Tensor] = None,
+ multimask_output: bool = True,
+ return_logits: bool = False,
+ ) -> Tuple[torch.Tensor, torch.Tensor, torch.Tensor]:
+ """
+ Predict masks for the given input prompts, using the currently set image.
+ Input prompts are batched torch tensors and are expected to already be
+ transformed to the input frame using ResizeLongestSide.
+
+ Arguments:
+ point_coords (torch.Tensor or None): A BxNx2 array of point prompts to the
+ model. Each point is in (X,Y) in pixels.
+ point_labels (torch.Tensor or None): A BxN array of labels for the
+ point prompts. 1 indicates a foreground point and 0 indicates a
+ background point.
+ boxes (np.ndarray or None): A Bx4 array given a box prompt to the
+ model, in XYXY format.
+ mask_input (np.ndarray): A low resolution mask input to the model, typically
+ coming from a previous prediction iteration. Has form Bx1xHxW, where
+ for SAM, H=W=256. Masks returned by a previous iteration of the
+ predict method do not need further transformation.
+ multimask_output (bool): If true, the model will return three masks.
+ For ambiguous input prompts (such as a single click), this will often
+ produce better masks than a single prediction. If only a single
+ mask is needed, the model's predicted quality score can be used
+ to select the best mask. For non-ambiguous prompts, such as multiple
+ input prompts, multimask_output=False can give better results.
+ return_logits (bool): If true, returns un-thresholded masks logits
+ instead of a binary mask.
+
+ Returns:
+ (torch.Tensor): The output masks in BxCxHxW format, where C is the
+ number of masks, and (H, W) is the original image size.
+ (torch.Tensor): An array of shape BxC containing the model's
+ predictions for the quality of each mask.
+ (torch.Tensor): An array of shape BxCxHxW, where C is the number
+ of masks and H=W=256. These low res logits can be passed to
+ a subsequent iteration as mask input.
+ """
+ if not self.is_image_set:
+ raise RuntimeError(
+ "An image must be set with .set_image(...) before mask prediction."
+ )
+
+ if point_coords is not None:
+ points = (point_coords, point_labels)
+ else:
+ points = None
+
+ # Embed prompts
+ sparse_embeddings, dense_embeddings = self.model.prompt_encoder(
+ points=points,
+ boxes=boxes,
+ masks=mask_input,
+ )
+
+ # Predict masks
+ low_res_masks, iou_predictions = self.model.mask_decoder(
+ image_embeddings=self.features,
+ image_pe=self.model.prompt_encoder.get_dense_pe(),
+ sparse_prompt_embeddings=sparse_embeddings,
+ dense_prompt_embeddings=dense_embeddings,
+ multimask_output=multimask_output,
+ )
+
+ # Upscale the masks to the original image resolution
+ masks = self.model.postprocess_masks(
+ low_res_masks, self.input_size, self.original_size
+ )
+
+ if not return_logits:
+ masks = masks > self.model.mask_threshold
+
+ return masks, iou_predictions, low_res_masks
+
+ def get_image_embedding(self) -> torch.Tensor:
+ """
+ Returns the image embeddings for the currently set image, with
+ shape 1xCxHxW, where C is the embedding dimension and (H,W) are
+ the embedding spatial dimension of SAM (typically C=256, H=W=64).
+ """
+ if not self.is_image_set:
+ raise RuntimeError(
+ "An image must be set with .set_image(...) to generate an embedding."
+ )
+ assert (
+ self.features is not None
+ ), "Features must exist if an image has been set."
+ return self.features
+
+ @property
+ def device(self) -> torch.device:
+ return self.model.device
+
+ def reset_image(self) -> None:
+ """Resets the currently set image."""
+ self.is_image_set = False
+ self.features = None
+ self.orig_h = None
+ self.orig_w = None
+ self.input_h = None
+ self.input_w = None
diff --git a/py/evf_sam/model/segment_anything/utils/__init__.py b/py/evf_sam/model/segment_anything/utils/__init__.py
new file mode 100644
index 0000000..5277f46
--- /dev/null
+++ b/py/evf_sam/model/segment_anything/utils/__init__.py
@@ -0,0 +1,5 @@
+# Copyright (c) Meta Platforms, Inc. and affiliates.
+# All rights reserved.
+
+# This source code is licensed under the license found in the
+# LICENSE file in the root directory of this source tree.
diff --git a/py/evf_sam/model/segment_anything/utils/amg.py b/py/evf_sam/model/segment_anything/utils/amg.py
new file mode 100644
index 0000000..5c3bc5d
--- /dev/null
+++ b/py/evf_sam/model/segment_anything/utils/amg.py
@@ -0,0 +1,346 @@
+# Copyright (c) Meta Platforms, Inc. and affiliates.
+# All rights reserved.
+
+# This source code is licensed under the license found in the
+# LICENSE file in the root directory of this source tree.
+
+import math
+from copy import deepcopy
+from itertools import product
+from typing import Any, Dict, Generator, ItemsView, List, Tuple
+
+import numpy as np
+import torch
+
+
+class MaskData:
+ """
+ A structure for storing masks and their related data in batched format.
+ Implements basic filtering and concatenation.
+ """
+
+ def __init__(self, **kwargs) -> None:
+ for v in kwargs.values():
+ assert isinstance(
+ v, (list, np.ndarray, torch.Tensor)
+ ), "MaskData only supports list, numpy arrays, and torch tensors."
+ self._stats = dict(**kwargs)
+
+ def __setitem__(self, key: str, item: Any) -> None:
+ assert isinstance(
+ item, (list, np.ndarray, torch.Tensor)
+ ), "MaskData only supports list, numpy arrays, and torch tensors."
+ self._stats[key] = item
+
+ def __delitem__(self, key: str) -> None:
+ del self._stats[key]
+
+ def __getitem__(self, key: str) -> Any:
+ return self._stats[key]
+
+ def items(self) -> ItemsView[str, Any]:
+ return self._stats.items()
+
+ def filter(self, keep: torch.Tensor) -> None:
+ for k, v in self._stats.items():
+ if v is None:
+ self._stats[k] = None
+ elif isinstance(v, torch.Tensor):
+ self._stats[k] = v[torch.as_tensor(keep, device=v.device)]
+ elif isinstance(v, np.ndarray):
+ self._stats[k] = v[keep.detach().cpu().numpy()]
+ elif isinstance(v, list) and keep.dtype == torch.bool:
+ self._stats[k] = [a for i, a in enumerate(v) if keep[i]]
+ elif isinstance(v, list):
+ self._stats[k] = [v[i] for i in keep]
+ else:
+ raise TypeError(f"MaskData key {k} has an unsupported type {type(v)}.")
+
+ def cat(self, new_stats: "MaskData") -> None:
+ for k, v in new_stats.items():
+ if k not in self._stats or self._stats[k] is None:
+ self._stats[k] = deepcopy(v)
+ elif isinstance(v, torch.Tensor):
+ self._stats[k] = torch.cat([self._stats[k], v], dim=0)
+ elif isinstance(v, np.ndarray):
+ self._stats[k] = np.concatenate([self._stats[k], v], axis=0)
+ elif isinstance(v, list):
+ self._stats[k] = self._stats[k] + deepcopy(v)
+ else:
+ raise TypeError(f"MaskData key {k} has an unsupported type {type(v)}.")
+
+ def to_numpy(self) -> None:
+ for k, v in self._stats.items():
+ if isinstance(v, torch.Tensor):
+ self._stats[k] = v.detach().cpu().numpy()
+
+
+def is_box_near_crop_edge(
+ boxes: torch.Tensor, crop_box: List[int], orig_box: List[int], atol: float = 20.0
+) -> torch.Tensor:
+ """Filter masks at the edge of a crop, but not at the edge of the original image."""
+ crop_box_torch = torch.as_tensor(crop_box, dtype=torch.float, device=boxes.device)
+ orig_box_torch = torch.as_tensor(orig_box, dtype=torch.float, device=boxes.device)
+ boxes = uncrop_boxes_xyxy(boxes, crop_box).float()
+ near_crop_edge = torch.isclose(boxes, crop_box_torch[None, :], atol=atol, rtol=0)
+ near_image_edge = torch.isclose(boxes, orig_box_torch[None, :], atol=atol, rtol=0)
+ near_crop_edge = torch.logical_and(near_crop_edge, ~near_image_edge)
+ return torch.any(near_crop_edge, dim=1)
+
+
+def box_xyxy_to_xywh(box_xyxy: torch.Tensor) -> torch.Tensor:
+ box_xywh = deepcopy(box_xyxy)
+ box_xywh[2] = box_xywh[2] - box_xywh[0]
+ box_xywh[3] = box_xywh[3] - box_xywh[1]
+ return box_xywh
+
+
+def batch_iterator(batch_size: int, *args) -> Generator[List[Any], None, None]:
+ assert len(args) > 0 and all(
+ len(a) == len(args[0]) for a in args
+ ), "Batched iteration must have inputs of all the same size."
+ n_batches = len(args[0]) // batch_size + int(len(args[0]) % batch_size != 0)
+ for b in range(n_batches):
+ yield [arg[b * batch_size : (b + 1) * batch_size] for arg in args]
+
+
+def mask_to_rle_pytorch(tensor: torch.Tensor) -> List[Dict[str, Any]]:
+ """
+ Encodes masks to an uncompressed RLE, in the format expected by
+ pycoco tools.
+ """
+ # Put in fortran order and flatten h,w
+ b, h, w = tensor.shape
+ tensor = tensor.permute(0, 2, 1).flatten(1)
+
+ # Compute change indices
+ diff = tensor[:, 1:] ^ tensor[:, :-1]
+ change_indices = diff.nonzero()
+
+ # Encode run length
+ out = []
+ for i in range(b):
+ cur_idxs = change_indices[change_indices[:, 0] == i, 1]
+ cur_idxs = torch.cat(
+ [
+ torch.tensor([0], dtype=cur_idxs.dtype, device=cur_idxs.device),
+ cur_idxs + 1,
+ torch.tensor([h * w], dtype=cur_idxs.dtype, device=cur_idxs.device),
+ ]
+ )
+ btw_idxs = cur_idxs[1:] - cur_idxs[:-1]
+ counts = [] if tensor[i, 0] == 0 else [0]
+ counts.extend(btw_idxs.detach().cpu().tolist())
+ out.append({"size": [h, w], "counts": counts})
+ return out
+
+
+def rle_to_mask(rle: Dict[str, Any]) -> np.ndarray:
+ """Compute a binary mask from an uncompressed RLE."""
+ h, w = rle["size"]
+ mask = np.empty(h * w, dtype=bool)
+ idx = 0
+ parity = False
+ for count in rle["counts"]:
+ mask[idx : idx + count] = parity
+ idx += count
+ parity ^= True
+ mask = mask.reshape(w, h)
+ return mask.transpose() # Put in C order
+
+
+def area_from_rle(rle: Dict[str, Any]) -> int:
+ return sum(rle["counts"][1::2])
+
+
+def calculate_stability_score(
+ masks: torch.Tensor, mask_threshold: float, threshold_offset: float
+) -> torch.Tensor:
+ """
+ Computes the stability score for a batch of masks. The stability
+ score is the IoU between the binary masks obtained by thresholding
+ the predicted mask logits at high and low values.
+ """
+ # One mask is always contained inside the other.
+ # Save memory by preventing unnecessary cast to torch.int64
+ intersections = (
+ (masks > (mask_threshold + threshold_offset))
+ .sum(-1, dtype=torch.int16)
+ .sum(-1, dtype=torch.int32)
+ )
+ unions = (
+ (masks > (mask_threshold - threshold_offset))
+ .sum(-1, dtype=torch.int16)
+ .sum(-1, dtype=torch.int32)
+ )
+ return intersections / unions
+
+
+def build_point_grid(n_per_side: int) -> np.ndarray:
+ """Generates a 2D grid of points evenly spaced in [0,1]x[0,1]."""
+ offset = 1 / (2 * n_per_side)
+ points_one_side = np.linspace(offset, 1 - offset, n_per_side)
+ points_x = np.tile(points_one_side[None, :], (n_per_side, 1))
+ points_y = np.tile(points_one_side[:, None], (1, n_per_side))
+ points = np.stack([points_x, points_y], axis=-1).reshape(-1, 2)
+ return points
+
+
+def build_all_layer_point_grids(
+ n_per_side: int, n_layers: int, scale_per_layer: int
+) -> List[np.ndarray]:
+ """Generates point grids for all crop layers."""
+ points_by_layer = []
+ for i in range(n_layers + 1):
+ n_points = int(n_per_side / (scale_per_layer**i))
+ points_by_layer.append(build_point_grid(n_points))
+ return points_by_layer
+
+
+def generate_crop_boxes(
+ im_size: Tuple[int, ...], n_layers: int, overlap_ratio: float
+) -> Tuple[List[List[int]], List[int]]:
+ """
+ Generates a list of crop boxes of different sizes. Each layer
+ has (2**i)**2 boxes for the ith layer.
+ """
+ crop_boxes, layer_idxs = [], []
+ im_h, im_w = im_size
+ short_side = min(im_h, im_w)
+
+ # Original image
+ crop_boxes.append([0, 0, im_w, im_h])
+ layer_idxs.append(0)
+
+ def crop_len(orig_len, n_crops, overlap):
+ return int(math.ceil((overlap * (n_crops - 1) + orig_len) / n_crops))
+
+ for i_layer in range(n_layers):
+ n_crops_per_side = 2 ** (i_layer + 1)
+ overlap = int(overlap_ratio * short_side * (2 / n_crops_per_side))
+
+ crop_w = crop_len(im_w, n_crops_per_side, overlap)
+ crop_h = crop_len(im_h, n_crops_per_side, overlap)
+
+ crop_box_x0 = [int((crop_w - overlap) * i) for i in range(n_crops_per_side)]
+ crop_box_y0 = [int((crop_h - overlap) * i) for i in range(n_crops_per_side)]
+
+ # Crops in XYWH format
+ for x0, y0 in product(crop_box_x0, crop_box_y0):
+ box = [x0, y0, min(x0 + crop_w, im_w), min(y0 + crop_h, im_h)]
+ crop_boxes.append(box)
+ layer_idxs.append(i_layer + 1)
+
+ return crop_boxes, layer_idxs
+
+
+def uncrop_boxes_xyxy(boxes: torch.Tensor, crop_box: List[int]) -> torch.Tensor:
+ x0, y0, _, _ = crop_box
+ offset = torch.tensor([[x0, y0, x0, y0]], device=boxes.device)
+ # Check if boxes has a channel dimension
+ if len(boxes.shape) == 3:
+ offset = offset.unsqueeze(1)
+ return boxes + offset
+
+
+def uncrop_points(points: torch.Tensor, crop_box: List[int]) -> torch.Tensor:
+ x0, y0, _, _ = crop_box
+ offset = torch.tensor([[x0, y0]], device=points.device)
+ # Check if points has a channel dimension
+ if len(points.shape) == 3:
+ offset = offset.unsqueeze(1)
+ return points + offset
+
+
+def uncrop_masks(
+ masks: torch.Tensor, crop_box: List[int], orig_h: int, orig_w: int
+) -> torch.Tensor:
+ x0, y0, x1, y1 = crop_box
+ if x0 == 0 and y0 == 0 and x1 == orig_w and y1 == orig_h:
+ return masks
+ # Coordinate transform masks
+ pad_x, pad_y = orig_w - (x1 - x0), orig_h - (y1 - y0)
+ pad = (x0, pad_x - x0, y0, pad_y - y0)
+ return torch.nn.functional.pad(masks, pad, value=0)
+
+
+def remove_small_regions(
+ mask: np.ndarray, area_thresh: float, mode: str
+) -> Tuple[np.ndarray, bool]:
+ """
+ Removes small disconnected regions and holes in a mask. Returns the
+ mask and an indicator of if the mask has been modified.
+ """
+ import cv2 # type: ignore
+
+ assert mode in ["holes", "islands"]
+ correct_holes = mode == "holes"
+ working_mask = (correct_holes ^ mask).astype(np.uint8)
+ n_labels, regions, stats, _ = cv2.connectedComponentsWithStats(working_mask, 8)
+ sizes = stats[:, -1][1:] # Row 0 is background label
+ small_regions = [i + 1 for i, s in enumerate(sizes) if s < area_thresh]
+ if len(small_regions) == 0:
+ return mask, False
+ fill_labels = [0] + small_regions
+ if not correct_holes:
+ fill_labels = [i for i in range(n_labels) if i not in fill_labels]
+ # If every region is below threshold, keep largest
+ if len(fill_labels) == 0:
+ fill_labels = [int(np.argmax(sizes)) + 1]
+ mask = np.isin(regions, fill_labels)
+ return mask, True
+
+
+def coco_encode_rle(uncompressed_rle: Dict[str, Any]) -> Dict[str, Any]:
+ from pycocotools import mask as mask_utils # type: ignore
+
+ h, w = uncompressed_rle["size"]
+ rle = mask_utils.frPyObjects(uncompressed_rle, h, w)
+ rle["counts"] = rle["counts"].decode("utf-8") # Necessary to serialize with json
+ return rle
+
+
+def batched_mask_to_box(masks: torch.Tensor) -> torch.Tensor:
+ """
+ Calculates boxes in XYXY format around masks. Return [0,0,0,0] for
+ an empty mask. For input shape C1xC2x...xHxW, the output shape is C1xC2x...x4.
+ """
+ # torch.max below raises an error on empty inputs, just skip in this case
+ if torch.numel(masks) == 0:
+ return torch.zeros(*masks.shape[:-2], 4, device=masks.device)
+
+ # Normalize shape to CxHxW
+ shape = masks.shape
+ h, w = shape[-2:]
+ if len(shape) > 2:
+ masks = masks.flatten(0, -3)
+ else:
+ masks = masks.unsqueeze(0)
+
+ # Get top and bottom edges
+ in_height, _ = torch.max(masks, dim=-1)
+ in_height_coords = in_height * torch.arange(h, device=in_height.device)[None, :]
+ bottom_edges, _ = torch.max(in_height_coords, dim=-1)
+ in_height_coords = in_height_coords + h * (~in_height)
+ top_edges, _ = torch.min(in_height_coords, dim=-1)
+
+ # Get left and right edges
+ in_width, _ = torch.max(masks, dim=-2)
+ in_width_coords = in_width * torch.arange(w, device=in_width.device)[None, :]
+ right_edges, _ = torch.max(in_width_coords, dim=-1)
+ in_width_coords = in_width_coords + w * (~in_width)
+ left_edges, _ = torch.min(in_width_coords, dim=-1)
+
+ # If the mask is empty the right edge will be to the left of the left edge.
+ # Replace these boxes with [0, 0, 0, 0]
+ empty_filter = (right_edges < left_edges) | (bottom_edges < top_edges)
+ out = torch.stack([left_edges, top_edges, right_edges, bottom_edges], dim=-1)
+ out = out * (~empty_filter).unsqueeze(-1)
+
+ # Return to original shape
+ if len(shape) > 2:
+ out = out.reshape(*shape[:-2], 4)
+ else:
+ out = out[0]
+
+ return out
diff --git a/py/evf_sam/model/segment_anything/utils/onnx.py b/py/evf_sam/model/segment_anything/utils/onnx.py
new file mode 100644
index 0000000..3521208
--- /dev/null
+++ b/py/evf_sam/model/segment_anything/utils/onnx.py
@@ -0,0 +1,157 @@
+# Copyright (c) Meta Platforms, Inc. and affiliates.
+# All rights reserved.
+
+# This source code is licensed under the license found in the
+# LICENSE file in the root directory of this source tree.
+
+from typing import Tuple
+
+import torch
+import torch.nn as nn
+from torch.nn import functional as F
+
+from ..modeling import Sam
+from .amg import calculate_stability_score
+
+
+class SamOnnxModel(nn.Module):
+ """
+ This model should not be called directly, but is used in ONNX export.
+ It combines the prompt encoder, mask decoder, and mask postprocessing of Sam,
+ with some functions modified to enable model tracing. Also supports extra
+ options controlling what information. See the ONNX export script for details.
+ """
+
+ def __init__(
+ self,
+ model: Sam,
+ return_single_mask: bool,
+ use_stability_score: bool = False,
+ return_extra_metrics: bool = False,
+ ) -> None:
+ super().__init__()
+ self.mask_decoder = model.mask_decoder
+ self.model = model
+ self.img_size = model.image_encoder.img_size
+ self.return_single_mask = return_single_mask
+ self.use_stability_score = use_stability_score
+ self.stability_score_offset = 1.0
+ self.return_extra_metrics = return_extra_metrics
+
+ @staticmethod
+ def resize_longest_image_size(
+ input_image_size: torch.Tensor, longest_side: int
+ ) -> torch.Tensor:
+ input_image_size = input_image_size.to(torch.float32)
+ scale = longest_side / torch.max(input_image_size)
+ transformed_size = scale * input_image_size
+ transformed_size = torch.floor(transformed_size + 0.5).to(torch.int64)
+ return transformed_size
+
+ def _embed_points(
+ self, point_coords: torch.Tensor, point_labels: torch.Tensor
+ ) -> torch.Tensor:
+ point_coords = point_coords + 0.5
+ point_coords = point_coords / self.img_size
+ point_embedding = self.model.prompt_encoder.pe_layer._pe_encoding(point_coords)
+ point_labels = point_labels.unsqueeze(-1).expand_as(point_embedding)
+
+ point_embedding = point_embedding * (point_labels != -1)
+ point_embedding = (
+ point_embedding
+ + self.model.prompt_encoder.not_a_point_embed.weight * (point_labels == -1)
+ )
+
+ for i in range(self.model.prompt_encoder.num_point_embeddings):
+ point_embedding = (
+ point_embedding
+ + self.model.prompt_encoder.point_embeddings[i].weight
+ * (point_labels == i)
+ )
+
+ return point_embedding
+
+ def _embed_masks(
+ self, input_mask: torch.Tensor, has_mask_input: torch.Tensor
+ ) -> torch.Tensor:
+ mask_embedding = has_mask_input * self.model.prompt_encoder.mask_downscaling(
+ input_mask
+ )
+ mask_embedding = mask_embedding + (
+ 1 - has_mask_input
+ ) * self.model.prompt_encoder.no_mask_embed.weight.reshape(1, -1, 1, 1)
+ return mask_embedding
+
+ def mask_postprocessing(
+ self, masks: torch.Tensor, orig_im_size: torch.Tensor
+ ) -> torch.Tensor:
+ masks = F.interpolate(
+ masks,
+ size=(self.img_size, self.img_size),
+ mode="bilinear",
+ align_corners=False,
+ )
+
+ prepadded_size = self.resize_longest_image_size(orig_im_size, self.img_size).to(
+ torch.int64
+ )
+ masks = masks[..., : prepadded_size[0], : prepadded_size[1]] # type: ignore
+
+ orig_im_size = orig_im_size.to(torch.int64)
+ h, w = orig_im_size[0], orig_im_size[1]
+ masks = F.interpolate(masks, size=(h, w), mode="bilinear", align_corners=False)
+ return masks
+
+ def select_masks(
+ self, masks: torch.Tensor, iou_preds: torch.Tensor, num_points: int
+ ) -> Tuple[torch.Tensor, torch.Tensor]:
+ # Determine if we should return the multiclick mask or not from the number of points.
+ # The reweighting is used to avoid control flow.
+ score_reweight = torch.tensor(
+ [[1000] + [0] * (self.model.mask_decoder.num_mask_tokens - 1)]
+ ).to(iou_preds.device)
+ score = iou_preds + (num_points - 2.5) * score_reweight
+ best_idx = torch.argmax(score, dim=1)
+ masks = masks[torch.arange(masks.shape[0]), best_idx, :, :].unsqueeze(1)
+ iou_preds = iou_preds[torch.arange(masks.shape[0]), best_idx].unsqueeze(1)
+
+ return masks, iou_preds
+
+ @torch.no_grad()
+ def forward(
+ self,
+ image_embeddings: torch.Tensor,
+ point_coords: torch.Tensor,
+ point_labels: torch.Tensor,
+ mask_input: torch.Tensor,
+ has_mask_input: torch.Tensor,
+ orig_im_size: torch.Tensor,
+ ):
+ sparse_embedding = self._embed_points(point_coords, point_labels)
+ dense_embedding = self._embed_masks(mask_input, has_mask_input)
+
+ masks, scores = self.model.mask_decoder.predict_masks(
+ image_embeddings=image_embeddings,
+ image_pe=self.model.prompt_encoder.get_dense_pe(),
+ sparse_prompt_embeddings=sparse_embedding,
+ dense_prompt_embeddings=dense_embedding,
+ )
+
+ if self.use_stability_score:
+ scores = calculate_stability_score(
+ masks, self.model.mask_threshold, self.stability_score_offset
+ )
+
+ if self.return_single_mask:
+ masks, scores = self.select_masks(masks, scores, point_coords.shape[1])
+
+ upscaled_masks = self.mask_postprocessing(masks, orig_im_size)
+
+ if self.return_extra_metrics:
+ stability_scores = calculate_stability_score(
+ upscaled_masks, self.model.mask_threshold, self.stability_score_offset
+ )
+ areas = (upscaled_masks > self.model.mask_threshold).sum(-1).sum(-1)
+ return upscaled_masks, scores, stability_scores, areas, masks
+
+ return upscaled_masks, scores, masks
diff --git a/py/evf_sam/model/segment_anything/utils/transforms.py b/py/evf_sam/model/segment_anything/utils/transforms.py
new file mode 100644
index 0000000..4232d84
--- /dev/null
+++ b/py/evf_sam/model/segment_anything/utils/transforms.py
@@ -0,0 +1,113 @@
+# Copyright (c) Meta Platforms, Inc. and affiliates.
+# All rights reserved.
+
+# This source code is licensed under the license found in the
+# LICENSE file in the root directory of this source tree.
+
+from copy import deepcopy
+from typing import Tuple
+
+import numpy as np
+import torch
+from torch.nn import functional as F
+from torchvision.transforms.functional import resize # type: ignore
+from torchvision.transforms.functional import to_pil_image
+
+
+class ResizeLongestSide:
+ """
+ Resizes images to the longest side 'target_length', as well as provides
+ methods for resizing coordinates and boxes. Provides methods for
+ transforming both numpy array and batched torch tensors.
+ """
+
+ def __init__(self, target_length: int) -> None:
+ self.target_length = target_length
+
+ def apply_image(self, image: np.ndarray) -> np.ndarray:
+ """
+ Expects a numpy array with shape HxWxC in uint8 format.
+ """
+ target_size = self.get_preprocess_shape(
+ image.shape[0], image.shape[1], self.target_length
+ )
+ return np.array(resize(to_pil_image(image), target_size))
+
+ def apply_coords(
+ self, coords: np.ndarray, original_size: Tuple[int, ...]
+ ) -> np.ndarray:
+ """
+ Expects a numpy array of length 2 in the final dimension. Requires the
+ original image size in (H, W) format.
+ """
+ old_h, old_w = original_size
+ new_h, new_w = self.get_preprocess_shape(
+ original_size[0], original_size[1], self.target_length
+ )
+ coords = deepcopy(coords).astype(float)
+ coords[..., 0] = coords[..., 0] * (new_w / old_w)
+ coords[..., 1] = coords[..., 1] * (new_h / old_h)
+ return coords
+
+ def apply_boxes(
+ self, boxes: np.ndarray, original_size: Tuple[int, ...]
+ ) -> np.ndarray:
+ """
+ Expects a numpy array shape Bx4. Requires the original image size
+ in (H, W) format.
+ """
+ boxes = self.apply_coords(boxes.reshape(-1, 2, 2), original_size)
+ return boxes.reshape(-1, 4)
+
+ def apply_image_torch(self, image: torch.Tensor) -> torch.Tensor:
+ """
+ Expects batched images with shape BxCxHxW and float format. This
+ transformation may not exactly match apply_image. apply_image is
+ the transformation expected by the model.
+ """
+ # Expects an image in BCHW format. May not exactly match apply_image.
+ target_size = self.get_preprocess_shape(
+ image.shape[0], image.shape[1], self.target_length
+ )
+ return F.interpolate(
+ image, target_size, mode="bilinear", align_corners=False, antialias=True
+ )
+
+ def apply_coords_torch(
+ self, coords: torch.Tensor, original_size: Tuple[int, ...]
+ ) -> torch.Tensor:
+ """
+ Expects a torch tensor with length 2 in the last dimension. Requires the
+ original image size in (H, W) format.
+ """
+ old_h, old_w = original_size
+ new_h, new_w = self.get_preprocess_shape(
+ original_size[0], original_size[1], self.target_length
+ )
+ coords = deepcopy(coords).to(torch.float)
+ coords[..., 0] = coords[..., 0] * (new_w / old_w)
+ coords[..., 1] = coords[..., 1] * (new_h / old_h)
+ return coords
+
+ def apply_boxes_torch(
+ self, boxes: torch.Tensor, original_size: Tuple[int, ...]
+ ) -> torch.Tensor:
+ """
+ Expects a torch tensor with shape Bx4. Requires the original image
+ size in (H, W) format.
+ """
+ boxes = self.apply_coords_torch(boxes.reshape(-1, 2, 2), original_size)
+ return boxes.reshape(-1, 4)
+
+ @staticmethod
+ def get_preprocess_shape(
+ oldh: int, oldw: int, long_side_length: int
+ ) -> Tuple[int, int]:
+ """
+ Compute the output size given input size and target long side length.
+ """
+ scale = long_side_length * 1.0 / max(oldh, oldw)
+ newh, neww = oldh * scale, oldw * scale
+ neww = int(neww + 0.5)
+ newh = int(newh + 0.5)
+ return (newh, neww)
diff --git a/py/evf_sam/model/segment_anything_2/sam2/__init__.py b/py/evf_sam/model/segment_anything_2/sam2/__init__.py
new file mode 100644
index 0000000..ad3406c
--- /dev/null
+++ b/py/evf_sam/model/segment_anything_2/sam2/__init__.py
@@ -0,0 +1,11 @@
+# Copyright (c) Meta Platforms, Inc. and affiliates.
+# All rights reserved.
+
+# This source code is licensed under the license found in the
+# LICENSE file in the root directory of this source tree.
+
+
+from hydra import initialize_config_module
+
+# initialize_config_module("model/segment_anything_2/sam2_configs", version_base="1.2")
+initialize_config_module("model/segment_anything_2/sam2_configs")
diff --git a/py/evf_sam/model/segment_anything_2/sam2/automatic_mask_generator.py b/py/evf_sam/model/segment_anything_2/sam2/automatic_mask_generator.py
new file mode 100644
index 0000000..91b76f1
--- /dev/null
+++ b/py/evf_sam/model/segment_anything_2/sam2/automatic_mask_generator.py
@@ -0,0 +1,434 @@
+# Copyright (c) Meta Platforms, Inc. and affiliates.
+# All rights reserved.
+
+# This source code is licensed under the license found in the
+# LICENSE file in the root directory of this source tree.
+
+# Adapted from https://github.com/facebookresearch/segment-anything/blob/main/segment_anything/automatic_mask_generator.py
+from typing import Any, Dict, List, Optional, Tuple
+
+import numpy as np
+import torch
+from torchvision.ops.boxes import batched_nms, box_area # type: ignore
+
+from model.segment_anything_2.sam2.modeling.sam2_base import SAM2Base
+from model.segment_anything_2.sam2.sam2_image_predictor import SAM2ImagePredictor
+from model.segment_anything_2.sam2.utils.amg import (
+ area_from_rle,
+ batch_iterator,
+ batched_mask_to_box,
+ box_xyxy_to_xywh,
+ build_all_layer_point_grids,
+ calculate_stability_score,
+ coco_encode_rle,
+ generate_crop_boxes,
+ is_box_near_crop_edge,
+ mask_to_rle_pytorch,
+ MaskData,
+ remove_small_regions,
+ rle_to_mask,
+ uncrop_boxes_xyxy,
+ uncrop_masks,
+ uncrop_points,
+)
+
+
+class SAM2AutomaticMaskGenerator:
+ def __init__(
+ self,
+ model: SAM2Base,
+ points_per_side: Optional[int] = 32,
+ points_per_batch: int = 64,
+ pred_iou_thresh: float = 0.8,
+ stability_score_thresh: float = 0.95,
+ stability_score_offset: float = 1.0,
+ mask_threshold: float = 0.0,
+ box_nms_thresh: float = 0.7,
+ crop_n_layers: int = 0,
+ crop_nms_thresh: float = 0.7,
+ crop_overlap_ratio: float = 512 / 1500,
+ crop_n_points_downscale_factor: int = 1,
+ point_grids: Optional[List[np.ndarray]] = None,
+ min_mask_region_area: int = 0,
+ output_mode: str = "binary_mask",
+ use_m2m: bool = False,
+ multimask_output: bool = True,
+ ) -> None:
+ """
+ Using a SAM 2 model, generates masks for the entire image.
+ Generates a grid of point prompts over the image, then filters
+ low quality and duplicate masks. The default settings are chosen
+ for SAM 2 with a HieraL backbone.
+
+ Arguments:
+ model (Sam): The SAM 2 model to use for mask prediction.
+ points_per_side (int or None): The number of points to be sampled
+ along one side of the image. The total number of points is
+ points_per_side**2. If None, 'point_grids' must provide explicit
+ point sampling.
+ points_per_batch (int): Sets the number of points run simultaneously
+ by the model. Higher numbers may be faster but use more GPU memory.
+ pred_iou_thresh (float): A filtering threshold in [0,1], using the
+ model's predicted mask quality.
+ stability_score_thresh (float): A filtering threshold in [0,1], using
+ the stability of the mask under changes to the cutoff used to binarize
+ the model's mask predictions.
+ stability_score_offset (float): The amount to shift the cutoff when
+ calculated the stability score.
+ mask_threshold (float): Threshold for binarizing the mask logits
+ box_nms_thresh (float): The box IoU cutoff used by non-maximal
+ suppression to filter duplicate masks.
+ crop_n_layers (int): If >0, mask prediction will be run again on
+ crops of the image. Sets the number of layers to run, where each
+ layer has 2**i_layer number of image crops.
+ crop_nms_thresh (float): The box IoU cutoff used by non-maximal
+ suppression to filter duplicate masks between different crops.
+ crop_overlap_ratio (float): Sets the degree to which crops overlap.
+ In the first crop layer, crops will overlap by this fraction of
+ the image length. Later layers with more crops scale down this overlap.
+ crop_n_points_downscale_factor (int): The number of points-per-side
+ sampled in layer n is scaled down by crop_n_points_downscale_factor**n.
+ point_grids (list(np.ndarray) or None): A list over explicit grids
+ of points used for sampling, normalized to [0,1]. The nth grid in the
+ list is used in the nth crop layer. Exclusive with points_per_side.
+ min_mask_region_area (int): If >0, postprocessing will be applied
+ to remove disconnected regions and holes in masks with area smaller
+ than min_mask_region_area. Requires opencv.
+ output_mode (str): The form masks are returned in. Can be 'binary_mask',
+ 'uncompressed_rle', or 'coco_rle'. 'coco_rle' requires pycocotools.
+ For large resolutions, 'binary_mask' may consume large amounts of
+ memory.
+ use_m2m (bool): Whether to add a one step refinement using previous mask predictions.
+ multimask_output (bool): Whether to output multimask at each point of the grid.
+ """
+
+ assert (points_per_side is None) != (
+ point_grids is None
+ ), "Exactly one of points_per_side or point_grid must be provided."
+ if points_per_side is not None:
+ self.point_grids = build_all_layer_point_grids(
+ points_per_side,
+ crop_n_layers,
+ crop_n_points_downscale_factor,
+ )
+ elif point_grids is not None:
+ self.point_grids = point_grids
+ else:
+ raise ValueError("Can't have both points_per_side and point_grid be None.")
+
+ assert output_mode in [
+ "binary_mask",
+ "uncompressed_rle",
+ "coco_rle",
+ ], f"Unknown output_mode {output_mode}."
+ if output_mode == "coco_rle":
+ try:
+ from pycocotools import mask as mask_utils # type: ignore # noqa: F401
+ except ImportError as e:
+ print("Please install pycocotools")
+ raise e
+
+ self.predictor = SAM2ImagePredictor(
+ model,
+ max_hole_area=min_mask_region_area,
+ max_sprinkle_area=min_mask_region_area,
+ )
+ self.points_per_batch = points_per_batch
+ self.pred_iou_thresh = pred_iou_thresh
+ self.stability_score_thresh = stability_score_thresh
+ self.stability_score_offset = stability_score_offset
+ self.mask_threshold = mask_threshold
+ self.box_nms_thresh = box_nms_thresh
+ self.crop_n_layers = crop_n_layers
+ self.crop_nms_thresh = crop_nms_thresh
+ self.crop_overlap_ratio = crop_overlap_ratio
+ self.crop_n_points_downscale_factor = crop_n_points_downscale_factor
+ self.min_mask_region_area = min_mask_region_area
+ self.output_mode = output_mode
+ self.use_m2m = use_m2m
+ self.multimask_output = multimask_output
+
+ @torch.no_grad()
+ def generate(self, image: np.ndarray) -> List[Dict[str, Any]]:
+ """
+ Generates masks for the given image.
+
+ Arguments:
+ image (np.ndarray): The image to generate masks for, in HWC uint8 format.
+
+ Returns:
+ list(dict(str, any)): A list over records for masks. Each record is
+ a dict containing the following keys:
+ segmentation (dict(str, any) or np.ndarray): The mask. If
+ output_mode='binary_mask', is an array of shape HW. Otherwise,
+ is a dictionary containing the RLE.
+ bbox (list(float)): The box around the mask, in XYWH format.
+ area (int): The area in pixels of the mask.
+ predicted_iou (float): The model's own prediction of the mask's
+ quality. This is filtered by the pred_iou_thresh parameter.
+ point_coords (list(list(float))): The point coordinates input
+ to the model to generate this mask.
+ stability_score (float): A measure of the mask's quality. This
+ is filtered on using the stability_score_thresh parameter.
+ crop_box (list(float)): The crop of the image used to generate
+ the mask, given in XYWH format.
+ """
+
+ # Generate masks
+ mask_data = self._generate_masks(image)
+
+ # Encode masks
+ if self.output_mode == "coco_rle":
+ mask_data["segmentations"] = [
+ coco_encode_rle(rle) for rle in mask_data["rles"]
+ ]
+ elif self.output_mode == "binary_mask":
+ mask_data["segmentations"] = [rle_to_mask(rle) for rle in mask_data["rles"]]
+ else:
+ mask_data["segmentations"] = mask_data["rles"]
+
+ # Write mask records
+ curr_anns = []
+ for idx in range(len(mask_data["segmentations"])):
+ ann = {
+ "segmentation": mask_data["segmentations"][idx],
+ "area": area_from_rle(mask_data["rles"][idx]),
+ "bbox": box_xyxy_to_xywh(mask_data["boxes"][idx]).tolist(),
+ "predicted_iou": mask_data["iou_preds"][idx].item(),
+ "point_coords": [mask_data["points"][idx].tolist()],
+ "stability_score": mask_data["stability_score"][idx].item(),
+ "crop_box": box_xyxy_to_xywh(mask_data["crop_boxes"][idx]).tolist(),
+ }
+ curr_anns.append(ann)
+
+ return curr_anns
+
+ def _generate_masks(self, image: np.ndarray) -> MaskData:
+ orig_size = image.shape[:2]
+ crop_boxes, layer_idxs = generate_crop_boxes(
+ orig_size, self.crop_n_layers, self.crop_overlap_ratio
+ )
+
+ # Iterate over image crops
+ data = MaskData()
+ for crop_box, layer_idx in zip(crop_boxes, layer_idxs):
+ crop_data = self._process_crop(image, crop_box, layer_idx, orig_size)
+ data.cat(crop_data)
+
+ # Remove duplicate masks between crops
+ if len(crop_boxes) > 1:
+ # Prefer masks from smaller crops
+ scores = 1 / box_area(data["crop_boxes"])
+ scores = scores.to(data["boxes"].device)
+ keep_by_nms = batched_nms(
+ data["boxes"].float(),
+ scores,
+ torch.zeros_like(data["boxes"][:, 0]), # categories
+ iou_threshold=self.crop_nms_thresh,
+ )
+ data.filter(keep_by_nms)
+ data.to_numpy()
+ return data
+
+ def _process_crop(
+ self,
+ image: np.ndarray,
+ crop_box: List[int],
+ crop_layer_idx: int,
+ orig_size: Tuple[int, ...],
+ ) -> MaskData:
+ # Crop the image and calculate embeddings
+ x0, y0, x1, y1 = crop_box
+ cropped_im = image[y0:y1, x0:x1, :]
+ cropped_im_size = cropped_im.shape[:2]
+ self.predictor.set_image(cropped_im)
+
+ # Get points for this crop
+ points_scale = np.array(cropped_im_size)[None, ::-1]
+ points_for_image = self.point_grids[crop_layer_idx] * points_scale
+
+ # Generate masks for this crop in batches
+ data = MaskData()
+ for (points,) in batch_iterator(self.points_per_batch, points_for_image):
+ batch_data = self._process_batch(
+ points, cropped_im_size, crop_box, orig_size, normalize=True
+ )
+ data.cat(batch_data)
+ del batch_data
+ self.predictor.reset_predictor()
+
+ # Remove duplicates within this crop.
+ keep_by_nms = batched_nms(
+ data["boxes"].float(),
+ data["iou_preds"],
+ torch.zeros_like(data["boxes"][:, 0]), # categories
+ iou_threshold=self.box_nms_thresh,
+ )
+ data.filter(keep_by_nms)
+
+ # Return to the original image frame
+ data["boxes"] = uncrop_boxes_xyxy(data["boxes"], crop_box)
+ data["points"] = uncrop_points(data["points"], crop_box)
+ data["crop_boxes"] = torch.tensor([crop_box for _ in range(len(data["rles"]))])
+
+ return data
+
+ def _process_batch(
+ self,
+ points: np.ndarray,
+ im_size: Tuple[int, ...],
+ crop_box: List[int],
+ orig_size: Tuple[int, ...],
+ normalize=False,
+ ) -> MaskData:
+ orig_h, orig_w = orig_size
+
+ # Run model on this batch
+ points = torch.as_tensor(points, device=self.predictor.device)
+ in_points = self.predictor._transforms.transform_coords(
+ points, normalize=normalize, orig_hw=im_size
+ )
+ in_labels = torch.ones(
+ in_points.shape[0], dtype=torch.int, device=in_points.device
+ )
+ masks, iou_preds, low_res_masks = self.predictor._predict(
+ in_points[:, None, :],
+ in_labels[:, None],
+ multimask_output=self.multimask_output,
+ return_logits=True,
+ )
+
+ # Serialize predictions and store in MaskData
+ data = MaskData(
+ masks=masks.flatten(0, 1),
+ iou_preds=iou_preds.flatten(0, 1),
+ points=points.repeat_interleave(masks.shape[1], dim=0),
+ low_res_masks=low_res_masks.flatten(0, 1),
+ )
+ del masks
+
+ if not self.use_m2m:
+ # Filter by predicted IoU
+ if self.pred_iou_thresh > 0.0:
+ keep_mask = data["iou_preds"] > self.pred_iou_thresh
+ data.filter(keep_mask)
+
+ # Calculate and filter by stability score
+ data["stability_score"] = calculate_stability_score(
+ data["masks"], self.mask_threshold, self.stability_score_offset
+ )
+ if self.stability_score_thresh > 0.0:
+ keep_mask = data["stability_score"] >= self.stability_score_thresh
+ data.filter(keep_mask)
+ else:
+ # One step refinement using previous mask predictions
+ in_points = self.predictor._transforms.transform_coords(
+ data["points"], normalize=normalize, orig_hw=im_size
+ )
+ labels = torch.ones(
+ in_points.shape[0], dtype=torch.int, device=in_points.device
+ )
+ masks, ious = self.refine_with_m2m(
+ in_points, labels, data["low_res_masks"], self.points_per_batch
+ )
+ data["masks"] = masks.squeeze(1)
+ data["iou_preds"] = ious.squeeze(1)
+
+ if self.pred_iou_thresh > 0.0:
+ keep_mask = data["iou_preds"] > self.pred_iou_thresh
+ data.filter(keep_mask)
+
+ data["stability_score"] = calculate_stability_score(
+ data["masks"], self.mask_threshold, self.stability_score_offset
+ )
+ if self.stability_score_thresh > 0.0:
+ keep_mask = data["stability_score"] >= self.stability_score_thresh
+ data.filter(keep_mask)
+
+ # Threshold masks and calculate boxes
+ data["masks"] = data["masks"] > self.mask_threshold
+ data["boxes"] = batched_mask_to_box(data["masks"])
+
+ # Filter boxes that touch crop boundaries
+ keep_mask = ~is_box_near_crop_edge(
+ data["boxes"], crop_box, [0, 0, orig_w, orig_h]
+ )
+ if not torch.all(keep_mask):
+ data.filter(keep_mask)
+
+ # Compress to RLE
+ data["masks"] = uncrop_masks(data["masks"], crop_box, orig_h, orig_w)
+ data["rles"] = mask_to_rle_pytorch(data["masks"])
+ del data["masks"]
+
+ return data
+
+ @staticmethod
+ def postprocess_small_regions(
+ mask_data: MaskData, min_area: int, nms_thresh: float
+ ) -> MaskData:
+ """
+ Removes small disconnected regions and holes in masks, then reruns
+ box NMS to remove any new duplicates.
+
+ Edits mask_data in place.
+
+ Requires open-cv as a dependency.
+ """
+ if len(mask_data["rles"]) == 0:
+ return mask_data
+
+ # Filter small disconnected regions and holes
+ new_masks = []
+ scores = []
+ for rle in mask_data["rles"]:
+ mask = rle_to_mask(rle)
+
+ mask, changed = remove_small_regions(mask, min_area, mode="holes")
+ unchanged = not changed
+ mask, changed = remove_small_regions(mask, min_area, mode="islands")
+ unchanged = unchanged and not changed
+
+ new_masks.append(torch.as_tensor(mask).unsqueeze(0))
+ # Give score=0 to changed masks and score=1 to unchanged masks
+ # so NMS will prefer ones that didn't need postprocessing
+ scores.append(float(unchanged))
+
+ # Recalculate boxes and remove any new duplicates
+ masks = torch.cat(new_masks, dim=0)
+ boxes = batched_mask_to_box(masks)
+ keep_by_nms = batched_nms(
+ boxes.float(),
+ torch.as_tensor(scores),
+ torch.zeros_like(boxes[:, 0]), # categories
+ iou_threshold=nms_thresh,
+ )
+
+ # Only recalculate RLEs for masks that have changed
+ for i_mask in keep_by_nms:
+ if scores[i_mask] == 0.0:
+ mask_torch = masks[i_mask].unsqueeze(0)
+ mask_data["rles"][i_mask] = mask_to_rle_pytorch(mask_torch)[0]
+ mask_data["boxes"][i_mask] = boxes[i_mask] # update res directly
+ mask_data.filter(keep_by_nms)
+
+ return mask_data
+
+ def refine_with_m2m(self, points, point_labels, low_res_masks, points_per_batch):
+ new_masks = []
+ new_iou_preds = []
+
+ for cur_points, cur_point_labels, low_res_mask in batch_iterator(
+ points_per_batch, points, point_labels, low_res_masks
+ ):
+ best_masks, best_iou_preds, _ = self.predictor._predict(
+ cur_points[:, None, :],
+ cur_point_labels[:, None],
+ mask_input=low_res_mask[:, None, :],
+ multimask_output=False,
+ return_logits=True,
+ )
+ new_masks.append(best_masks)
+ new_iou_preds.append(best_iou_preds)
+ masks = torch.cat(new_masks, dim=0)
+ return masks, torch.cat(new_iou_preds, dim=0)
diff --git a/py/evf_sam/model/segment_anything_2/sam2/build_sam.py b/py/evf_sam/model/segment_anything_2/sam2/build_sam.py
new file mode 100644
index 0000000..c0cfe12
--- /dev/null
+++ b/py/evf_sam/model/segment_anything_2/sam2/build_sam.py
@@ -0,0 +1,90 @@
+# Copyright (c) Meta Platforms, Inc. and affiliates.
+# All rights reserved.
+
+# This source code is licensed under the license found in the
+# LICENSE file in the root directory of this source tree.
+
+import logging
+
+import torch
+from hydra import compose
+from hydra.utils import instantiate
+from omegaconf import OmegaConf
+
+def build_sam2(
+ config_file,
+ ckpt_path=None,
+ device="cuda",
+ mode="eval",
+ hydra_overrides_extra=[],
+ apply_postprocessing=True,
+):
+
+ if apply_postprocessing:
+ hydra_overrides_extra = hydra_overrides_extra.copy()
+ hydra_overrides_extra += [
+ # dynamically fall back to multi-mask if the single mask is not stable
+ "++model.sam_mask_decoder_extra_args.dynamic_multimask_via_stability=true",
+ "++model.sam_mask_decoder_extra_args.dynamic_multimask_stability_delta=0.05",
+ "++model.sam_mask_decoder_extra_args.dynamic_multimask_stability_thresh=0.98",
+ ]
+ # Read config and init model
+ cfg = compose(config_name=config_file, overrides=hydra_overrides_extra)
+ OmegaConf.resolve(cfg)
+ model = instantiate(cfg.model, _recursive_=True)
+ _load_checkpoint(model, ckpt_path)
+ if device:
+ model = model.to(device)
+ if mode == "eval":
+ model.eval()
+ return model
+
+
+def build_sam2_video_predictor(
+ config_file,
+ ckpt_path=None,
+ device="cuda",
+ mode="eval",
+ hydra_overrides_extra=[],
+ apply_postprocessing=True,
+):
+ hydra_overrides = [
+ "++model._target_=model.segment_anything_2.sam2.sam2_video_predictor.SAM2VideoPredictor",
+ ]
+ if apply_postprocessing:
+ hydra_overrides_extra = hydra_overrides_extra.copy()
+ hydra_overrides_extra += [
+ # dynamically fall back to multi-mask if the single mask is not stable
+ "++model.sam_mask_decoder_extra_args.dynamic_multimask_via_stability=true",
+ "++model.sam_mask_decoder_extra_args.dynamic_multimask_stability_delta=0.05",
+ "++model.sam_mask_decoder_extra_args.dynamic_multimask_stability_thresh=0.98",
+ # the sigmoid mask logits on interacted frames with clicks in the memory encoder so that the encoded masks are exactly as what users see from clicking
+ "++model.binarize_mask_from_pts_for_mem_enc=true",
+ # fill small holes in the low-res masks up to `fill_hole_area` (before resizing them to the original video resolution)
+ "++model.fill_hole_area=8",
+ ]
+ hydra_overrides.extend(hydra_overrides_extra)
+
+ # Read config and init model
+ cfg = compose(config_name=config_file, overrides=hydra_overrides)
+ OmegaConf.resolve(cfg)
+ model = instantiate(cfg.model, _recursive_=True)
+ _load_checkpoint(model, ckpt_path)
+ if device:
+ model = model.to(device)
+ if mode == "eval":
+ model.eval()
+ return model
+
+
+def _load_checkpoint(model, ckpt_path):
+ if ckpt_path is not None:
+ sd = torch.load(ckpt_path, map_location="cpu")["model"]
+ missing_keys, unexpected_keys = model.load_state_dict(sd)
+ if missing_keys:
+ logging.error(missing_keys)
+ raise RuntimeError()
+ if unexpected_keys:
+ logging.error(unexpected_keys)
+ raise RuntimeError()
+ logging.info("Loaded checkpoint sucessfully")
diff --git a/py/evf_sam/model/segment_anything_2/sam2/csrc/connected_components.cu b/py/evf_sam/model/segment_anything_2/sam2/csrc/connected_components.cu
new file mode 100644
index 0000000..eb83231
--- /dev/null
+++ b/py/evf_sam/model/segment_anything_2/sam2/csrc/connected_components.cu
@@ -0,0 +1,289 @@
+// Copyright (c) Meta Platforms, Inc. and affiliates.
+// All rights reserved.
+
+// This source code is licensed under the license found in the
+// LICENSE file in the root directory of this source tree.
+
+// adapted from https://github.com/zsef123/Connected_components_PyTorch
+// with license found in the LICENSE_cctorch file in the root directory.
+#include
+#include
+#include
+#include
+#include
+#include
+
+// 2d
+#define BLOCK_ROWS 16
+#define BLOCK_COLS 16
+
+namespace cc2d {
+
+template
+__device__ __forceinline__ unsigned char hasBit(T bitmap, unsigned char pos) {
+ return (bitmap >> pos) & 1;
+}
+
+__device__ int32_t find(const int32_t* s_buf, int32_t n) {
+ while (s_buf[n] != n)
+ n = s_buf[n];
+ return n;
+}
+
+__device__ int32_t find_n_compress(int32_t* s_buf, int32_t n) {
+ const int32_t id = n;
+ while (s_buf[n] != n) {
+ n = s_buf[n];
+ s_buf[id] = n;
+ }
+ return n;
+}
+
+__device__ void union_(int32_t* s_buf, int32_t a, int32_t b) {
+ bool done;
+ do {
+ a = find(s_buf, a);
+ b = find(s_buf, b);
+
+ if (a < b) {
+ int32_t old = atomicMin(s_buf + b, a);
+ done = (old == b);
+ b = old;
+ } else if (b < a) {
+ int32_t old = atomicMin(s_buf + a, b);
+ done = (old == a);
+ a = old;
+ } else
+ done = true;
+
+ } while (!done);
+}
+
+__global__ void
+init_labeling(int32_t* label, const uint32_t W, const uint32_t H) {
+ const uint32_t row = (blockIdx.y * blockDim.y + threadIdx.y) * 2;
+ const uint32_t col = (blockIdx.x * blockDim.x + threadIdx.x) * 2;
+ const uint32_t idx = row * W + col;
+
+ if (row < H && col < W)
+ label[idx] = idx;
+}
+
+__global__ void
+merge(uint8_t* img, int32_t* label, const uint32_t W, const uint32_t H) {
+ const uint32_t row = (blockIdx.y * blockDim.y + threadIdx.y) * 2;
+ const uint32_t col = (blockIdx.x * blockDim.x + threadIdx.x) * 2;
+ const uint32_t idx = row * W + col;
+
+ if (row >= H || col >= W)
+ return;
+
+ uint32_t P = 0;
+
+ if (img[idx])
+ P |= 0x777;
+ if (row + 1 < H && img[idx + W])
+ P |= 0x777 << 4;
+ if (col + 1 < W && img[idx + 1])
+ P |= 0x777 << 1;
+
+ if (col == 0)
+ P &= 0xEEEE;
+ if (col + 1 >= W)
+ P &= 0x3333;
+ else if (col + 2 >= W)
+ P &= 0x7777;
+
+ if (row == 0)
+ P &= 0xFFF0;
+ if (row + 1 >= H)
+ P &= 0xFF;
+
+ if (P > 0) {
+ // If need check about top-left pixel(if flag the first bit) and hit the
+ // top-left pixel
+ if (hasBit(P, 0) && img[idx - W - 1]) {
+ union_(label, idx, idx - 2 * W - 2); // top left block
+ }
+
+ if ((hasBit(P, 1) && img[idx - W]) || (hasBit(P, 2) && img[idx - W + 1]))
+ union_(label, idx, idx - 2 * W); // top bottom block
+
+ if (hasBit(P, 3) && img[idx + 2 - W])
+ union_(label, idx, idx - 2 * W + 2); // top right block
+
+ if ((hasBit(P, 4) && img[idx - 1]) || (hasBit(P, 8) && img[idx + W - 1]))
+ union_(label, idx, idx - 2); // just left block
+ }
+}
+
+__global__ void compression(int32_t* label, const int32_t W, const int32_t H) {
+ const uint32_t row = (blockIdx.y * blockDim.y + threadIdx.y) * 2;
+ const uint32_t col = (blockIdx.x * blockDim.x + threadIdx.x) * 2;
+ const uint32_t idx = row * W + col;
+
+ if (row < H && col < W)
+ find_n_compress(label, idx);
+}
+
+__global__ void final_labeling(
+ const uint8_t* img,
+ int32_t* label,
+ const int32_t W,
+ const int32_t H) {
+ const uint32_t row = (blockIdx.y * blockDim.y + threadIdx.y) * 2;
+ const uint32_t col = (blockIdx.x * blockDim.x + threadIdx.x) * 2;
+ const uint32_t idx = row * W + col;
+
+ if (row >= H || col >= W)
+ return;
+
+ int32_t y = label[idx] + 1;
+
+ if (img[idx])
+ label[idx] = y;
+ else
+ label[idx] = 0;
+
+ if (col + 1 < W) {
+ if (img[idx + 1])
+ label[idx + 1] = y;
+ else
+ label[idx + 1] = 0;
+
+ if (row + 1 < H) {
+ if (img[idx + W + 1])
+ label[idx + W + 1] = y;
+ else
+ label[idx + W + 1] = 0;
+ }
+ }
+
+ if (row + 1 < H) {
+ if (img[idx + W])
+ label[idx + W] = y;
+ else
+ label[idx + W] = 0;
+ }
+}
+
+__global__ void init_counting(
+ const int32_t* label,
+ int32_t* count_init,
+ const int32_t W,
+ const int32_t H) {
+ const uint32_t row = (blockIdx.y * blockDim.y + threadIdx.y);
+ const uint32_t col = (blockIdx.x * blockDim.x + threadIdx.x);
+ const uint32_t idx = row * W + col;
+
+ if (row >= H || col >= W)
+ return;
+
+ int32_t y = label[idx];
+ if (y > 0) {
+ int32_t count_idx = y - 1;
+ atomicAdd(count_init + count_idx, 1);
+ }
+}
+
+__global__ void final_counting(
+ const int32_t* label,
+ const int32_t* count_init,
+ int32_t* count_final,
+ const int32_t W,
+ const int32_t H) {
+ const uint32_t row = (blockIdx.y * blockDim.y + threadIdx.y);
+ const uint32_t col = (blockIdx.x * blockDim.x + threadIdx.x);
+ const uint32_t idx = row * W + col;
+
+ if (row >= H || col >= W)
+ return;
+
+ int32_t y = label[idx];
+ if (y > 0) {
+ int32_t count_idx = y - 1;
+ count_final[idx] = count_init[count_idx];
+ } else {
+ count_final[idx] = 0;
+ }
+}
+
+} // namespace cc2d
+
+std::vector get_connected_componnets(
+ const torch::Tensor& inputs) {
+ AT_ASSERTM(inputs.is_cuda(), "inputs must be a CUDA tensor");
+ AT_ASSERTM(inputs.ndimension() == 4, "inputs must be [N, 1, H, W] shape");
+ AT_ASSERTM(
+ inputs.scalar_type() == torch::kUInt8, "inputs must be a uint8 type");
+
+ const uint32_t N = inputs.size(0);
+ const uint32_t C = inputs.size(1);
+ const uint32_t H = inputs.size(2);
+ const uint32_t W = inputs.size(3);
+
+ AT_ASSERTM(C == 1, "inputs must be [N, 1, H, W] shape");
+ AT_ASSERTM((H % 2) == 0, "height must be a even number");
+ AT_ASSERTM((W % 2) == 0, "width must be a even number");
+
+ // label must be uint32_t
+ auto label_options =
+ torch::TensorOptions().dtype(torch::kInt32).device(inputs.device());
+ torch::Tensor labels = torch::zeros({N, C, H, W}, label_options);
+ torch::Tensor counts_init = torch::zeros({N, C, H, W}, label_options);
+ torch::Tensor counts_final = torch::zeros({N, C, H, W}, label_options);
+
+ dim3 grid = dim3(
+ ((W + 1) / 2 + BLOCK_COLS - 1) / BLOCK_COLS,
+ ((H + 1) / 2 + BLOCK_ROWS - 1) / BLOCK_ROWS);
+ dim3 block = dim3(BLOCK_COLS, BLOCK_ROWS);
+ dim3 grid_count =
+ dim3((W + BLOCK_COLS) / BLOCK_COLS, (H + BLOCK_ROWS) / BLOCK_ROWS);
+ dim3 block_count = dim3(BLOCK_COLS, BLOCK_ROWS);
+ cudaStream_t stream = at::cuda::getCurrentCUDAStream();
+
+ for (int n = 0; n < N; n++) {
+ uint32_t offset = n * H * W;
+
+ cc2d::init_labeling<<>>(
+ labels.data_ptr() + offset, W, H);
+ cc2d::merge<<>>(
+ inputs.data_ptr() + offset,
+ labels.data_ptr() + offset,
+ W,
+ H);
+ cc2d::compression<<>>(
+ labels.data_ptr() + offset, W, H);
+ cc2d::final_labeling<<>>(
+ inputs.data_ptr() + offset,
+ labels.data_ptr() + offset,
+ W,
+ H);
+
+ // get the counting of each pixel
+ cc2d::init_counting<<>>(
+ labels.data_ptr() + offset,
+ counts_init.data_ptr() + offset,
+ W,
+ H);
+ cc2d::final_counting<<>>(
+ labels.data_ptr() + offset,
+ counts_init.data_ptr() + offset,
+ counts_final.data_ptr() + offset,
+ W,
+ H);
+ }
+
+ // returned values are [labels, counts]
+ std::vector outputs;
+ outputs.push_back(labels);
+ outputs.push_back(counts_final);
+ return outputs;
+}
+
+PYBIND11_MODULE(TORCH_EXTENSION_NAME, m) {
+ m.def(
+ "get_connected_componnets",
+ &get_connected_componnets,
+ "get_connected_componnets");
+}
diff --git a/py/evf_sam/model/segment_anything_2/sam2/modeling/__init__.py b/py/evf_sam/model/segment_anything_2/sam2/modeling/__init__.py
new file mode 100644
index 0000000..5277f46
--- /dev/null
+++ b/py/evf_sam/model/segment_anything_2/sam2/modeling/__init__.py
@@ -0,0 +1,5 @@
+# Copyright (c) Meta Platforms, Inc. and affiliates.
+# All rights reserved.
+
+# This source code is licensed under the license found in the
+# LICENSE file in the root directory of this source tree.
diff --git a/py/evf_sam/model/segment_anything_2/sam2/modeling/backbones/__init__.py b/py/evf_sam/model/segment_anything_2/sam2/modeling/backbones/__init__.py
new file mode 100644
index 0000000..5277f46
--- /dev/null
+++ b/py/evf_sam/model/segment_anything_2/sam2/modeling/backbones/__init__.py
@@ -0,0 +1,5 @@
+# Copyright (c) Meta Platforms, Inc. and affiliates.
+# All rights reserved.
+
+# This source code is licensed under the license found in the
+# LICENSE file in the root directory of this source tree.
diff --git a/py/evf_sam/model/segment_anything_2/sam2/modeling/backbones/hieradet.py b/py/evf_sam/model/segment_anything_2/sam2/modeling/backbones/hieradet.py
new file mode 100644
index 0000000..407598e
--- /dev/null
+++ b/py/evf_sam/model/segment_anything_2/sam2/modeling/backbones/hieradet.py
@@ -0,0 +1,295 @@
+# Copyright (c) Meta Platforms, Inc. and affiliates.
+# All rights reserved.
+
+# This source code is licensed under the license found in the
+# LICENSE file in the root directory of this source tree.
+
+from functools import partial
+from typing import List, Tuple, Union
+
+import torch
+import torch.nn as nn
+import torch.nn.functional as F
+
+from model.segment_anything_2.sam2.modeling.backbones.utils import (
+ PatchEmbed,
+ window_partition,
+ window_unpartition,
+)
+
+from model.segment_anything_2.sam2.modeling.sam2_utils import DropPath, MLP
+
+
+def do_pool(x: torch.Tensor, pool: nn.Module, norm: nn.Module = None) -> torch.Tensor:
+ if pool is None:
+ return x
+ # (B, H, W, C) -> (B, C, H, W)
+ x = x.permute(0, 3, 1, 2)
+ x = pool(x)
+ # (B, C, H', W') -> (B, H', W', C)
+ x = x.permute(0, 2, 3, 1)
+ if norm:
+ x = norm(x)
+
+ return x
+
+
+class MultiScaleAttention(nn.Module):
+ def __init__(
+ self,
+ dim: int,
+ dim_out: int,
+ num_heads: int,
+ q_pool: nn.Module = None,
+ ):
+ super().__init__()
+
+ self.dim = dim
+ self.dim_out = dim_out
+
+ self.num_heads = num_heads
+ head_dim = dim_out // num_heads
+ self.scale = head_dim**-0.5
+
+ self.q_pool = q_pool
+ self.qkv = nn.Linear(dim, dim_out * 3)
+ self.proj = nn.Linear(dim_out, dim_out)
+
+ def forward(self, x: torch.Tensor) -> torch.Tensor:
+ B, H, W, _ = x.shape
+ # qkv with shape (B, H * W, 3, nHead, C)
+ qkv = self.qkv(x).reshape(B, H * W, 3, self.num_heads, -1)
+ # q, k, v with shape (B, H * W, nheads, C)
+ q, k, v = torch.unbind(qkv, 2)
+
+ # Q pooling (for downsample at stage changes)
+ if self.q_pool:
+ q = do_pool(q.reshape(B, H, W, -1), self.q_pool)
+ H, W = q.shape[1:3] # downsampled shape
+ q = q.reshape(B, H * W, self.num_heads, -1)
+
+ # Torch's SDPA expects [B, nheads, H*W, C] so we transpose
+ x = F.scaled_dot_product_attention(
+ q.transpose(1, 2),
+ k.transpose(1, 2),
+ v.transpose(1, 2),
+ )
+ # Transpose back
+ x = x.transpose(1, 2)
+ x = x.reshape(B, H, W, -1)
+
+ x = self.proj(x)
+
+ return x
+
+
+class MultiScaleBlock(nn.Module):
+ def __init__(
+ self,
+ dim: int,
+ dim_out: int,
+ num_heads: int,
+ mlp_ratio: float = 4.0,
+ drop_path: float = 0.0,
+ norm_layer: Union[nn.Module, str] = "LayerNorm",
+ q_stride: Tuple[int, int] = None,
+ act_layer: nn.Module = nn.GELU,
+ window_size: int = 0,
+ ):
+ super().__init__()
+
+ if isinstance(norm_layer, str):
+ norm_layer = partial(getattr(nn, norm_layer), eps=1e-6)
+
+ self.dim = dim
+ self.dim_out = dim_out
+ self.norm1 = norm_layer(dim)
+
+ self.window_size = window_size
+
+ self.pool, self.q_stride = None, q_stride
+ if self.q_stride:
+ self.pool = nn.MaxPool2d(
+ kernel_size=q_stride, stride=q_stride, ceil_mode=False
+ )
+
+ self.attn = MultiScaleAttention(
+ dim,
+ dim_out,
+ num_heads=num_heads,
+ q_pool=self.pool,
+ )
+ self.drop_path = DropPath(drop_path) if drop_path > 0.0 else nn.Identity()
+
+ self.norm2 = norm_layer(dim_out)
+ self.mlp = MLP(
+ dim_out,
+ int(dim_out * mlp_ratio),
+ dim_out,
+ num_layers=2,
+ activation=act_layer,
+ )
+
+ if dim != dim_out:
+ self.proj = nn.Linear(dim, dim_out)
+
+ def forward(self, x: torch.Tensor) -> torch.Tensor:
+ shortcut = x # B, H, W, C
+ x = self.norm1(x)
+
+ # Skip connection
+ if self.dim != self.dim_out:
+ shortcut = do_pool(self.proj(x), self.pool)
+
+ # Window partition
+ window_size = self.window_size
+ if window_size > 0:
+ H, W = x.shape[1], x.shape[2]
+ x, pad_hw = window_partition(x, window_size)
+
+ # Window Attention + Q Pooling (if stage change)
+ x = self.attn(x)
+ if self.q_stride:
+ # Shapes have changed due to Q pooling
+ window_size = self.window_size // self.q_stride[0]
+ H, W = shortcut.shape[1:3]
+
+ pad_h = (window_size - H % window_size) % window_size
+ pad_w = (window_size - W % window_size) % window_size
+ pad_hw = (H + pad_h, W + pad_w)
+
+ # Reverse window partition
+ if self.window_size > 0:
+ x = window_unpartition(x, window_size, pad_hw, (H, W))
+
+ x = shortcut + self.drop_path(x)
+ # MLP
+ x = x + self.drop_path(self.mlp(self.norm2(x)))
+ return x
+
+
+class Hiera(nn.Module):
+ """
+ Reference: https://arxiv.org/abs/2306.00989
+ """
+
+ def __init__(
+ self,
+ embed_dim: int = 96, # initial embed dim
+ num_heads: int = 1, # initial number of heads
+ drop_path_rate: float = 0.0, # stochastic depth
+ q_pool: int = 3, # number of q_pool stages
+ q_stride: Tuple[int, int] = (2, 2), # downsample stride bet. stages
+ stages: Tuple[int, ...] = (2, 3, 16, 3), # blocks per stage
+ dim_mul: float = 2.0, # dim_mul factor at stage shift
+ head_mul: float = 2.0, # head_mul factor at stage shift
+ window_pos_embed_bkg_spatial_size: Tuple[int, int] = (14, 14),
+ # window size per stage, when not using global att.
+ window_spec: Tuple[int, ...] = (
+ 8,
+ 4,
+ 14,
+ 7,
+ ),
+ # global attn in these blocks
+ global_att_blocks: Tuple[int, ...] = (
+ 12,
+ 16,
+ 20,
+ ),
+ return_interm_layers=True, # return feats from every stage
+ ):
+ super().__init__()
+
+ assert len(stages) == len(window_spec)
+ self.window_spec = window_spec
+
+ depth = sum(stages)
+ self.q_stride = q_stride
+ self.stage_ends = [sum(stages[:i]) - 1 for i in range(1, len(stages) + 1)]
+ assert 0 <= q_pool <= len(self.stage_ends[:-1])
+ self.q_pool_blocks = [x + 1 for x in self.stage_ends[:-1]][:q_pool]
+ self.return_interm_layers = return_interm_layers
+
+ self.patch_embed = PatchEmbed(
+ embed_dim=embed_dim,
+ )
+ # Which blocks have global att?
+ self.global_att_blocks = global_att_blocks
+
+ # Windowed positional embedding (https://arxiv.org/abs/2311.05613)
+ self.window_pos_embed_bkg_spatial_size = window_pos_embed_bkg_spatial_size
+ self.pos_embed = nn.Parameter(
+ torch.zeros(1, embed_dim, *self.window_pos_embed_bkg_spatial_size)
+ )
+ self.pos_embed_window = nn.Parameter(
+ torch.zeros(1, embed_dim, self.window_spec[0], self.window_spec[0])
+ )
+
+ dpr = [
+ x.item() for x in torch.linspace(0, drop_path_rate, depth)
+ ] # stochastic depth decay rule
+
+ cur_stage = 1
+ self.blocks = nn.ModuleList()
+
+ for i in range(depth):
+ dim_out = embed_dim
+ # lags by a block, so first block of
+ # next stage uses an initial window size
+ # of previous stage and final window size of current stage
+ window_size = self.window_spec[cur_stage - 1]
+
+ if self.global_att_blocks is not None:
+ window_size = 0 if i in self.global_att_blocks else window_size
+
+ if i - 1 in self.stage_ends:
+ dim_out = int(embed_dim * dim_mul)
+ num_heads = int(num_heads * head_mul)
+ cur_stage += 1
+
+ block = MultiScaleBlock(
+ dim=embed_dim,
+ dim_out=dim_out,
+ num_heads=num_heads,
+ drop_path=dpr[i],
+ q_stride=self.q_stride if i in self.q_pool_blocks else None,
+ window_size=window_size,
+ )
+
+ embed_dim = dim_out
+ self.blocks.append(block)
+
+ self.channel_list = (
+ [self.blocks[i].dim_out for i in self.stage_ends[::-1]]
+ if return_interm_layers
+ else [self.blocks[-1].dim_out]
+ )
+
+ def _get_pos_embed(self, hw: Tuple[int, int]) -> torch.Tensor:
+ h, w = hw
+ window_embed = self.pos_embed_window
+ pos_embed = F.interpolate(self.pos_embed, size=(h, w), mode="bicubic")
+ pos_embed = pos_embed + window_embed.tile(
+ [x // y for x, y in zip(pos_embed.shape, window_embed.shape)]
+ )
+ pos_embed = pos_embed.permute(0, 2, 3, 1)
+ return pos_embed
+
+ def forward(self, x: torch.Tensor) -> List[torch.Tensor]:
+ x = self.patch_embed(x)
+ # x: (B, H, W, C)
+
+ # Add pos embed
+ x = x + self._get_pos_embed(x.shape[1:3])
+
+ outputs = []
+ for i, blk in enumerate(self.blocks):
+ x = blk(x)
+ if (i == self.stage_ends[-1]) or (
+ i in self.stage_ends and self.return_interm_layers
+ ):
+ feats = x.permute(0, 3, 1, 2)
+ outputs.append(feats)
+
+ return outputs
diff --git a/py/evf_sam/model/segment_anything_2/sam2/modeling/backbones/image_encoder.py b/py/evf_sam/model/segment_anything_2/sam2/modeling/backbones/image_encoder.py
new file mode 100644
index 0000000..5f92baf
--- /dev/null
+++ b/py/evf_sam/model/segment_anything_2/sam2/modeling/backbones/image_encoder.py
@@ -0,0 +1,133 @@
+# Copyright (c) Meta Platforms, Inc. and affiliates.
+# All rights reserved.
+
+# This source code is licensed under the license found in the
+# LICENSE file in the root directory of this source tree.
+
+from typing import List, Optional
+
+import torch
+import torch.nn as nn
+import torch.nn.functional as F
+
+
+class ImageEncoder(nn.Module):
+ def __init__(
+ self,
+ trunk: nn.Module,
+ neck: nn.Module,
+ scalp: int = 0,
+ ):
+ super().__init__()
+ self.trunk = trunk
+ self.neck = neck
+ self.scalp = scalp
+ assert (
+ self.trunk.channel_list == self.neck.backbone_channel_list
+ ), f"Channel dims of trunk and neck do not match. Trunk: {self.trunk.channel_list}, neck: {self.neck.backbone_channel_list}"
+
+ def forward(self, sample: torch.Tensor):
+ # Forward through backbone
+ features, pos = self.neck(self.trunk(sample))
+ if self.scalp > 0:
+ # Discard the lowest resolution features
+ features, pos = features[: -self.scalp], pos[: -self.scalp]
+
+ src = features[-1]
+ output = {
+ "vision_features": src,
+ "vision_pos_enc": pos,
+ "backbone_fpn": features,
+ }
+ return output
+
+
+class FpnNeck(nn.Module):
+ """
+ A modified variant of Feature Pyramid Network (FPN) neck
+ (we remove output conv and also do bicubic interpolation similar to ViT
+ pos embed interpolation)
+ """
+
+ def __init__(
+ self,
+ position_encoding: nn.Module,
+ d_model: int,
+ backbone_channel_list: List[int],
+ kernel_size: int = 1,
+ stride: int = 1,
+ padding: int = 0,
+ fpn_interp_model: str = "bilinear",
+ fuse_type: str = "sum",
+ fpn_top_down_levels: Optional[List[int]] = None,
+ ):
+ """Initialize the neck
+ :param trunk: the backbone
+ :param position_encoding: the positional encoding to use
+ :param d_model: the dimension of the model
+ :param neck_norm: the normalization to use
+ """
+ super().__init__()
+ self.position_encoding = position_encoding
+ self.convs = nn.ModuleList()
+ self.backbone_channel_list = backbone_channel_list
+ for dim in backbone_channel_list:
+ current = nn.Sequential()
+ current.add_module(
+ "conv",
+ nn.Conv2d(
+ in_channels=dim,
+ out_channels=d_model,
+ kernel_size=kernel_size,
+ stride=stride,
+ padding=padding,
+ ),
+ )
+
+ self.convs.append(current)
+ self.fpn_interp_model = fpn_interp_model
+ assert fuse_type in ["sum", "avg"]
+ self.fuse_type = fuse_type
+
+ # levels to have top-down features in its outputs
+ # e.g. if fpn_top_down_levels is [2, 3], then only outputs of level 2 and 3
+ # have top-down propagation, while outputs of level 0 and level 1 have only
+ # lateral features from the same backbone level.
+ if fpn_top_down_levels is None:
+ # default is to have top-down features on all levels
+ fpn_top_down_levels = range(len(self.convs))
+ self.fpn_top_down_levels = list(fpn_top_down_levels)
+
+ def forward(self, xs: List[torch.Tensor]):
+
+ out = [None] * len(self.convs)
+ pos = [None] * len(self.convs)
+ assert len(xs) == len(self.convs)
+ # fpn forward pass
+ # see https://github.com/facebookresearch/detectron2/blob/main/detectron2/modeling/backbone/fpn.py
+ prev_features = None
+ # forward in top-down order (from low to high resolution)
+ n = len(self.convs) - 1
+ for i in range(n, -1, -1):
+ x = xs[i]
+ lateral_features = self.convs[n - i](x)
+ if i in self.fpn_top_down_levels and prev_features is not None:
+ top_down_features = F.interpolate(
+ prev_features.to(dtype=torch.float32),
+ scale_factor=2.0,
+ mode=self.fpn_interp_model,
+ align_corners=(
+ None if self.fpn_interp_model == "nearest" else False
+ ),
+ antialias=False,
+ )
+ prev_features = lateral_features + top_down_features
+ if self.fuse_type == "avg":
+ prev_features /= 2
+ else:
+ prev_features = lateral_features
+ x_out = prev_features
+ out[i] = x_out
+ pos[i] = self.position_encoding(x_out).to(x_out.dtype)
+
+ return out, pos
diff --git a/py/evf_sam/model/segment_anything_2/sam2/modeling/backbones/utils.py b/py/evf_sam/model/segment_anything_2/sam2/modeling/backbones/utils.py
new file mode 100644
index 0000000..32d55c7
--- /dev/null
+++ b/py/evf_sam/model/segment_anything_2/sam2/modeling/backbones/utils.py
@@ -0,0 +1,95 @@
+# Copyright (c) Meta Platforms, Inc. and affiliates.
+# All rights reserved.
+
+# This source code is licensed under the license found in the
+# LICENSE file in the root directory of this source tree.
+
+"""Some utilities for backbones, in particular for windowing"""
+
+from typing import Tuple
+
+import torch
+import torch.nn as nn
+import torch.nn.functional as F
+
+
+def window_partition(x, window_size):
+ """
+ Partition into non-overlapping windows with padding if needed.
+ Args:
+ x (tensor): input tokens with [B, H, W, C].
+ window_size (int): window size.
+ Returns:
+ windows: windows after partition with [B * num_windows, window_size, window_size, C].
+ (Hp, Wp): padded height and width before partition
+ """
+ B, H, W, C = x.shape
+
+ pad_h = (window_size - H % window_size) % window_size
+ pad_w = (window_size - W % window_size) % window_size
+ if pad_h > 0 or pad_w > 0:
+ x = F.pad(x, (0, 0, 0, pad_w, 0, pad_h))
+ Hp, Wp = H + pad_h, W + pad_w
+
+ x = x.view(B, Hp // window_size, window_size, Wp // window_size, window_size, C)
+ windows = (
+ x.permute(0, 1, 3, 2, 4, 5).contiguous().view(-1, window_size, window_size, C)
+ )
+ return windows, (Hp, Wp)
+
+
+def window_unpartition(windows, window_size, pad_hw, hw):
+ """
+ Window unpartition into original sequences and removing padding.
+ Args:
+ x (tensor): input tokens with [B * num_windows, window_size, window_size, C].
+ window_size (int): window size.
+ pad_hw (Tuple): padded height and width (Hp, Wp).
+ hw (Tuple): original height and width (H, W) before padding.
+ Returns:
+ x: unpartitioned sequences with [B, H, W, C].
+ """
+ Hp, Wp = pad_hw
+ H, W = hw
+ B = windows.shape[0] // (Hp * Wp // window_size // window_size)
+ x = windows.view(
+ B, Hp // window_size, Wp // window_size, window_size, window_size, -1
+ )
+ x = x.permute(0, 1, 3, 2, 4, 5).contiguous().view(B, Hp, Wp, -1)
+
+ if Hp > H or Wp > W:
+ x = x[:, :H, :W, :].contiguous()
+ return x
+
+
+class PatchEmbed(nn.Module):
+ """
+ Image to Patch Embedding.
+ """
+
+ def __init__(
+ self,
+ kernel_size: Tuple[int, ...] = (7, 7),
+ stride: Tuple[int, ...] = (4, 4),
+ padding: Tuple[int, ...] = (3, 3),
+ in_chans: int = 3,
+ embed_dim: int = 768,
+ ):
+ """
+ Args:
+ kernel_size (Tuple): kernel size of the projection layer.
+ stride (Tuple): stride of the projection layer.
+ padding (Tuple): padding size of the projection layer.
+ in_chans (int): Number of input image channels.
+ embed_dim (int): embed_dim (int): Patch embedding dimension.
+ """
+ super().__init__()
+ self.proj = nn.Conv2d(
+ in_chans, embed_dim, kernel_size=kernel_size, stride=stride, padding=padding
+ )
+
+ def forward(self, x: torch.Tensor) -> torch.Tensor:
+ x = self.proj(x)
+ # B C H W -> B H W C
+ x = x.permute(0, 2, 3, 1)
+ return x
diff --git a/py/evf_sam/model/segment_anything_2/sam2/modeling/memory_attention.py b/py/evf_sam/model/segment_anything_2/sam2/modeling/memory_attention.py
new file mode 100644
index 0000000..42f56ef
--- /dev/null
+++ b/py/evf_sam/model/segment_anything_2/sam2/modeling/memory_attention.py
@@ -0,0 +1,169 @@
+# Copyright (c) Meta Platforms, Inc. and affiliates.
+# All rights reserved.
+
+# This source code is licensed under the license found in the
+# LICENSE file in the root directory of this source tree.
+
+from typing import Optional
+
+import torch
+from torch import nn, Tensor
+
+from model.segment_anything_2.sam2.modeling.sam.transformer import RoPEAttention
+
+from model.segment_anything_2.sam2.modeling.sam2_utils import get_activation_fn, get_clones
+
+
+class MemoryAttentionLayer(nn.Module):
+
+ def __init__(
+ self,
+ activation: str,
+ cross_attention: nn.Module,
+ d_model: int,
+ dim_feedforward: int,
+ dropout: float,
+ pos_enc_at_attn: bool,
+ pos_enc_at_cross_attn_keys: bool,
+ pos_enc_at_cross_attn_queries: bool,
+ self_attention: nn.Module,
+ ):
+ super().__init__()
+ self.d_model = d_model
+ self.dim_feedforward = dim_feedforward
+ self.dropout_value = dropout
+ self.self_attn = self_attention
+ self.cross_attn_image = cross_attention
+
+ # Implementation of Feedforward model
+ self.linear1 = nn.Linear(d_model, dim_feedforward)
+ self.dropout = nn.Dropout(dropout)
+ self.linear2 = nn.Linear(dim_feedforward, d_model)
+
+ self.norm1 = nn.LayerNorm(d_model)
+ self.norm2 = nn.LayerNorm(d_model)
+ self.norm3 = nn.LayerNorm(d_model)
+ self.dropout1 = nn.Dropout(dropout)
+ self.dropout2 = nn.Dropout(dropout)
+ self.dropout3 = nn.Dropout(dropout)
+
+ self.activation_str = activation
+ self.activation = get_activation_fn(activation)
+
+ # Where to add pos enc
+ self.pos_enc_at_attn = pos_enc_at_attn
+ self.pos_enc_at_cross_attn_queries = pos_enc_at_cross_attn_queries
+ self.pos_enc_at_cross_attn_keys = pos_enc_at_cross_attn_keys
+
+ def _forward_sa(self, tgt, query_pos):
+ # Self-Attention
+ tgt2 = self.norm1(tgt)
+ q = k = tgt2 + query_pos if self.pos_enc_at_attn else tgt2
+ tgt2 = self.self_attn(q, k, v=tgt2)
+ tgt = tgt + self.dropout1(tgt2)
+ return tgt
+
+ def _forward_ca(self, tgt, memory, query_pos, pos, num_k_exclude_rope=0):
+ kwds = {}
+ if num_k_exclude_rope > 0:
+ assert isinstance(self.cross_attn_image, RoPEAttention)
+ kwds = {"num_k_exclude_rope": num_k_exclude_rope}
+
+ # Cross-Attention
+ tgt2 = self.norm2(tgt)
+ tgt2 = self.cross_attn_image(
+ q=tgt2 + query_pos if self.pos_enc_at_cross_attn_queries else tgt2,
+ k=memory + pos if self.pos_enc_at_cross_attn_keys else memory,
+ v=memory,
+ **kwds,
+ )
+ tgt = tgt + self.dropout2(tgt2)
+ return tgt
+
+ def forward(
+ self,
+ tgt,
+ memory,
+ pos: Optional[Tensor] = None,
+ query_pos: Optional[Tensor] = None,
+ num_k_exclude_rope: int = 0,
+ ) -> torch.Tensor:
+
+ # Self-Attn, Cross-Attn
+ tgt = self._forward_sa(tgt, query_pos)
+ tgt = self._forward_ca(tgt, memory, query_pos, pos, num_k_exclude_rope)
+ # MLP
+ tgt2 = self.norm3(tgt)
+ tgt2 = self.linear2(self.dropout(self.activation(self.linear1(tgt2))))
+ tgt = tgt + self.dropout3(tgt2)
+ return tgt
+
+
+class MemoryAttention(nn.Module):
+ def __init__(
+ self,
+ d_model: int,
+ pos_enc_at_input: bool,
+ layer: nn.Module,
+ num_layers: int,
+ batch_first: bool = True, # Do layers expect batch first input?
+ ):
+ super().__init__()
+ self.d_model = d_model
+ self.layers = get_clones(layer, num_layers)
+ self.num_layers = num_layers
+ self.norm = nn.LayerNorm(d_model)
+ self.pos_enc_at_input = pos_enc_at_input
+ self.batch_first = batch_first
+
+ def forward(
+ self,
+ curr: torch.Tensor, # self-attention inputs
+ memory: torch.Tensor, # cross-attention inputs
+ curr_pos: Optional[Tensor] = None, # pos_enc for self-attention inputs
+ memory_pos: Optional[Tensor] = None, # pos_enc for cross-attention inputs
+ num_obj_ptr_tokens: int = 0, # number of object pointer *tokens*
+ ):
+ if isinstance(curr, list):
+ assert isinstance(curr_pos, list)
+ assert len(curr) == len(curr_pos) == 1
+ curr, curr_pos = (
+ curr[0],
+ curr_pos[0],
+ )
+
+ assert (
+ curr.shape[1] == memory.shape[1]
+ ), "Batch size must be the same for curr and memory"
+
+ output = curr
+ if self.pos_enc_at_input and curr_pos is not None:
+ output = output + 0.1 * curr_pos
+
+ if self.batch_first:
+ # Convert to batch first
+ output = output.transpose(0, 1)
+ curr_pos = curr_pos.transpose(0, 1)
+ memory = memory.transpose(0, 1)
+ memory_pos = memory_pos.transpose(0, 1)
+
+ for layer in self.layers:
+ kwds = {}
+ if isinstance(layer.cross_attn_image, RoPEAttention):
+ kwds = {"num_k_exclude_rope": num_obj_ptr_tokens}
+
+ output = layer(
+ tgt=output,
+ memory=memory,
+ pos=memory_pos,
+ query_pos=curr_pos,
+ **kwds,
+ )
+ normed_output = self.norm(output)
+
+ if self.batch_first:
+ # Convert back to seq first
+ normed_output = normed_output.transpose(0, 1)
+ curr_pos = curr_pos.transpose(0, 1)
+
+ return normed_output
diff --git a/py/evf_sam/model/segment_anything_2/sam2/modeling/memory_encoder.py b/py/evf_sam/model/segment_anything_2/sam2/modeling/memory_encoder.py
new file mode 100644
index 0000000..fb11cbf
--- /dev/null
+++ b/py/evf_sam/model/segment_anything_2/sam2/modeling/memory_encoder.py
@@ -0,0 +1,182 @@
+# Copyright (c) Meta Platforms, Inc. and affiliates.
+# All rights reserved.
+
+# This source code is licensed under the license found in the
+# LICENSE file in the root directory of this source tree.
+
+import math
+from typing import Tuple
+
+import torch
+import torch.nn as nn
+import torch.nn.functional as F
+
+from model.segment_anything_2.sam2.modeling.sam2_utils import DropPath, get_clones, LayerNorm2d
+
+
+class MaskDownSampler(nn.Module):
+ """
+ Progressively downsample a mask by total_stride, each time by stride.
+ Note that LayerNorm is applied per *token*, like in ViT.
+
+ With each downsample (by a factor stride**2), channel capacity increases by the same factor.
+ In the end, we linearly project to embed_dim channels.
+ """
+
+ def __init__(
+ self,
+ embed_dim=256,
+ kernel_size=4,
+ stride=4,
+ padding=0,
+ total_stride=16,
+ activation=nn.GELU,
+ ):
+ super().__init__()
+ num_layers = int(math.log2(total_stride) // math.log2(stride))
+ assert stride**num_layers == total_stride
+ self.encoder = nn.Sequential()
+ mask_in_chans, mask_out_chans = 1, 1
+ for _ in range(num_layers):
+ mask_out_chans = mask_in_chans * (stride**2)
+ self.encoder.append(
+ nn.Conv2d(
+ mask_in_chans,
+ mask_out_chans,
+ kernel_size=kernel_size,
+ stride=stride,
+ padding=padding,
+ )
+ )
+ self.encoder.append(LayerNorm2d(mask_out_chans))
+ self.encoder.append(activation())
+ mask_in_chans = mask_out_chans
+
+ self.encoder.append(nn.Conv2d(mask_out_chans, embed_dim, kernel_size=1))
+
+ def forward(self, x):
+ return self.encoder(x)
+
+
+# Lightly adapted from ConvNext (https://github.com/facebookresearch/ConvNeXt)
+class CXBlock(nn.Module):
+ r"""ConvNeXt Block. There are two equivalent implementations:
+ (1) DwConv -> LayerNorm (channels_first) -> 1x1 Conv -> GELU -> 1x1 Conv; all in (N, C, H, W)
+ (2) DwConv -> Permute to (N, H, W, C); LayerNorm (channels_last) -> Linear -> GELU -> Linear; Permute back
+ We use (2) as we find it slightly faster in PyTorch
+
+ Args:
+ dim (int): Number of input channels.
+ drop_path (float): Stochastic depth rate. Default: 0.0
+ layer_scale_init_value (float): Init value for Layer Scale. Default: 1e-6.
+ """
+
+ def __init__(
+ self,
+ dim,
+ kernel_size=7,
+ padding=3,
+ drop_path=0.0,
+ layer_scale_init_value=1e-6,
+ use_dwconv=True,
+ ):
+ super().__init__()
+ self.dwconv = nn.Conv2d(
+ dim,
+ dim,
+ kernel_size=kernel_size,
+ padding=padding,
+ groups=dim if use_dwconv else 1,
+ ) # depthwise conv
+ self.norm = LayerNorm2d(dim, eps=1e-6)
+ self.pwconv1 = nn.Linear(
+ dim, 4 * dim
+ ) # pointwise/1x1 convs, implemented with linear layers
+ self.act = nn.GELU()
+ self.pwconv2 = nn.Linear(4 * dim, dim)
+ # modified by ZhangYx from self.gamma to self.weight. Due to (https://github.com/facebookresearch/segment-anything-2/issues/85)
+ self.weight = (
+ nn.Parameter(layer_scale_init_value * torch.ones((dim)), requires_grad=True)
+ if layer_scale_init_value > 0
+ else None
+ )
+ self.drop_path = DropPath(drop_path) if drop_path > 0.0 else nn.Identity()
+
+ def forward(self, x):
+ input = x
+ x = self.dwconv(x)
+ x = self.norm(x)
+ x = x.permute(0, 2, 3, 1) # (N, C, H, W) -> (N, H, W, C)
+ x = self.pwconv1(x)
+ x = self.act(x)
+ x = self.pwconv2(x)
+ if self.weight is not None:
+ x = self.weight * x
+ x = x.permute(0, 3, 1, 2) # (N, H, W, C) -> (N, C, H, W)
+
+ x = input + self.drop_path(x)
+ return x
+
+
+class Fuser(nn.Module):
+ def __init__(self, layer, num_layers, dim=None, input_projection=False):
+ super().__init__()
+ self.proj = nn.Identity()
+ self.layers = get_clones(layer, num_layers)
+
+ if input_projection:
+ assert dim is not None
+ self.proj = nn.Conv2d(dim, dim, kernel_size=1)
+
+ def forward(self, x):
+ # normally x: (N, C, H, W)
+ x = self.proj(x)
+ for layer in self.layers:
+ x = layer(x)
+ return x
+
+
+class MemoryEncoder(nn.Module):
+ def __init__(
+ self,
+ out_dim,
+ mask_downsampler,
+ fuser,
+ position_encoding,
+ in_dim=256, # in_dim of pix_feats
+ ):
+ super().__init__()
+
+ self.mask_downsampler = mask_downsampler
+
+ self.pix_feat_proj = nn.Conv2d(in_dim, in_dim, kernel_size=1)
+ self.fuser = fuser
+ self.position_encoding = position_encoding
+ self.out_proj = nn.Identity()
+ if out_dim != in_dim:
+ self.out_proj = nn.Conv2d(in_dim, out_dim, kernel_size=1)
+
+ def forward(
+ self,
+ pix_feat: torch.Tensor,
+ masks: torch.Tensor,
+ skip_mask_sigmoid: bool = False,
+ ) -> Tuple[torch.Tensor, torch.Tensor]:
+ ## Process masks
+ # sigmoid, so that less domain shift from gt masks which are bool
+ if not skip_mask_sigmoid:
+ masks = F.sigmoid(masks)
+ masks = self.mask_downsampler(masks)
+
+ ## Fuse pix_feats and downsampled masks
+ # in case the visual features are on CPU, cast them to CUDA
+ pix_feat = pix_feat.to(masks.device)
+
+ x = self.pix_feat_proj(pix_feat)
+ x = x + masks
+ x = self.fuser(x)
+ x = self.out_proj(x)
+
+ pos = self.position_encoding(x).to(x.dtype)
+
+ return {"vision_features": x, "vision_pos_enc": [pos]}
diff --git a/py/evf_sam/model/segment_anything_2/sam2/modeling/position_encoding.py b/py/evf_sam/model/segment_anything_2/sam2/modeling/position_encoding.py
new file mode 100644
index 0000000..f4b57ae
--- /dev/null
+++ b/py/evf_sam/model/segment_anything_2/sam2/modeling/position_encoding.py
@@ -0,0 +1,216 @@
+# Copyright (c) Meta Platforms, Inc. and affiliates.
+# All rights reserved.
+
+# This source code is licensed under the license found in the
+# LICENSE file in the root directory of this source tree.
+
+import math
+from typing import Any, Optional, Tuple
+
+import numpy as np
+
+import torch
+from torch import nn
+
+
+class PositionEmbeddingSine(nn.Module):
+ """
+ This is a more standard version of the position embedding, very similar to the one
+ used by the Attention is all you need paper, generalized to work on images.
+ """
+
+ def __init__(
+ self,
+ num_pos_feats,
+ temperature: int = 10000,
+ normalize: bool = True,
+ scale: Optional[float] = None,
+ ):
+ super().__init__()
+ assert num_pos_feats % 2 == 0, "Expecting even model width"
+ self.num_pos_feats = num_pos_feats // 2
+ self.temperature = temperature
+ self.normalize = normalize
+ if scale is not None and normalize is False:
+ raise ValueError("normalize should be True if scale is passed")
+ if scale is None:
+ scale = 2 * math.pi
+ self.scale = scale
+
+ self.cache = {}
+
+ def _encode_xy(self, x, y):
+ # The positions are expected to be normalized
+ assert len(x) == len(y) and x.ndim == y.ndim == 1
+ x_embed = x * self.scale
+ y_embed = y * self.scale
+
+ dim_t = torch.arange(self.num_pos_feats, dtype=torch.float32, device=x.device)
+ dim_t = self.temperature ** (2 * (dim_t // 2) / self.num_pos_feats)
+
+ pos_x = x_embed[:, None] / dim_t
+ pos_y = y_embed[:, None] / dim_t
+ pos_x = torch.stack(
+ (pos_x[:, 0::2].sin(), pos_x[:, 1::2].cos()), dim=2
+ ).flatten(1)
+ pos_y = torch.stack(
+ (pos_y[:, 0::2].sin(), pos_y[:, 1::2].cos()), dim=2
+ ).flatten(1)
+ return pos_x, pos_y
+
+ @torch.no_grad()
+ def encode_boxes(self, x, y, w, h):
+ pos_x, pos_y = self._encode_xy(x, y)
+ pos = torch.cat((pos_y, pos_x, h[:, None], w[:, None]), dim=1)
+ return pos
+
+ encode = encode_boxes # Backwards compatibility
+
+ @torch.no_grad()
+ def encode_points(self, x, y, labels):
+ (bx, nx), (by, ny), (bl, nl) = x.shape, y.shape, labels.shape
+ assert bx == by and nx == ny and bx == bl and nx == nl
+ pos_x, pos_y = self._encode_xy(x.flatten(), y.flatten())
+ pos_x, pos_y = pos_x.reshape(bx, nx, -1), pos_y.reshape(by, ny, -1)
+ pos = torch.cat((pos_y, pos_x, labels[:, :, None]), dim=2)
+ return pos
+
+ @torch.no_grad()
+ def forward(self, x: torch.Tensor):
+ cache_key = (x.shape[-2], x.shape[-1])
+ if cache_key in self.cache:
+ return self.cache[cache_key][None].repeat(x.shape[0], 1, 1, 1)
+ y_embed = (
+ torch.arange(1, x.shape[-2] + 1, dtype=torch.float32, device=x.device)
+ .view(1, -1, 1)
+ .repeat(x.shape[0], 1, x.shape[-1])
+ )
+ x_embed = (
+ torch.arange(1, x.shape[-1] + 1, dtype=torch.float32, device=x.device)
+ .view(1, 1, -1)
+ .repeat(x.shape[0], x.shape[-2], 1)
+ )
+
+ if self.normalize:
+ eps = 1e-6
+ y_embed = y_embed / (y_embed[:, -1:, :] + eps) * self.scale
+ x_embed = x_embed / (x_embed[:, :, -1:] + eps) * self.scale
+
+ dim_t = torch.arange(self.num_pos_feats, dtype=torch.float32, device=x.device)
+ dim_t = self.temperature ** (2 * (dim_t // 2) / self.num_pos_feats)
+
+ pos_x = x_embed[:, :, :, None] / dim_t
+ pos_y = y_embed[:, :, :, None] / dim_t
+ pos_x = torch.stack(
+ (pos_x[:, :, :, 0::2].sin(), pos_x[:, :, :, 1::2].cos()), dim=4
+ ).flatten(3)
+ pos_y = torch.stack(
+ (pos_y[:, :, :, 0::2].sin(), pos_y[:, :, :, 1::2].cos()), dim=4
+ ).flatten(3)
+ pos = torch.cat((pos_y, pos_x), dim=3).permute(0, 3, 1, 2)
+ self.cache[cache_key] = pos[0]
+ return pos
+
+
+class PositionEmbeddingRandom(nn.Module):
+ """
+ Positional encoding using random spatial frequencies.
+ """
+
+ def __init__(self, num_pos_feats: int = 64, scale: Optional[float] = None) -> None:
+ super().__init__()
+ if scale is None or scale <= 0.0:
+ scale = 1.0
+ self.register_buffer(
+ "positional_encoding_gaussian_matrix",
+ scale * torch.randn((2, num_pos_feats)),
+ )
+
+ def _pe_encoding(self, coords: torch.Tensor) -> torch.Tensor:
+ """Positionally encode points that are normalized to [0,1]."""
+ # assuming coords are in [0, 1]^2 square and have d_1 x ... x d_n x 2 shape
+ coords = 2 * coords - 1
+ coords = coords @ self.positional_encoding_gaussian_matrix
+ coords = 2 * np.pi * coords
+ # outputs d_1 x ... x d_n x C shape
+ return torch.cat([torch.sin(coords), torch.cos(coords)], dim=-1)
+
+ def forward(self, size: Tuple[int, int]) -> torch.Tensor:
+ """Generate positional encoding for a grid of the specified size."""
+ h, w = size
+ device: Any = self.positional_encoding_gaussian_matrix.device
+ grid = torch.ones((h, w), device=device, dtype=torch.float32)
+ y_embed = grid.cumsum(dim=0) - 0.5
+ x_embed = grid.cumsum(dim=1) - 0.5
+ y_embed = y_embed / h
+ x_embed = x_embed / w
+
+ pe = self._pe_encoding(torch.stack([x_embed, y_embed], dim=-1))
+ return pe.permute(2, 0, 1) # C x H x W
+
+ def forward_with_coords(
+ self, coords_input: torch.Tensor, image_size: Tuple[int, int]
+ ) -> torch.Tensor:
+ """Positionally encode points that are not normalized to [0,1]."""
+ coords = coords_input.clone()
+ coords[:, :, 0] = coords[:, :, 0] / image_size[1]
+ coords[:, :, 1] = coords[:, :, 1] / image_size[0]
+ return self._pe_encoding(coords.to(torch.float)) # B x N x C
+
+
+# Rotary Positional Encoding, adapted from:
+# 1. https://github.com/meta-llama/codellama/blob/main/llama/model.py
+# 2. https://github.com/naver-ai/rope-vit
+# 3. https://github.com/lucidrains/rotary-embedding-torch
+
+
+def init_t_xy(end_x: int, end_y: int):
+ t = torch.arange(end_x * end_y, dtype=torch.float32)
+ t_x = (t % end_x).float()
+ t_y = torch.div(t, end_x, rounding_mode="floor").float()
+ return t_x, t_y
+
+
+def compute_axial_cis(dim: int, end_x: int, end_y: int, theta: float = 10000.0):
+ freqs_x = 1.0 / (theta ** (torch.arange(0, dim, 4)[: (dim // 4)].float() / dim))
+ freqs_y = 1.0 / (theta ** (torch.arange(0, dim, 4)[: (dim // 4)].float() / dim))
+
+ t_x, t_y = init_t_xy(end_x, end_y)
+ freqs_x = torch.outer(t_x, freqs_x)
+ freqs_y = torch.outer(t_y, freqs_y)
+ freqs_cis_x = torch.polar(torch.ones_like(freqs_x), freqs_x)
+ freqs_cis_y = torch.polar(torch.ones_like(freqs_y), freqs_y)
+ return torch.cat([freqs_cis_x, freqs_cis_y], dim=-1)
+
+
+def reshape_for_broadcast(freqs_cis: torch.Tensor, x: torch.Tensor):
+ ndim = x.ndim
+ assert 0 <= 1 < ndim
+ assert freqs_cis.shape == (x.shape[-2], x.shape[-1])
+ shape = [d if i >= ndim - 2 else 1 for i, d in enumerate(x.shape)]
+ return freqs_cis.view(*shape)
+
+
+def apply_rotary_enc(
+ xq: torch.Tensor,
+ xk: torch.Tensor,
+ freqs_cis: torch.Tensor,
+ repeat_freqs_k: bool = False,
+):
+ xq_ = torch.view_as_complex(xq.float().reshape(*xq.shape[:-1], -1, 2))
+ xk_ = (
+ torch.view_as_complex(xk.float().reshape(*xk.shape[:-1], -1, 2))
+ if xk.shape[-2] != 0
+ else None
+ )
+ freqs_cis = reshape_for_broadcast(freqs_cis, xq_)
+ xq_out = torch.view_as_real(xq_ * freqs_cis).flatten(3)
+ if xk_ is None:
+ # no keys to rotate, due to dropout
+ return xq_out.type_as(xq).to(xq.device), xk
+ # repeat freqs along seq_len dim to match k seq_len
+ if repeat_freqs_k:
+ r = xk_.shape[-2] // xq_.shape[-2]
+ freqs_cis = freqs_cis.repeat(*([1] * (freqs_cis.ndim - 2)), r, 1)
+ xk_out = torch.view_as_real(xk_ * freqs_cis).flatten(3)
+ return xq_out.type_as(xq).to(xq.device), xk_out.type_as(xk).to(xk.device)
diff --git a/py/evf_sam/model/segment_anything_2/sam2/modeling/sam/__init__.py b/py/evf_sam/model/segment_anything_2/sam2/modeling/sam/__init__.py
new file mode 100644
index 0000000..5277f46
--- /dev/null
+++ b/py/evf_sam/model/segment_anything_2/sam2/modeling/sam/__init__.py
@@ -0,0 +1,5 @@
+# Copyright (c) Meta Platforms, Inc. and affiliates.
+# All rights reserved.
+
+# This source code is licensed under the license found in the
+# LICENSE file in the root directory of this source tree.
diff --git a/py/evf_sam/model/segment_anything_2/sam2/modeling/sam/mask_decoder.py b/py/evf_sam/model/segment_anything_2/sam2/modeling/sam/mask_decoder.py
new file mode 100644
index 0000000..1992910
--- /dev/null
+++ b/py/evf_sam/model/segment_anything_2/sam2/modeling/sam/mask_decoder.py
@@ -0,0 +1,295 @@
+# Copyright (c) Meta Platforms, Inc. and affiliates.
+# All rights reserved.
+
+# This source code is licensed under the license found in the
+# LICENSE file in the root directory of this source tree.
+
+from typing import List, Optional, Tuple, Type
+
+import torch
+from torch import nn
+
+from model.segment_anything_2.sam2.modeling.sam2_utils import LayerNorm2d, MLP
+
+
+class MaskDecoder(nn.Module):
+ def __init__(
+ self,
+ *,
+ transformer_dim: int,
+ transformer: nn.Module,
+ num_multimask_outputs: int = 3,
+ activation: Type[nn.Module] = nn.GELU,
+ iou_head_depth: int = 3,
+ iou_head_hidden_dim: int = 256,
+ use_high_res_features: bool = False,
+ iou_prediction_use_sigmoid=False,
+ dynamic_multimask_via_stability=False,
+ dynamic_multimask_stability_delta=0.05,
+ dynamic_multimask_stability_thresh=0.98,
+ pred_obj_scores: bool = False,
+ pred_obj_scores_mlp: bool = False,
+ use_multimask_token_for_obj_ptr: bool = False,
+ ) -> None:
+ """
+ Predicts masks given an image and prompt embeddings, using a
+ transformer architecture.
+
+ Arguments:
+ transformer_dim (int): the channel dimension of the transformer
+ transformer (nn.Module): the transformer used to predict masks
+ num_multimask_outputs (int): the number of masks to predict
+ when disambiguating masks
+ activation (nn.Module): the type of activation to use when
+ upscaling masks
+ iou_head_depth (int): the depth of the MLP used to predict
+ mask quality
+ iou_head_hidden_dim (int): the hidden dimension of the MLP
+ used to predict mask quality
+ """
+ super().__init__()
+ self.transformer_dim = transformer_dim
+ self.transformer = transformer
+
+ self.num_multimask_outputs = num_multimask_outputs
+
+ self.iou_token = nn.Embedding(1, transformer_dim)
+ self.num_mask_tokens = num_multimask_outputs + 1
+ self.mask_tokens = nn.Embedding(self.num_mask_tokens, transformer_dim)
+
+ self.pred_obj_scores = pred_obj_scores
+ if self.pred_obj_scores:
+ self.obj_score_token = nn.Embedding(1, transformer_dim)
+ self.use_multimask_token_for_obj_ptr = use_multimask_token_for_obj_ptr
+
+ self.output_upscaling = nn.Sequential(
+ nn.ConvTranspose2d(
+ transformer_dim, transformer_dim // 4, kernel_size=2, stride=2
+ ),
+ LayerNorm2d(transformer_dim // 4),
+ activation(),
+ nn.ConvTranspose2d(
+ transformer_dim // 4, transformer_dim // 8, kernel_size=2, stride=2
+ ),
+ activation(),
+ )
+ self.use_high_res_features = use_high_res_features
+ if use_high_res_features:
+ self.conv_s0 = nn.Conv2d(
+ transformer_dim, transformer_dim // 8, kernel_size=1, stride=1
+ )
+ self.conv_s1 = nn.Conv2d(
+ transformer_dim, transformer_dim // 4, kernel_size=1, stride=1
+ )
+
+ self.output_hypernetworks_mlps = nn.ModuleList(
+ [
+ MLP(transformer_dim, transformer_dim, transformer_dim // 8, 3)
+ for i in range(self.num_mask_tokens)
+ ]
+ )
+
+ self.iou_prediction_head = MLP(
+ transformer_dim,
+ iou_head_hidden_dim,
+ self.num_mask_tokens,
+ iou_head_depth,
+ sigmoid_output=iou_prediction_use_sigmoid,
+ )
+ if self.pred_obj_scores:
+ self.pred_obj_score_head = nn.Linear(transformer_dim, 1)
+ if pred_obj_scores_mlp:
+ self.pred_obj_score_head = MLP(transformer_dim, transformer_dim, 1, 3)
+
+ # When outputting a single mask, optionally we can dynamically fall back to the best
+ # multimask output token if the single mask output token gives low stability scores.
+ self.dynamic_multimask_via_stability = dynamic_multimask_via_stability
+ self.dynamic_multimask_stability_delta = dynamic_multimask_stability_delta
+ self.dynamic_multimask_stability_thresh = dynamic_multimask_stability_thresh
+
+ def forward(
+ self,
+ image_embeddings: torch.Tensor,
+ image_pe: torch.Tensor,
+ sparse_prompt_embeddings: torch.Tensor,
+ dense_prompt_embeddings: torch.Tensor,
+ multimask_output: bool,
+ repeat_image: bool,
+ high_res_features: Optional[List[torch.Tensor]] = None,
+ ) -> Tuple[torch.Tensor, torch.Tensor]:
+ """
+ Predict masks given image and prompt embeddings.
+
+ Arguments:
+ image_embeddings (torch.Tensor): the embeddings from the image encoder
+ image_pe (torch.Tensor): positional encoding with the shape of image_embeddings
+ sparse_prompt_embeddings (torch.Tensor): the embeddings of the points and boxes
+ dense_prompt_embeddings (torch.Tensor): the embeddings of the mask inputs
+ multimask_output (bool): Whether to return multiple masks or a single
+ mask.
+
+ Returns:
+ torch.Tensor: batched predicted masks
+ torch.Tensor: batched predictions of mask quality
+ torch.Tensor: batched SAM token for mask output
+ """
+ masks, iou_pred, mask_tokens_out, object_score_logits = self.predict_masks(
+ image_embeddings=image_embeddings,
+ image_pe=image_pe,
+ sparse_prompt_embeddings=sparse_prompt_embeddings,
+ dense_prompt_embeddings=dense_prompt_embeddings,
+ repeat_image=repeat_image,
+ high_res_features=high_res_features,
+ )
+
+ # Select the correct mask or masks for output
+ if multimask_output:
+ masks = masks[:, 1:, :, :]
+ iou_pred = iou_pred[:, 1:]
+ elif self.dynamic_multimask_via_stability and not self.training:
+ masks, iou_pred = self._dynamic_multimask_via_stability(masks, iou_pred)
+ else:
+ masks = masks[:, 0:1, :, :]
+ iou_pred = iou_pred[:, 0:1]
+
+ if multimask_output and self.use_multimask_token_for_obj_ptr:
+ sam_tokens_out = mask_tokens_out[:, 1:] # [b, 3, c] shape
+ else:
+ # Take the mask output token. Here we *always* use the token for single mask output.
+ # At test time, even if we track after 1-click (and using multimask_output=True),
+ # we still take the single mask token here. The rationale is that we always track
+ # after multiple clicks during training, so the past tokens seen during training
+ # are always the single mask token (and we'll let it be the object-memory token).
+ sam_tokens_out = mask_tokens_out[:, 0:1] # [b, 1, c] shape
+
+ # Prepare output
+ return masks, iou_pred, sam_tokens_out, object_score_logits
+
+ def predict_masks(
+ self,
+ image_embeddings: torch.Tensor,
+ image_pe: torch.Tensor,
+ sparse_prompt_embeddings: torch.Tensor,
+ dense_prompt_embeddings: torch.Tensor,
+ repeat_image: bool,
+ high_res_features: Optional[List[torch.Tensor]] = None,
+ ) -> Tuple[torch.Tensor, torch.Tensor]:
+ """Predicts masks. See 'forward' for more details."""
+ # Concatenate output tokens
+ s = 0
+ if self.pred_obj_scores:
+ output_tokens = torch.cat(
+ [
+ self.obj_score_token.weight,
+ self.iou_token.weight,
+ self.mask_tokens.weight,
+ ],
+ dim=0,
+ )
+ s = 1
+ else:
+ output_tokens = torch.cat(
+ [self.iou_token.weight, self.mask_tokens.weight], dim=0
+ )
+ output_tokens = output_tokens.unsqueeze(0).expand(
+ sparse_prompt_embeddings.size(0), -1, -1
+ )
+ tokens = torch.cat((output_tokens, sparse_prompt_embeddings), dim=1)
+
+ # Expand per-image data in batch direction to be per-mask
+ if repeat_image:
+ src = torch.repeat_interleave(image_embeddings, tokens.shape[0], dim=0)
+ else:
+ assert image_embeddings.shape[0] == tokens.shape[0]
+ src = image_embeddings
+ src = src + dense_prompt_embeddings
+ assert (
+ image_pe.size(0) == 1
+ ), "image_pe should have size 1 in batch dim (from `get_dense_pe()`)"
+ pos_src = torch.repeat_interleave(image_pe, tokens.shape[0], dim=0)
+ b, c, h, w = src.shape
+
+ # Run the transformer
+ hs, src = self.transformer(src, pos_src, tokens)
+ iou_token_out = hs[:, s, :]
+ mask_tokens_out = hs[:, s + 1 : (s + 1 + self.num_mask_tokens), :]
+
+ # Upscale mask embeddings and predict masks using the mask tokens
+ src = src.transpose(1, 2).view(b, c, h, w)
+ if not self.use_high_res_features:
+ upscaled_embedding = self.output_upscaling(src)
+ else:
+ dc1, ln1, act1, dc2, act2 = self.output_upscaling
+ feat_s0, feat_s1 = high_res_features
+ upscaled_embedding = act1(ln1(dc1(src) + feat_s1))
+ upscaled_embedding = act2(dc2(upscaled_embedding) + feat_s0)
+
+ hyper_in_list: List[torch.Tensor] = []
+ for i in range(self.num_mask_tokens):
+ hyper_in_list.append(
+ self.output_hypernetworks_mlps[i](mask_tokens_out[:, i, :])
+ )
+ hyper_in = torch.stack(hyper_in_list, dim=1)
+ b, c, h, w = upscaled_embedding.shape
+ masks = (hyper_in @ upscaled_embedding.view(b, c, h * w)).view(b, -1, h, w)
+
+ # Generate mask quality predictions
+ iou_pred = self.iou_prediction_head(iou_token_out)
+ if self.pred_obj_scores:
+ assert s == 1
+ object_score_logits = self.pred_obj_score_head(hs[:, 0, :])
+ else:
+ # Obj scores logits - default to 10.0, i.e. assuming the object is present, sigmoid(10)=1
+ object_score_logits = 10.0 * iou_pred.new_ones(iou_pred.shape[0], 1)
+
+ return masks, iou_pred, mask_tokens_out, object_score_logits
+
+ def _get_stability_scores(self, mask_logits):
+ """
+ Compute stability scores of the mask logits based on the IoU between upper and
+ lower thresholds, similar to https://github.com/fairinternal/onevision/pull/568.
+ """
+ mask_logits = mask_logits.flatten(-2)
+ stability_delta = self.dynamic_multimask_stability_delta
+ area_i = torch.sum(mask_logits > stability_delta, dim=-1).float()
+ area_u = torch.sum(mask_logits > -stability_delta, dim=-1).float()
+ stability_scores = torch.where(area_u > 0, area_i / area_u, 1.0)
+ return stability_scores
+
+ def _dynamic_multimask_via_stability(self, all_mask_logits, all_iou_scores):
+ """
+ When outputting a single mask, if the stability score from the current single-mask
+ output (based on output token 0) falls below a threshold, we instead select from
+ multi-mask outputs (based on output token 1~3) the mask with the highest predicted
+ IoU score. This is intended to ensure a valid mask for both clicking and tracking.
+ """
+ # The best mask from multimask output tokens (1~3)
+ multimask_logits = all_mask_logits[:, 1:, :, :]
+ multimask_iou_scores = all_iou_scores[:, 1:]
+ best_scores_inds = torch.argmax(multimask_iou_scores, dim=-1)
+ batch_inds = torch.arange(
+ multimask_iou_scores.size(0), device=all_iou_scores.device
+ )
+ best_multimask_logits = multimask_logits[batch_inds, best_scores_inds]
+ best_multimask_logits = best_multimask_logits.unsqueeze(1)
+ best_multimask_iou_scores = multimask_iou_scores[batch_inds, best_scores_inds]
+ best_multimask_iou_scores = best_multimask_iou_scores.unsqueeze(1)
+
+ # The mask from singlemask output token 0 and its stability score
+ singlemask_logits = all_mask_logits[:, 0:1, :, :]
+ singlemask_iou_scores = all_iou_scores[:, 0:1]
+ stability_scores = self._get_stability_scores(singlemask_logits)
+ is_stable = stability_scores >= self.dynamic_multimask_stability_thresh
+
+ # Dynamically fall back to best multimask output upon low stability scores.
+ mask_logits_out = torch.where(
+ is_stable[..., None, None].expand_as(singlemask_logits),
+ singlemask_logits,
+ best_multimask_logits,
+ )
+ iou_scores_out = torch.where(
+ is_stable.expand_as(singlemask_iou_scores),
+ singlemask_iou_scores,
+ best_multimask_iou_scores,
+ )
+ return mask_logits_out, iou_scores_out
diff --git a/py/evf_sam/model/segment_anything_2/sam2/modeling/sam/prompt_encoder.py b/py/evf_sam/model/segment_anything_2/sam2/modeling/sam/prompt_encoder.py
new file mode 100644
index 0000000..44fb6c5
--- /dev/null
+++ b/py/evf_sam/model/segment_anything_2/sam2/modeling/sam/prompt_encoder.py
@@ -0,0 +1,241 @@
+# Copyright (c) Meta Platforms, Inc. and affiliates.
+# All rights reserved.
+
+# This source code is licensed under the license found in the
+# LICENSE file in the root directory of this source tree.
+
+from typing import Any, Optional, Tuple, Type
+
+import numpy as np
+import torch
+from torch import nn
+# from model.segment_anything_2.sam2.modeling.position_encoding import PositionEmbeddingRandom
+
+from model.segment_anything_2.sam2.modeling.sam2_utils import LayerNorm2d
+
+
+class PromptEncoder(nn.Module):
+ def __init__(
+ self,
+ embed_dim: int,
+ image_embedding_size: Tuple[int, int],
+ input_image_size: Tuple[int, int],
+ mask_in_chans: int,
+ activation: Type[nn.Module] = nn.GELU,
+ ) -> None:
+ """
+ Encodes prompts for input to SAM's mask decoder.
+
+ Arguments:
+ embed_dim (int): The prompts' embedding dimension
+ image_embedding_size (tuple(int, int)): The spatial size of the
+ image embedding, as (H, W).
+ input_image_size (int): The padded size of the image as input
+ to the image encoder, as (H, W).
+ mask_in_chans (int): The number of hidden channels used for
+ encoding input masks.
+ activation (nn.Module): The activation to use when encoding
+ input masks.
+ """
+ super().__init__()
+ self.embed_dim = embed_dim
+ self.input_image_size = input_image_size
+ self.image_embedding_size = image_embedding_size
+ self.pe_layer = PositionEmbeddingRandom(embed_dim // 2)
+
+ self.num_point_embeddings: int = 4 # pos/neg point + 2 box corners
+ point_embeddings = [
+ nn.Embedding(1, embed_dim) for i in range(self.num_point_embeddings)
+ ]
+ self.point_embeddings = nn.ModuleList(point_embeddings)
+ self.not_a_point_embed = nn.Embedding(1, embed_dim)
+
+ self.mask_input_size = (
+ 4 * image_embedding_size[0],
+ 4 * image_embedding_size[1],
+ )
+ self.mask_downscaling = nn.Sequential(
+ nn.Conv2d(1, mask_in_chans // 4, kernel_size=2, stride=2),
+ LayerNorm2d(mask_in_chans // 4),
+ activation(),
+ nn.Conv2d(mask_in_chans // 4, mask_in_chans, kernel_size=2, stride=2),
+ LayerNorm2d(mask_in_chans),
+ activation(),
+ nn.Conv2d(mask_in_chans, embed_dim, kernel_size=1),
+ )
+ self.no_mask_embed = nn.Embedding(1, embed_dim)
+
+ def get_dense_pe(self) -> torch.Tensor:
+ """
+ Returns the positional encoding used to encode point prompts,
+ applied to a dense set of points the shape of the image encoding.
+
+ Returns:
+ torch.Tensor: Positional encoding with shape
+ 1x(embed_dim)x(embedding_h)x(embedding_w)
+ """
+ return self.pe_layer(self.image_embedding_size).unsqueeze(0)
+
+ def _embed_points(
+ self,
+ points: torch.Tensor,
+ labels: torch.Tensor,
+ pad: bool,
+ ) -> torch.Tensor:
+ """Embeds point prompts."""
+ points = points + 0.5 # Shift to center of pixel
+ if pad:
+ padding_point = torch.zeros((points.shape[0], 1, 2), device=points.device)
+ padding_label = -torch.ones((labels.shape[0], 1), device=labels.device)
+ points = torch.cat([points, padding_point], dim=1)
+ labels = torch.cat([labels, padding_label], dim=1)
+ point_embedding = self.pe_layer.forward_with_coords(
+ points, self.input_image_size
+ )
+ point_embedding[labels == -1] = 0.0
+ point_embedding[labels == -1] += self.not_a_point_embed.weight
+ point_embedding[labels == 0] += self.point_embeddings[0].weight
+ point_embedding[labels == 1] += self.point_embeddings[1].weight
+ point_embedding[labels == 2] += self.point_embeddings[2].weight
+ point_embedding[labels == 3] += self.point_embeddings[3].weight
+ return point_embedding
+
+ def _embed_boxes(self, boxes: torch.Tensor) -> torch.Tensor:
+ """Embeds box prompts."""
+ boxes = boxes + 0.5 # Shift to center of pixel
+ coords = boxes.reshape(-1, 2, 2)
+ corner_embedding = self.pe_layer.forward_with_coords(
+ coords, self.input_image_size
+ )
+ corner_embedding[:, 0, :] += self.point_embeddings[2].weight
+ corner_embedding[:, 1, :] += self.point_embeddings[3].weight
+ return corner_embedding
+
+ def _embed_masks(self, masks: torch.Tensor) -> torch.Tensor:
+ """Embeds mask inputs."""
+ mask_embedding = self.mask_downscaling(masks)
+ return mask_embedding
+
+ def _get_batch_size(
+ self,
+ points: Optional[Tuple[torch.Tensor, torch.Tensor]],
+ boxes: Optional[torch.Tensor],
+ masks: Optional[torch.Tensor],
+ text_embeds: Optional[torch.Tensor],
+ ) -> int:
+ """
+ Gets the batch size of the output given the batch size of the input prompts.
+ """
+ if points is not None:
+ return points[0].shape[0]
+ elif boxes is not None:
+ return boxes.shape[0]
+ elif masks is not None:
+ return masks.shape[0]
+ elif text_embeds is not None:
+ return text_embeds.shape[0]
+ else:
+ return 1
+
+ def _get_device(self) -> torch.device:
+ return self.point_embeddings[0].weight.device
+
+ def forward(
+ self,
+ points: Optional[Tuple[torch.Tensor, torch.Tensor]],
+ boxes: Optional[torch.Tensor],
+ masks: Optional[torch.Tensor],
+ text_embeds: Optional[torch.Tensor],
+ ) -> Tuple[torch.Tensor, torch.Tensor]:
+ """
+ Embeds different types of prompts, returning both sparse and dense
+ embeddings.
+
+ Arguments:
+ points (tuple(torch.Tensor, torch.Tensor) or none): point coordinates
+ and labels to embed.
+ boxes (torch.Tensor or none): boxes to embed
+ masks (torch.Tensor or none): masks to embed
+
+ Returns:
+ torch.Tensor: sparse embeddings for the points and boxes, with shape
+ BxNx(embed_dim), where N is determined by the number of input points
+ and boxes.
+ torch.Tensor: dense embeddings for the masks, in the shape
+ Bx(embed_dim)x(embed_H)x(embed_W)
+ """
+ bs = self._get_batch_size(points, boxes, masks, text_embeds)
+ sparse_embeddings = torch.empty(
+ (bs, 0, self.embed_dim), device=self._get_device()
+ )
+ if points is not None:
+ coords, labels = points
+ point_embeddings = self._embed_points(coords, labels, pad=(boxes is None))
+ sparse_embeddings = torch.cat([sparse_embeddings, point_embeddings], dim=1)
+ if boxes is not None:
+ box_embeddings = self._embed_boxes(boxes)
+ sparse_embeddings = torch.cat([sparse_embeddings, box_embeddings], dim=1)
+
+ if text_embeds is not None:
+ sparse_embeddings = torch.cat([sparse_embeddings, text_embeds], dim=1)
+
+ if masks is not None:
+ dense_embeddings = self._embed_masks(masks)
+ else:
+ dense_embeddings = self.no_mask_embed.weight.reshape(1, -1, 1, 1).expand(
+ bs, -1, self.image_embedding_size[0], self.image_embedding_size[1]
+ )
+
+ return sparse_embeddings, dense_embeddings
+
+
+class PositionEmbeddingRandom(nn.Module):
+ """
+ Positional encoding using random spatial frequencies.
+ """
+
+ def __init__(self, num_pos_feats: int = 64, scale: Optional[float] = None) -> None:
+ super().__init__()
+ if scale is None or scale <= 0.0:
+ scale = 1.0
+ self.register_buffer(
+ "positional_encoding_gaussian_matrix",
+ scale * torch.randn((2, num_pos_feats)),
+ )
+
+ def _pe_encoding(self, coords: torch.Tensor) -> torch.Tensor:
+ """Positionally encode points that are normalized to [0,1]."""
+ # assuming coords are in [0, 1]^2 square and have d_1 x ... x d_n x 2 shape
+ coords = 2 * coords - 1
+
+ if coords.dtype != self.positional_encoding_gaussian_matrix.dtype:
+ coords = coords.to(self.positional_encoding_gaussian_matrix.dtype)
+
+ coords = coords @ self.positional_encoding_gaussian_matrix
+ coords = 2 * np.pi * coords
+ # outputs d_1 x ... x d_n x C shape
+ return torch.cat([torch.sin(coords), torch.cos(coords)], dim=-1)
+
+ def forward(self, size: Tuple[int, int]) -> torch.Tensor:
+ """Generate positional encoding for a grid of the specified size."""
+ h, w = size
+ device: Any = self.positional_encoding_gaussian_matrix.device
+ grid = torch.ones(
+ (h, w), device=device, dtype=self.positional_encoding_gaussian_matrix.dtype
+ )
+ y_embed = grid.cumsum(dim=0) - 0.5
+ x_embed = grid.cumsum(dim=1) - 0.5
+ y_embed = y_embed / h
+ x_embed = x_embed / w
+
+ pe = self._pe_encoding(torch.stack([x_embed, y_embed], dim=-1))
+ return pe.permute(2, 0, 1) # C x H x W
+
+ def forward_with_coords(
+ self, coords_input: torch.Tensor, image_size: Tuple[int, int]
+ ) -> torch.Tensor:
+ """Positionally encode points that are not normalized to [0,1]."""
+ coords = coords_input.clone()
+ coords[:, :, 0] = coords[:, :, 0] / image_size[1]
+ coords[:, :, 1] = coords[:, :, 1] / image_size[0]
+ return self._pe_encoding(coords.to(torch.float)) # B x N x C
diff --git a/py/evf_sam/model/segment_anything_2/sam2/modeling/sam/transformer.py b/py/evf_sam/model/segment_anything_2/sam2/modeling/sam/transformer.py
new file mode 100644
index 0000000..dae74f5
--- /dev/null
+++ b/py/evf_sam/model/segment_anything_2/sam2/modeling/sam/transformer.py
@@ -0,0 +1,330 @@
+# Copyright (c) Meta Platforms, Inc. and affiliates.
+# All rights reserved.
+
+# This source code is licensed under the license found in the
+# LICENSE file in the root directory of this source tree.
+
+import math
+import warnings
+from functools import partial
+from typing import Tuple, Type
+
+import torch
+import torch.nn.functional as F
+from torch import nn, Tensor
+
+from model.segment_anything_2.sam2.modeling.position_encoding import apply_rotary_enc, compute_axial_cis
+
+from model.segment_anything_2.sam2.modeling.sam2_utils import MLP
+from model.segment_anything_2.sam2.utils.misc import get_sdpa_settings
+
+warnings.simplefilter(action="ignore", category=FutureWarning)
+# OLD_GPU, USE_FLASH_ATTN, MATH_KERNEL_ON = get_sdpa_settings()
+USE_FLASH_ATTN = False
+MATH_KERNEL_ON = True
+OLD_GPU = True
+
+
+class TwoWayTransformer(nn.Module):
+ def __init__(
+ self,
+ depth: int,
+ embedding_dim: int,
+ num_heads: int,
+ mlp_dim: int,
+ activation: Type[nn.Module] = nn.ReLU,
+ attention_downsample_rate: int = 2,
+ ) -> None:
+ """
+ A transformer decoder that attends to an input image using
+ queries whose positional embedding is supplied.
+
+ Args:
+ depth (int): number of layers in the transformer
+ embedding_dim (int): the channel dimension for the input embeddings
+ num_heads (int): the number of heads for multihead attention. Must
+ divide embedding_dim
+ mlp_dim (int): the channel dimension internal to the MLP block
+ activation (nn.Module): the activation to use in the MLP block
+ """
+ super().__init__()
+ self.depth = depth
+ self.embedding_dim = embedding_dim
+ self.num_heads = num_heads
+ self.mlp_dim = mlp_dim
+ self.layers = nn.ModuleList()
+
+ for i in range(depth):
+ self.layers.append(
+ TwoWayAttentionBlock(
+ embedding_dim=embedding_dim,
+ num_heads=num_heads,
+ mlp_dim=mlp_dim,
+ activation=activation,
+ attention_downsample_rate=attention_downsample_rate,
+ skip_first_layer_pe=(i == 0),
+ )
+ )
+
+ self.final_attn_token_to_image = Attention(
+ embedding_dim, num_heads, downsample_rate=attention_downsample_rate
+ )
+ self.norm_final_attn = nn.LayerNorm(embedding_dim)
+
+ def forward(
+ self,
+ image_embedding: Tensor,
+ image_pe: Tensor,
+ point_embedding: Tensor,
+ ) -> Tuple[Tensor, Tensor]:
+ """
+ Args:
+ image_embedding (torch.Tensor): image to attend to. Should be shape
+ B x embedding_dim x h x w for any h and w.
+ image_pe (torch.Tensor): the positional encoding to add to the image. Must
+ have the same shape as image_embedding.
+ point_embedding (torch.Tensor): the embedding to add to the query points.
+ Must have shape B x N_points x embedding_dim for any N_points.
+
+ Returns:
+ torch.Tensor: the processed point_embedding
+ torch.Tensor: the processed image_embedding
+ """
+ # BxCxHxW -> BxHWxC == B x N_image_tokens x C
+ bs, c, h, w = image_embedding.shape
+ image_embedding = image_embedding.flatten(2).permute(0, 2, 1)
+ image_pe = image_pe.flatten(2).permute(0, 2, 1)
+
+ # Prepare queries
+ queries = point_embedding
+ keys = image_embedding
+
+ # Apply transformer blocks and final layernorm
+ for layer in self.layers:
+ queries, keys = layer(
+ queries=queries,
+ keys=keys,
+ query_pe=point_embedding,
+ key_pe=image_pe,
+ )
+
+ # Apply the final attention layer from the points to the image
+ q = queries + point_embedding
+ k = keys + image_pe
+ attn_out = self.final_attn_token_to_image(q=q, k=k, v=keys)
+ queries = queries + attn_out
+ queries = self.norm_final_attn(queries)
+
+ return queries, keys
+
+
+class TwoWayAttentionBlock(nn.Module):
+ def __init__(
+ self,
+ embedding_dim: int,
+ num_heads: int,
+ mlp_dim: int = 2048,
+ activation: Type[nn.Module] = nn.ReLU,
+ attention_downsample_rate: int = 2,
+ skip_first_layer_pe: bool = False,
+ ) -> None:
+ """
+ A transformer block with four layers: (1) self-attention of sparse
+ inputs, (2) cross attention of sparse inputs to dense inputs, (3) mlp
+ block on sparse inputs, and (4) cross attention of dense inputs to sparse
+ inputs.
+
+ Arguments:
+ embedding_dim (int): the channel dimension of the embeddings
+ num_heads (int): the number of heads in the attention layers
+ mlp_dim (int): the hidden dimension of the mlp block
+ activation (nn.Module): the activation of the mlp block
+ skip_first_layer_pe (bool): skip the PE on the first layer
+ """
+ super().__init__()
+ self.self_attn = Attention(embedding_dim, num_heads)
+ self.norm1 = nn.LayerNorm(embedding_dim)
+
+ self.cross_attn_token_to_image = Attention(
+ embedding_dim, num_heads, downsample_rate=attention_downsample_rate
+ )
+ self.norm2 = nn.LayerNorm(embedding_dim)
+
+ self.mlp = MLP(
+ embedding_dim, mlp_dim, embedding_dim, num_layers=2, activation=activation
+ )
+ self.norm3 = nn.LayerNorm(embedding_dim)
+
+ self.norm4 = nn.LayerNorm(embedding_dim)
+ self.cross_attn_image_to_token = Attention(
+ embedding_dim, num_heads, downsample_rate=attention_downsample_rate
+ )
+
+ self.skip_first_layer_pe = skip_first_layer_pe
+
+ def forward(
+ self, queries: Tensor, keys: Tensor, query_pe: Tensor, key_pe: Tensor
+ ) -> Tuple[Tensor, Tensor]:
+ # Self attention block
+ if self.skip_first_layer_pe:
+ queries = self.self_attn(q=queries, k=queries, v=queries)
+ else:
+ q = queries + query_pe
+ attn_out = self.self_attn(q=q, k=q, v=queries)
+ queries = queries + attn_out
+ queries = self.norm1(queries)
+
+ # Cross attention block, tokens attending to image embedding
+ q = queries + query_pe
+ k = keys + key_pe
+ attn_out = self.cross_attn_token_to_image(q=q, k=k, v=keys)
+ queries = queries + attn_out
+ queries = self.norm2(queries)
+
+ # MLP block
+ mlp_out = self.mlp(queries)
+ queries = queries + mlp_out
+ queries = self.norm3(queries)
+
+ # Cross attention block, image embedding attending to tokens
+ q = queries + query_pe
+ k = keys + key_pe
+ attn_out = self.cross_attn_image_to_token(q=k, k=q, v=queries)
+ keys = keys + attn_out
+ keys = self.norm4(keys)
+
+ return queries, keys
+
+
+class Attention(nn.Module):
+ """
+ An attention layer that allows for downscaling the size of the embedding
+ after projection to queries, keys, and values.
+ """
+
+ def __init__(
+ self,
+ embedding_dim: int,
+ num_heads: int,
+ downsample_rate: int = 1,
+ dropout: float = 0.0,
+ kv_in_dim: int = None,
+ ) -> None:
+ super().__init__()
+ self.embedding_dim = embedding_dim
+ self.kv_in_dim = kv_in_dim if kv_in_dim is not None else embedding_dim
+ self.internal_dim = embedding_dim // downsample_rate
+ self.num_heads = num_heads
+ assert (
+ self.internal_dim % num_heads == 0
+ ), "num_heads must divide embedding_dim."
+
+ self.q_proj = nn.Linear(embedding_dim, self.internal_dim)
+ self.k_proj = nn.Linear(self.kv_in_dim, self.internal_dim)
+ self.v_proj = nn.Linear(self.kv_in_dim, self.internal_dim)
+ self.out_proj = nn.Linear(self.internal_dim, embedding_dim)
+
+ self.dropout_p = dropout
+
+ def _separate_heads(self, x: Tensor, num_heads: int) -> Tensor:
+ b, n, c = x.shape
+ x = x.reshape(b, n, num_heads, c // num_heads)
+ return x.transpose(1, 2) # B x N_heads x N_tokens x C_per_head
+
+ def _recombine_heads(self, x: Tensor) -> Tensor:
+ b, n_heads, n_tokens, c_per_head = x.shape
+ x = x.transpose(1, 2)
+ return x.reshape(b, n_tokens, n_heads * c_per_head) # B x N_tokens x C
+
+ def forward(self, q: Tensor, k: Tensor, v: Tensor) -> Tensor:
+ # Input projections
+ q = self.q_proj(q)
+ k = self.k_proj(k)
+ v = self.v_proj(v)
+
+ # Separate into heads
+ q = self._separate_heads(q, self.num_heads)
+ k = self._separate_heads(k, self.num_heads)
+ v = self._separate_heads(v, self.num_heads)
+
+ dropout_p = self.dropout_p if self.training else 0.0
+ # Attention
+ with torch.backends.cuda.sdp_kernel(
+ enable_flash=USE_FLASH_ATTN,
+ # if Flash attention kernel is off, then math kernel needs to be enabled
+ enable_math=(OLD_GPU and dropout_p > 0.0) or MATH_KERNEL_ON,
+ enable_mem_efficient=OLD_GPU,
+ ):
+ out = F.scaled_dot_product_attention(q, k, v, dropout_p=dropout_p)
+
+ out = self._recombine_heads(out)
+ out = self.out_proj(out)
+
+ return out
+
+
+class RoPEAttention(Attention):
+ """Attention with rotary position encoding."""
+
+ def __init__(
+ self,
+ *args,
+ rope_theta=10000.0,
+ # whether to repeat q rope to match k length
+ # this is needed for cross-attention to memories
+ rope_k_repeat=False,
+ feat_sizes=(32, 32), # [w, h] for stride 16 feats at 512 resolution
+ **kwargs,
+ ):
+ super().__init__(*args, **kwargs)
+
+ self.compute_cis = partial(
+ compute_axial_cis, dim=self.internal_dim // self.num_heads, theta=rope_theta
+ )
+ freqs_cis = self.compute_cis(end_x=feat_sizes[0], end_y=feat_sizes[1])
+ self.freqs_cis = freqs_cis
+ self.rope_k_repeat = rope_k_repeat
+
+ def forward(
+ self, q: Tensor, k: Tensor, v: Tensor, num_k_exclude_rope: int = 0
+ ) -> Tensor:
+ # Input projections
+ q = self.q_proj(q)
+ k = self.k_proj(k)
+ v = self.v_proj(v)
+
+ # Separate into heads
+ q = self._separate_heads(q, self.num_heads)
+ k = self._separate_heads(k, self.num_heads)
+ v = self._separate_heads(v, self.num_heads)
+
+ # Apply rotary position encoding
+ w = h = math.sqrt(q.shape[-2])
+ self.freqs_cis = self.freqs_cis.to(q.device)
+ if self.freqs_cis.shape[0] != q.shape[-2]:
+ self.freqs_cis = self.compute_cis(end_x=w, end_y=h).to(q.device)
+ if q.shape[-2] != k.shape[-2]:
+ assert self.rope_k_repeat
+
+ num_k_rope = k.size(-2) - num_k_exclude_rope
+ q, k[:, :, :num_k_rope] = apply_rotary_enc(
+ q,
+ k[:, :, :num_k_rope],
+ freqs_cis=self.freqs_cis,
+ repeat_freqs_k=self.rope_k_repeat,
+ )
+
+ dropout_p = self.dropout_p if self.training else 0.0
+ # Attention
+ with torch.backends.cuda.sdp_kernel(
+ enable_flash=USE_FLASH_ATTN,
+ # if Flash attention kernel is off, then math kernel needs to be enabled
+ enable_math=(OLD_GPU and dropout_p > 0.0) or MATH_KERNEL_ON,
+ enable_mem_efficient=OLD_GPU,
+ ):
+ out = F.scaled_dot_product_attention(q, k, v, dropout_p=dropout_p)
+
+ out = self._recombine_heads(out)
+ out = self.out_proj(out)
+
+ return out
diff --git a/py/evf_sam/model/segment_anything_2/sam2/modeling/sam2_base.py b/py/evf_sam/model/segment_anything_2/sam2/modeling/sam2_base.py
new file mode 100644
index 0000000..746f2ed
--- /dev/null
+++ b/py/evf_sam/model/segment_anything_2/sam2/modeling/sam2_base.py
@@ -0,0 +1,833 @@
+# Copyright (c) Meta Platforms, Inc. and affiliates.
+# All rights reserved.
+
+# This source code is licensed under the license found in the
+# LICENSE file in the root directory of this source tree.
+
+import torch
+import torch.distributed
+import torch.nn.functional as F
+
+from torch.nn.init import trunc_normal_
+
+from model.segment_anything_2.sam2.modeling.sam.mask_decoder import MaskDecoder
+from model.segment_anything_2.sam2.modeling.sam.prompt_encoder import PromptEncoder
+from model.segment_anything_2.sam2.modeling.sam.transformer import TwoWayTransformer
+from model.segment_anything_2.sam2.modeling.sam2_utils import get_1d_sine_pe, MLP, select_closest_cond_frames
+
+# a large negative value as a placeholder score for missing objects
+NO_OBJ_SCORE = -1024.0
+
+
+class SAM2Base(torch.nn.Module):
+ def __init__(
+ self,
+ image_encoder,
+ memory_attention,
+ memory_encoder,
+ num_maskmem=7, # default 1 input frame + 6 previous frames
+ image_size=512,
+ backbone_stride=16, # stride of the image backbone output
+ sigmoid_scale_for_mem_enc=1.0, # scale factor for mask sigmoid prob
+ sigmoid_bias_for_mem_enc=0.0, # bias factor for mask sigmoid prob
+ # During evaluation, whether to binarize the sigmoid mask logits on interacted frames with clicks
+ binarize_mask_from_pts_for_mem_enc=False,
+ use_mask_input_as_output_without_sam=False, # on frames with mask input, whether to directly output the input mask without using a SAM prompt encoder + mask decoder
+ # The maximum number of conditioning frames to participate in the memory attention (-1 means no limit; if there are more conditioning frames than this limit,
+ # we only cross-attend to the temporally closest `max_cond_frames_in_attn` conditioning frames in the encoder when tracking each frame). This gives the model
+ # a temporal locality when handling a large number of annotated frames (since closer frames should be more important) and also avoids GPU OOM.
+ max_cond_frames_in_attn=-1,
+ # on the first frame, whether to directly add the no-memory embedding to the image feature
+ # (instead of using the transformer encoder)
+ directly_add_no_mem_embed=False,
+ # whether to use high-resolution feature maps in the SAM mask decoder
+ use_high_res_features_in_sam=False,
+ # whether to output multiple (3) masks for the first click on initial conditioning frames
+ multimask_output_in_sam=False,
+ # the minimum and maximum number of clicks to use multimask_output_in_sam (only relevant when `multimask_output_in_sam=True`;
+ # default is 1 for both, meaning that only the first click gives multimask output; also note that a box counts as two points)
+ multimask_min_pt_num=1,
+ multimask_max_pt_num=1,
+ # whether to also use multimask output for tracking (not just for the first click on initial conditioning frames; only relevant when `multimask_output_in_sam=True`)
+ multimask_output_for_tracking=False,
+ # Whether to use multimask tokens for obj ptr; Only relevant when both
+ # use_obj_ptrs_in_encoder=True and multimask_output_for_tracking=True
+ use_multimask_token_for_obj_ptr: bool = False,
+ # whether to use sigmoid to restrict ious prediction to [0-1]
+ iou_prediction_use_sigmoid=False,
+ # The memory bank's temporal stride during evaluation (i.e. the `r` parameter in XMem and Cutie; XMem and Cutie use r=5).
+ # For r>1, the (self.num_maskmem - 1) non-conditioning memory frames consist of
+ # (self.num_maskmem - 2) nearest frames from every r-th frames, plus the last frame.
+ memory_temporal_stride_for_eval=1,
+ # if `add_all_frames_to_correct_as_cond` is True, we also append to the conditioning frame list any frame that receives a later correction click
+ # if `add_all_frames_to_correct_as_cond` is False, we conditioning frame list to only use those initial conditioning frames
+ add_all_frames_to_correct_as_cond=False,
+ # whether to apply non-overlapping constraints on the object masks in the memory encoder during evaluation (to avoid/alleviate superposing masks)
+ non_overlap_masks_for_mem_enc=False,
+ # whether to cross-attend to object pointers from other frames (based on SAM output tokens) in the encoder
+ use_obj_ptrs_in_encoder=False,
+ # the maximum number of object pointers from other frames in encoder cross attention (only relevant when `use_obj_ptrs_in_encoder=True`)
+ max_obj_ptrs_in_encoder=16,
+ # whether to add temporal positional encoding to the object pointers in the encoder (only relevant when `use_obj_ptrs_in_encoder=True`)
+ add_tpos_enc_to_obj_ptrs=True,
+ # whether to add an extra linear projection layer for the temporal positional encoding in the object pointers to avoid potential interference
+ # with spatial positional encoding (only relevant when both `use_obj_ptrs_in_encoder=True` and `add_tpos_enc_to_obj_ptrs=True`)
+ proj_tpos_enc_in_obj_ptrs=False,
+ # whether to only attend to object pointers in the past (before the current frame) in the encoder during evaluation
+ # (only relevant when `use_obj_ptrs_in_encoder=True`; this might avoid pointer information too far in the future to distract the initial tracking)
+ only_obj_ptrs_in_the_past_for_eval=False,
+ # Whether to predict if there is an object in the frame
+ pred_obj_scores: bool = False,
+ # Whether to use an MLP to predict object scores
+ pred_obj_scores_mlp: bool = False,
+ # Only relevant if pred_obj_scores=True and use_obj_ptrs_in_encoder=True;
+ # Whether to have a fixed no obj pointer when there is no object present
+ # or to use it as an additive embedding with obj_ptr produced by decoder
+ fixed_no_obj_ptr: bool = False,
+ # Soft no object, i.e. mix in no_obj_ptr softly,
+ # hope to make recovery easier if there is a mistake and mitigate accumulation of errors
+ soft_no_obj_ptr: bool = False,
+ use_mlp_for_obj_ptr_proj: bool = False,
+ # extra arguments used to construct the SAM mask decoder; if not None, it should be a dict of kwargs to be passed into `MaskDecoder` class.
+ sam_mask_decoder_extra_args=None,
+ compile_image_encoder: bool = False,
+ ):
+ super().__init__()
+
+ # Part 1: the image backbone
+ self.image_encoder = image_encoder
+ # Use level 0, 1, 2 for high-res setting, or just level 2 for the default setting
+ self.use_high_res_features_in_sam = use_high_res_features_in_sam
+ self.num_feature_levels = 3 if use_high_res_features_in_sam else 1
+ self.use_obj_ptrs_in_encoder = use_obj_ptrs_in_encoder
+ self.max_obj_ptrs_in_encoder = max_obj_ptrs_in_encoder
+ if use_obj_ptrs_in_encoder:
+ # A conv layer to downsample the mask prompt to stride 4 (the same stride as
+ # low-res SAM mask logits) and to change its scales from 0~1 to SAM logit scale,
+ # so that it can be fed into the SAM mask decoder to generate a pointer.
+ self.mask_downsample = torch.nn.Conv2d(1, 1, kernel_size=4, stride=4)
+ self.add_tpos_enc_to_obj_ptrs = add_tpos_enc_to_obj_ptrs
+ if proj_tpos_enc_in_obj_ptrs:
+ assert add_tpos_enc_to_obj_ptrs # these options need to be used together
+ self.proj_tpos_enc_in_obj_ptrs = proj_tpos_enc_in_obj_ptrs
+ self.only_obj_ptrs_in_the_past_for_eval = only_obj_ptrs_in_the_past_for_eval
+
+ # Part 2: memory attention to condition current frame's visual features
+ # with memories (and obj ptrs) from past frames
+ self.memory_attention = memory_attention
+ self.hidden_dim = memory_attention.d_model
+
+ # Part 3: memory encoder for the previous frame's outputs
+ self.memory_encoder = memory_encoder
+ self.mem_dim = self.hidden_dim
+ if hasattr(self.memory_encoder, "out_proj") and hasattr(
+ self.memory_encoder.out_proj, "weight"
+ ):
+ # if there is compression of memories along channel dim
+ self.mem_dim = self.memory_encoder.out_proj.weight.shape[0]
+ self.num_maskmem = num_maskmem # Number of memories accessible
+ # Temporal encoding of the memories
+ self.maskmem_tpos_enc = torch.nn.Parameter(
+ torch.zeros(num_maskmem, 1, 1, self.mem_dim)
+ )
+ trunc_normal_(self.maskmem_tpos_enc, std=0.02)
+ # a single token to indicate no memory embedding from previous frames
+ self.no_mem_embed = torch.nn.Parameter(torch.zeros(1, 1, self.hidden_dim))
+ self.no_mem_pos_enc = torch.nn.Parameter(torch.zeros(1, 1, self.hidden_dim))
+ trunc_normal_(self.no_mem_embed, std=0.02)
+ trunc_normal_(self.no_mem_pos_enc, std=0.02)
+ self.directly_add_no_mem_embed = directly_add_no_mem_embed
+ # Apply sigmoid to the output raw mask logits (to turn them from
+ # range (-inf, +inf) to range (0, 1)) before feeding them into the memory encoder
+ self.sigmoid_scale_for_mem_enc = sigmoid_scale_for_mem_enc
+ self.sigmoid_bias_for_mem_enc = sigmoid_bias_for_mem_enc
+ self.binarize_mask_from_pts_for_mem_enc = binarize_mask_from_pts_for_mem_enc
+ self.non_overlap_masks_for_mem_enc = non_overlap_masks_for_mem_enc
+ self.memory_temporal_stride_for_eval = memory_temporal_stride_for_eval
+ # On frames with mask input, whether to directly output the input mask without
+ # using a SAM prompt encoder + mask decoder
+ self.use_mask_input_as_output_without_sam = use_mask_input_as_output_without_sam
+ self.multimask_output_in_sam = multimask_output_in_sam
+ self.multimask_min_pt_num = multimask_min_pt_num
+ self.multimask_max_pt_num = multimask_max_pt_num
+ self.multimask_output_for_tracking = multimask_output_for_tracking
+ self.use_multimask_token_for_obj_ptr = use_multimask_token_for_obj_ptr
+ self.iou_prediction_use_sigmoid = iou_prediction_use_sigmoid
+
+ # Part 4: SAM-style prompt encoder (for both mask and point inputs)
+ # and SAM-style mask decoder for the final mask output
+ self.image_size = image_size
+ self.backbone_stride = backbone_stride
+ self.sam_mask_decoder_extra_args = sam_mask_decoder_extra_args
+ self.pred_obj_scores = pred_obj_scores
+ self.pred_obj_scores_mlp = pred_obj_scores_mlp
+ self.fixed_no_obj_ptr = fixed_no_obj_ptr
+ self.soft_no_obj_ptr = soft_no_obj_ptr
+ if self.fixed_no_obj_ptr:
+ assert self.pred_obj_scores
+ assert self.use_obj_ptrs_in_encoder
+ if self.pred_obj_scores and self.use_obj_ptrs_in_encoder:
+ self.no_obj_ptr = torch.nn.Parameter(torch.zeros(1, self.hidden_dim))
+ trunc_normal_(self.no_obj_ptr, std=0.02)
+ self.use_mlp_for_obj_ptr_proj = use_mlp_for_obj_ptr_proj
+
+ self._build_sam_heads()
+ self.add_all_frames_to_correct_as_cond = add_all_frames_to_correct_as_cond
+ self.max_cond_frames_in_attn = max_cond_frames_in_attn
+
+ # Model compilation
+ if compile_image_encoder:
+ # Compile the forward function (not the full module) to allow loading checkpoints.
+ print(
+ "Image encoder compilation is enabled. First forward pass will be slow."
+ )
+ self.image_encoder.forward = torch.compile(
+ self.image_encoder.forward,
+ mode="max-autotune",
+ fullgraph=True,
+ dynamic=False,
+ )
+
+ @property
+ def device(self):
+ return next(self.parameters()).device
+
+ def forward(self, *args, **kwargs):
+ raise NotImplementedError(
+ "Please use the corresponding methods in SAM2VideoPredictor for inference."
+ "See notebooks/video_predictor_example.ipynb for an example."
+ )
+
+ def _build_sam_heads(self):
+ """Build SAM-style prompt encoder and mask decoder."""
+ self.sam_prompt_embed_dim = self.hidden_dim
+ self.sam_image_embedding_size = self.image_size // self.backbone_stride
+
+ # build PromptEncoder and MaskDecoder from SAM
+ # (their hyperparameters like `mask_in_chans=16` are from SAM code)
+ self.sam_prompt_encoder = PromptEncoder(
+ embed_dim=self.sam_prompt_embed_dim,
+ image_embedding_size=(
+ self.sam_image_embedding_size,
+ self.sam_image_embedding_size,
+ ),
+ input_image_size=(self.image_size, self.image_size),
+ mask_in_chans=16,
+ )
+ self.sam_mask_decoder = MaskDecoder(
+ num_multimask_outputs=3,
+ transformer=TwoWayTransformer(
+ depth=2,
+ embedding_dim=self.sam_prompt_embed_dim,
+ mlp_dim=2048,
+ num_heads=8,
+ ),
+ transformer_dim=self.sam_prompt_embed_dim,
+ iou_head_depth=3,
+ iou_head_hidden_dim=256,
+ use_high_res_features=self.use_high_res_features_in_sam,
+ iou_prediction_use_sigmoid=self.iou_prediction_use_sigmoid,
+ pred_obj_scores=self.pred_obj_scores,
+ pred_obj_scores_mlp=self.pred_obj_scores_mlp,
+ use_multimask_token_for_obj_ptr=self.use_multimask_token_for_obj_ptr,
+ **(self.sam_mask_decoder_extra_args or {}),
+ )
+ if self.use_obj_ptrs_in_encoder:
+ # a linear projection on SAM output tokens to turn them into object pointers
+ self.obj_ptr_proj = torch.nn.Linear(self.hidden_dim, self.hidden_dim)
+ if self.use_mlp_for_obj_ptr_proj:
+ self.obj_ptr_proj = MLP(
+ self.hidden_dim, self.hidden_dim, self.hidden_dim, 3
+ )
+ else:
+ self.obj_ptr_proj = torch.nn.Identity()
+ if self.proj_tpos_enc_in_obj_ptrs:
+ # a linear projection on temporal positional encoding in object pointers to
+ # avoid potential interference with spatial positional encoding
+ self.obj_ptr_tpos_proj = torch.nn.Linear(self.hidden_dim, self.mem_dim)
+ else:
+ self.obj_ptr_tpos_proj = torch.nn.Identity()
+
+ def _forward_sam_heads(
+ self,
+ backbone_features,
+ point_inputs=None,
+ mask_inputs=None,
+ text_inputs=None,
+ high_res_features=None,
+ multimask_output=False,
+ ):
+ """
+ Forward SAM prompt encoders and mask heads.
+
+ Inputs:
+ - backbone_features: image features of [B, C, H, W] shape
+ - point_inputs: a dictionary with "point_coords" and "point_labels", where
+ 1) "point_coords" has [B, P, 2] shape and float32 dtype and contains the
+ absolute pixel-unit coordinate in (x, y) format of the P input points
+ 2) "point_labels" has shape [B, P] and int32 dtype, where 1 means
+ positive clicks, 0 means negative clicks, and -1 means padding
+ - mask_inputs: a mask of [B, 1, H*16, W*16] shape, float or bool, with the
+ same spatial size as the image.
+ - high_res_features: either 1) None or 2) or a list of length 2 containing
+ two feature maps of [B, C, 4*H, 4*W] and [B, C, 2*H, 2*W] shapes respectively,
+ which will be used as high-resolution feature maps for SAM decoder.
+ - multimask_output: if it's True, we output 3 candidate masks and their 3
+ corresponding IoU estimates, and if it's False, we output only 1 mask and
+ its corresponding IoU estimate.
+
+ Outputs:
+ - low_res_multimasks: [B, M, H*4, W*4] shape (where M = 3 if
+ `multimask_output=True` and M = 1 if `multimask_output=False`), the SAM
+ output mask logits (before sigmoid) for the low-resolution masks, with 4x
+ the resolution (1/4 stride) of the input backbone_features.
+ - high_res_multimasks: [B, M, H*16, W*16] shape (where M = 3
+ if `multimask_output=True` and M = 1 if `multimask_output=False`),
+ upsampled from the low-resolution masks, with shape size as the image
+ (stride is 1 pixel).
+ - ious, [B, M] shape, where (where M = 3 if `multimask_output=True` and M = 1
+ if `multimask_output=False`), the estimated IoU of each output mask.
+ - low_res_masks: [B, 1, H*4, W*4] shape, the best mask in `low_res_multimasks`.
+ If `multimask_output=True`, it's the mask with the highest IoU estimate.
+ If `multimask_output=False`, it's the same as `low_res_multimasks`.
+ - high_res_masks: [B, 1, H*16, W*16] shape, the best mask in `high_res_multimasks`.
+ If `multimask_output=True`, it's the mask with the highest IoU estimate.
+ If `multimask_output=False`, it's the same as `high_res_multimasks`.
+ - obj_ptr: [B, C] shape, the object pointer vector for the output mask, extracted
+ based on the output token from the SAM mask decoder.
+ """
+ B = backbone_features.size(0)
+ device = backbone_features.device
+ assert backbone_features.size(1) == self.sam_prompt_embed_dim
+ assert backbone_features.size(2) == self.sam_image_embedding_size
+ assert backbone_features.size(3) == self.sam_image_embedding_size
+
+ # a) Handle point prompts
+ if point_inputs is not None:
+ sam_point_coords = point_inputs["point_coords"]
+ sam_point_labels = point_inputs["point_labels"]
+ assert sam_point_coords.size(0) == B and sam_point_labels.size(0) == B
+ else:
+ # If no points are provide, pad with an empty point (with label -1)
+ sam_point_coords = torch.zeros(B, 1, 2, device=device)
+ sam_point_labels = -torch.ones(B, 1, dtype=torch.int32, device=device)
+
+ # b) Handle mask prompts
+ if mask_inputs is not None:
+ # If mask_inputs is provided, downsize it into low-res mask input if needed
+ # and feed it as a dense mask prompt into the SAM mask encoder
+ assert len(mask_inputs.shape) == 4 and mask_inputs.shape[:2] == (B, 1)
+ if mask_inputs.shape[-2:] != self.sam_prompt_encoder.mask_input_size:
+ sam_mask_prompt = F.interpolate(
+ mask_inputs.float(),
+ size=self.sam_prompt_encoder.mask_input_size,
+ align_corners=False,
+ mode="bilinear",
+ antialias=True, # use antialias for downsampling
+ )
+ else:
+ sam_mask_prompt = mask_inputs
+ else:
+ # Otherwise, simply feed None (and SAM's prompt encoder will add
+ # a learned `no_mask_embed` to indicate no mask input in this case).
+ sam_mask_prompt = None
+
+ sparse_embeddings, dense_embeddings = self.sam_prompt_encoder(
+ points=(sam_point_coords, sam_point_labels),
+ boxes=None,
+ masks=sam_mask_prompt,
+ text_embeds=text_inputs
+ )
+ (
+ low_res_multimasks,
+ ious,
+ sam_output_tokens,
+ object_score_logits,
+ ) = self.sam_mask_decoder(
+ image_embeddings=backbone_features,
+ image_pe=self.sam_prompt_encoder.get_dense_pe(),
+ sparse_prompt_embeddings=sparse_embeddings,
+ dense_prompt_embeddings=dense_embeddings,
+ multimask_output=multimask_output,
+ repeat_image=False, # the image is already batched
+ high_res_features=high_res_features,
+ )
+ if self.pred_obj_scores:
+ is_obj_appearing = object_score_logits > 0
+
+ # Mask used for spatial memories is always a *hard* choice between obj and no obj,
+ # consistent with the actual mask prediction
+ low_res_multimasks = torch.where(
+ is_obj_appearing[:, None, None],
+ low_res_multimasks,
+ NO_OBJ_SCORE,
+ )
+
+ # convert masks from possibly bfloat16 (or float16) to float32
+ # (older PyTorch versions before 2.1 don't support `interpolate` on bf16)
+ low_res_multimasks = low_res_multimasks.float()
+ high_res_multimasks = F.interpolate(
+ low_res_multimasks,
+ size=(self.image_size, self.image_size),
+ mode="bilinear",
+ align_corners=False,
+ )
+
+ sam_output_token = sam_output_tokens[:, 0]
+ if multimask_output:
+ # take the best mask prediction (with the highest IoU estimation)
+ best_iou_inds = torch.argmax(ious, dim=-1)
+ batch_inds = torch.arange(B, device=device)
+ low_res_masks = low_res_multimasks[batch_inds, best_iou_inds].unsqueeze(1)
+ high_res_masks = high_res_multimasks[batch_inds, best_iou_inds].unsqueeze(1)
+ if sam_output_tokens.size(1) > 1:
+ sam_output_token = sam_output_tokens[batch_inds, best_iou_inds]
+ else:
+ low_res_masks, high_res_masks = low_res_multimasks, high_res_multimasks
+
+ # Extract object pointer from the SAM output token (with occlusion handling)
+ obj_ptr = self.obj_ptr_proj(sam_output_token)
+ if self.pred_obj_scores:
+ # Allow *soft* no obj ptr, unlike for masks
+ if self.soft_no_obj_ptr:
+ # Only hard possible with gt
+ assert not self.teacher_force_obj_scores_for_mem
+ lambda_is_obj_appearing = object_score_logits.sigmoid()
+ else:
+ lambda_is_obj_appearing = is_obj_appearing.float()
+
+ if self.fixed_no_obj_ptr:
+ obj_ptr = lambda_is_obj_appearing * obj_ptr
+ obj_ptr = obj_ptr + (1 - lambda_is_obj_appearing) * self.no_obj_ptr
+
+ return (
+ low_res_multimasks,
+ high_res_multimasks,
+ ious,
+ low_res_masks,
+ high_res_masks,
+ obj_ptr,
+ object_score_logits,
+ )
+
+ def _use_mask_as_output(self, backbone_features, high_res_features, mask_inputs):
+ """
+ Directly turn binary `mask_inputs` into a output mask logits without using SAM.
+ (same input and output shapes as in _forward_sam_heads above).
+ """
+ # Use -10/+10 as logits for neg/pos pixels (very close to 0/1 in prob after sigmoid).
+ out_scale, out_bias = 20.0, -10.0 # sigmoid(-10.0)=4.5398e-05
+ mask_inputs_float = mask_inputs.float()
+ high_res_masks = mask_inputs_float * out_scale + out_bias
+ low_res_masks = F.interpolate(
+ high_res_masks,
+ size=(high_res_masks.size(-2) // 4, high_res_masks.size(-1) // 4),
+ align_corners=False,
+ mode="bilinear",
+ antialias=True, # use antialias for downsampling
+ )
+ # a dummy IoU prediction of all 1's under mask input
+ ious = mask_inputs.new_ones(mask_inputs.size(0), 1).float()
+ if not self.use_obj_ptrs_in_encoder:
+ # all zeros as a dummy object pointer (of shape [B, C])
+ obj_ptr = torch.zeros(
+ mask_inputs.size(0), self.hidden_dim, device=mask_inputs.device
+ )
+ else:
+ # produce an object pointer using the SAM decoder from the mask input
+ _, _, _, _, _, obj_ptr, _ = self._forward_sam_heads(
+ backbone_features=backbone_features,
+ mask_inputs=self.mask_downsample(mask_inputs_float),
+ high_res_features=high_res_features,
+ )
+ # In this method, we are treating mask_input as output, e.g. using it directly to create spatial mem;
+ # Below, we follow the same design axiom to use mask_input to decide if obj appears or not instead of relying
+ # on the object_scores from the SAM decoder.
+ is_obj_appearing = torch.any(mask_inputs.flatten(1).float() > 0.0, dim=1)
+ is_obj_appearing = is_obj_appearing[..., None]
+ lambda_is_obj_appearing = is_obj_appearing.float()
+ object_score_logits = out_scale * lambda_is_obj_appearing + out_bias
+ if self.pred_obj_scores:
+ if self.fixed_no_obj_ptr:
+ obj_ptr = lambda_is_obj_appearing * obj_ptr
+ obj_ptr = obj_ptr + (1 - lambda_is_obj_appearing) * self.no_obj_ptr
+
+ return (
+ low_res_masks,
+ high_res_masks,
+ ious,
+ low_res_masks,
+ high_res_masks,
+ obj_ptr,
+ object_score_logits,
+ )
+
+ def forward_image(self, img_batch: torch.Tensor):
+ """Get the image feature on the input batch."""
+ backbone_out = self.image_encoder(img_batch)
+ if self.use_high_res_features_in_sam:
+ # precompute projected level 0 and level 1 features in SAM decoder
+ # to avoid running it again on every SAM click
+ backbone_out["backbone_fpn"][0] = self.sam_mask_decoder.conv_s0(
+ backbone_out["backbone_fpn"][0]
+ )
+ backbone_out["backbone_fpn"][1] = self.sam_mask_decoder.conv_s1(
+ backbone_out["backbone_fpn"][1]
+ )
+ return backbone_out
+
+ def _prepare_backbone_features(self, backbone_out):
+ """Prepare and flatten visual features."""
+ backbone_out = backbone_out.copy()
+ assert len(backbone_out["backbone_fpn"]) == len(backbone_out["vision_pos_enc"])
+ assert len(backbone_out["backbone_fpn"]) >= self.num_feature_levels
+
+ feature_maps = backbone_out["backbone_fpn"][-self.num_feature_levels :]
+ vision_pos_embeds = backbone_out["vision_pos_enc"][-self.num_feature_levels :]
+
+ feat_sizes = [(x.shape[-2], x.shape[-1]) for x in vision_pos_embeds]
+ # flatten NxCxHxW to HWxNxC
+ vision_feats = [x.flatten(2).permute(2, 0, 1) for x in feature_maps]
+ vision_pos_embeds = [x.flatten(2).permute(2, 0, 1) for x in vision_pos_embeds]
+
+ return backbone_out, vision_feats, vision_pos_embeds, feat_sizes
+
+ def _prepare_memory_conditioned_features(
+ self,
+ frame_idx,
+ is_init_cond_frame,
+ current_vision_feats,
+ current_vision_pos_embeds,
+ feat_sizes,
+ output_dict,
+ num_frames,
+ track_in_reverse=False, # tracking in reverse time order (for demo usage)
+ ):
+ """Fuse the current frame's visual feature map with previous memory."""
+ B = current_vision_feats[-1].size(1) # batch size on this frame
+ C = self.hidden_dim
+ H, W = feat_sizes[-1] # top-level (lowest-resolution) feature size
+ device = current_vision_feats[-1].device
+ # The case of `self.num_maskmem == 0` below is primarily used for reproducing SAM on images.
+ # In this case, we skip the fusion with any memory.
+ if self.num_maskmem == 0: # Disable memory and skip fusion
+ pix_feat = current_vision_feats[-1].permute(1, 2, 0).view(B, C, H, W)
+ return pix_feat
+
+ num_obj_ptr_tokens = 0
+ # Step 1: condition the visual features of the current frame on previous memories
+ if not is_init_cond_frame:
+ # Retrieve the memories encoded with the maskmem backbone
+ to_cat_memory, to_cat_memory_pos_embed = [], []
+ # Add conditioning frames's output first (all cond frames have t_pos=0 for
+ # when getting temporal positional embedding below)
+ assert len(output_dict["cond_frame_outputs"]) > 0
+ # Select a maximum number of temporally closest cond frames for cross attention
+ cond_outputs = output_dict["cond_frame_outputs"]
+ selected_cond_outputs, unselected_cond_outputs = select_closest_cond_frames(
+ frame_idx, cond_outputs, self.max_cond_frames_in_attn
+ )
+ t_pos_and_prevs = [(0, out) for out in selected_cond_outputs.values()]
+ # Add last (self.num_maskmem - 1) frames before current frame for non-conditioning memory
+ # the earliest one has t_pos=1 and the latest one has t_pos=self.num_maskmem-1
+ # We also allow taking the memory frame non-consecutively (with r>1), in which case
+ # we take (self.num_maskmem - 2) frames among every r-th frames plus the last frame.
+ r = self.memory_temporal_stride_for_eval
+ for t_pos in range(1, self.num_maskmem):
+ t_rel = self.num_maskmem - t_pos # how many frames before current frame
+ if t_rel == 1:
+ # for t_rel == 1, we take the last frame (regardless of r)
+ if not track_in_reverse:
+ # the frame immediately before this frame (i.e. frame_idx - 1)
+ prev_frame_idx = frame_idx - t_rel
+ else:
+ # the frame immediately after this frame (i.e. frame_idx + 1)
+ prev_frame_idx = frame_idx + t_rel
+ else:
+ # for t_rel >= 2, we take the memory frame from every r-th frames
+ if not track_in_reverse:
+ # first find the nearest frame among every r-th frames before this frame
+ # for r=1, this would be (frame_idx - 2)
+ prev_frame_idx = ((frame_idx - 2) // r) * r
+ # then seek further among every r-th frames
+ prev_frame_idx = prev_frame_idx - (t_rel - 2) * r
+ else:
+ # first find the nearest frame among every r-th frames after this frame
+ # for r=1, this would be (frame_idx + 2)
+ prev_frame_idx = -(-(frame_idx + 2) // r) * r
+ # then seek further among every r-th frames
+ prev_frame_idx = prev_frame_idx + (t_rel - 2) * r
+ out = output_dict["non_cond_frame_outputs"].get(prev_frame_idx, None)
+ if out is None:
+ # If an unselected conditioning frame is among the last (self.num_maskmem - 1)
+ # frames, we still attend to it as if it's a non-conditioning frame.
+ out = unselected_cond_outputs.get(prev_frame_idx, None)
+ t_pos_and_prevs.append((t_pos, out))
+
+ for t_pos, prev in t_pos_and_prevs:
+ if prev is None:
+ continue # skip padding frames
+ # "maskmem_features" might have been offloaded to CPU in demo use cases,
+ # so we load it back to GPU (it's a no-op if it's already on GPU).
+ feats = prev["maskmem_features"].cuda(non_blocking=True)
+ to_cat_memory.append(feats.flatten(2).permute(2, 0, 1))
+ # Spatial positional encoding (it might have been offloaded to CPU in eval)
+ maskmem_enc = prev["maskmem_pos_enc"][-1].cuda()
+ maskmem_enc = maskmem_enc.flatten(2).permute(2, 0, 1)
+ # Temporal positional encoding
+ maskmem_enc = (
+ maskmem_enc + self.maskmem_tpos_enc[self.num_maskmem - t_pos - 1]
+ )
+ to_cat_memory_pos_embed.append(maskmem_enc)
+
+ # Construct the list of past object pointers
+ if self.use_obj_ptrs_in_encoder:
+ max_obj_ptrs_in_encoder = min(num_frames, self.max_obj_ptrs_in_encoder)
+ # First add those object pointers from selected conditioning frames
+ # (optionally, only include object pointers in the past during evaluation)
+ if not self.training and self.only_obj_ptrs_in_the_past_for_eval:
+ ptr_cond_outputs = {
+ t: out
+ for t, out in selected_cond_outputs.items()
+ if (t >= frame_idx if track_in_reverse else t <= frame_idx)
+ }
+ else:
+ ptr_cond_outputs = selected_cond_outputs
+ pos_and_ptrs = [
+ # Temporal pos encoding contains how far away each pointer is from current frame
+ (abs(frame_idx - t), out["obj_ptr"])
+ for t, out in ptr_cond_outputs.items()
+ ]
+ # Add up to (max_obj_ptrs_in_encoder - 1) non-conditioning frames before current frame
+ for t_diff in range(1, max_obj_ptrs_in_encoder):
+ t = frame_idx + t_diff if track_in_reverse else frame_idx - t_diff
+ if t < 0 or (num_frames is not None and t >= num_frames):
+ break
+ out = output_dict["non_cond_frame_outputs"].get(
+ t, unselected_cond_outputs.get(t, None)
+ )
+ if out is not None:
+ pos_and_ptrs.append((t_diff, out["obj_ptr"]))
+ # If we have at least one object pointer, add them to the across attention
+ if len(pos_and_ptrs) > 0:
+ pos_list, ptrs_list = zip(*pos_and_ptrs)
+ # stack object pointers along dim=0 into [ptr_seq_len, B, C] shape
+ obj_ptrs = torch.stack(ptrs_list, dim=0)
+ # a temporal positional embedding based on how far each object pointer is from
+ # the current frame (sine embedding normalized by the max pointer num).
+ if self.add_tpos_enc_to_obj_ptrs:
+ t_diff_max = max_obj_ptrs_in_encoder - 1
+ tpos_dim = C if self.proj_tpos_enc_in_obj_ptrs else self.mem_dim
+ obj_pos = torch.tensor(pos_list, device=device)
+ obj_pos = get_1d_sine_pe(obj_pos / t_diff_max, dim=tpos_dim)
+ obj_pos = self.obj_ptr_tpos_proj(obj_pos)
+ obj_pos = obj_pos.unsqueeze(1).expand(-1, B, self.mem_dim)
+ else:
+ obj_pos = obj_ptrs.new_zeros(len(pos_list), B, self.mem_dim)
+ if self.mem_dim < C:
+ # split a pointer into (C // self.mem_dim) tokens for self.mem_dim < C
+ obj_ptrs = obj_ptrs.reshape(
+ -1, B, C // self.mem_dim, self.mem_dim
+ )
+ obj_ptrs = obj_ptrs.permute(0, 2, 1, 3).flatten(0, 1)
+ obj_pos = obj_pos.repeat_interleave(C // self.mem_dim, dim=0)
+ to_cat_memory.append(obj_ptrs)
+ to_cat_memory_pos_embed.append(obj_pos)
+ num_obj_ptr_tokens = obj_ptrs.shape[0]
+ else:
+ num_obj_ptr_tokens = 0
+ else:
+ # for initial conditioning frames, encode them without using any previous memory
+ if self.directly_add_no_mem_embed:
+ # directly add no-mem embedding (instead of using the transformer encoder)
+ pix_feat_with_mem = current_vision_feats[-1] + self.no_mem_embed
+ pix_feat_with_mem = pix_feat_with_mem.permute(1, 2, 0).view(B, C, H, W)
+ return pix_feat_with_mem
+
+ # Use a dummy token on the first frame (to avoid emtpy memory input to tranformer encoder)
+ to_cat_memory = [self.no_mem_embed.expand(1, B, self.mem_dim)]
+ to_cat_memory_pos_embed = [self.no_mem_pos_enc.expand(1, B, self.mem_dim)]
+
+ # Step 2: Concatenate the memories and forward through the transformer encoder
+ memory = torch.cat(to_cat_memory, dim=0)
+ memory_pos_embed = torch.cat(to_cat_memory_pos_embed, dim=0)
+
+ pix_feat_with_mem = self.memory_attention(
+ curr=current_vision_feats,
+ curr_pos=current_vision_pos_embeds,
+ memory=memory,
+ memory_pos=memory_pos_embed,
+ num_obj_ptr_tokens=num_obj_ptr_tokens,
+ )
+ # reshape the output (HW)BC => BCHW
+ pix_feat_with_mem = pix_feat_with_mem.permute(1, 2, 0).view(B, C, H, W)
+ return pix_feat_with_mem
+
+ def _encode_new_memory(
+ self,
+ current_vision_feats,
+ feat_sizes,
+ pred_masks_high_res,
+ is_mask_from_pts,
+ ):
+ """Encode the current image and its prediction into a memory feature."""
+ B = current_vision_feats[-1].size(1) # batch size on this frame
+ C = self.hidden_dim
+ H, W = feat_sizes[-1] # top-level (lowest-resolution) feature size
+ # top-level feature, (HW)BC => BCHW
+ pix_feat = current_vision_feats[-1].permute(1, 2, 0).view(B, C, H, W)
+ if self.non_overlap_masks_for_mem_enc and not self.training:
+ # optionally, apply non-overlapping constraints to the masks (it's applied
+ # in the batch dimension and should only be used during eval, where all
+ # the objects come from the same video under batch size 1).
+ pred_masks_high_res = self._apply_non_overlapping_constraints(
+ pred_masks_high_res
+ )
+ # scale the raw mask logits with a temperature before applying sigmoid
+ binarize = self.binarize_mask_from_pts_for_mem_enc and is_mask_from_pts
+ if binarize and not self.training:
+ mask_for_mem = (pred_masks_high_res > 0).float()
+ else:
+ # apply sigmoid on the raw mask logits to turn them into range (0, 1)
+ mask_for_mem = torch.sigmoid(pred_masks_high_res)
+ # apply scale and bias terms to the sigmoid probabilities
+ if self.sigmoid_scale_for_mem_enc != 1.0:
+ mask_for_mem = mask_for_mem * self.sigmoid_scale_for_mem_enc
+ if self.sigmoid_bias_for_mem_enc != 0.0:
+ mask_for_mem = mask_for_mem + self.sigmoid_bias_for_mem_enc
+ maskmem_out = self.memory_encoder(
+ pix_feat, mask_for_mem, skip_mask_sigmoid=True # sigmoid already applied
+ )
+ maskmem_features = maskmem_out["vision_features"]
+ maskmem_pos_enc = maskmem_out["vision_pos_enc"]
+
+ return maskmem_features, maskmem_pos_enc
+
+ def track_step(
+ self,
+ frame_idx,
+ is_init_cond_frame,
+ current_vision_feats,
+ current_vision_pos_embeds,
+ feat_sizes,
+ point_inputs,
+ mask_inputs,
+ output_dict,
+ num_frames,
+ track_in_reverse=False, # tracking in reverse time order (for demo usage)
+ # Whether to run the memory encoder on the predicted masks. Sometimes we might want
+ # to skip the memory encoder with `run_mem_encoder=False`. For example,
+ # in demo we might call `track_step` multiple times for each user click,
+ # and only encode the memory when the user finalizes their clicks. And in ablation
+ # settings like SAM training on static images, we don't need the memory encoder.
+ run_mem_encoder=True,
+ # The previously predicted SAM mask logits (which can be fed together with new clicks in demo).
+ prev_sam_mask_logits=None,
+ text_inputs=None,
+ ):
+ current_out = {"point_inputs": point_inputs, "mask_inputs": mask_inputs}
+ # High-resolution feature maps for the SAM head, reshape (HW)BC => BCHW
+ if len(current_vision_feats) > 1:
+ high_res_features = [
+ x.permute(1, 2, 0).view(x.size(1), x.size(2), *s)
+ for x, s in zip(current_vision_feats[:-1], feat_sizes[:-1])
+ ]
+ else:
+ high_res_features = None
+ if mask_inputs is not None and self.use_mask_input_as_output_without_sam:
+ # When use_mask_input_as_output_without_sam=True, we directly output the mask input
+ # (see it as a GT mask) without using a SAM prompt encoder + mask decoder.
+ pix_feat = current_vision_feats[-1].permute(1, 2, 0)
+ pix_feat = pix_feat.view(-1, self.hidden_dim, *feat_sizes[-1])
+ sam_outputs = self._use_mask_as_output(
+ pix_feat, high_res_features, mask_inputs
+ )
+ else:
+ # fused the visual feature with previous memory features in the memory bank
+ pix_feat_with_mem = self._prepare_memory_conditioned_features(
+ frame_idx=frame_idx,
+ is_init_cond_frame=is_init_cond_frame,
+ current_vision_feats=current_vision_feats[-1:],
+ current_vision_pos_embeds=current_vision_pos_embeds[-1:],
+ feat_sizes=feat_sizes[-1:],
+ output_dict=output_dict,
+ num_frames=num_frames,
+ track_in_reverse=track_in_reverse,
+ )
+ # apply SAM-style segmentation head
+ # here we might feed previously predicted low-res SAM mask logits into the SAM mask decoder,
+ # e.g. in demo where such logits come from earlier interaction instead of correction sampling
+ # (in this case, any `mask_inputs` shouldn't reach here as they are sent to _use_mask_as_output instead)
+ if prev_sam_mask_logits is not None:
+ assert point_inputs is not None and mask_inputs is None
+ mask_inputs = prev_sam_mask_logits
+ multimask_output = self._use_multimask(is_init_cond_frame, point_inputs)
+ sam_outputs = self._forward_sam_heads(
+ backbone_features=pix_feat_with_mem,
+ point_inputs=point_inputs,
+ mask_inputs=mask_inputs,
+ high_res_features=high_res_features,
+ multimask_output=multimask_output,
+ text_inputs=text_inputs
+ )
+ (
+ _,
+ _,
+ _,
+ low_res_masks,
+ high_res_masks,
+ obj_ptr,
+ _,
+ ) = sam_outputs
+
+ current_out["pred_masks"] = low_res_masks
+ current_out["pred_masks_high_res"] = high_res_masks
+ current_out["obj_ptr"] = obj_ptr
+
+ # Finally run the memory encoder on the predicted mask to encode
+ # it into a new memory feature (that can be used in future frames)
+ if run_mem_encoder and self.num_maskmem > 0:
+ high_res_masks_for_mem_enc = high_res_masks
+ maskmem_features, maskmem_pos_enc = self._encode_new_memory(
+ current_vision_feats=current_vision_feats,
+ feat_sizes=feat_sizes,
+ pred_masks_high_res=high_res_masks_for_mem_enc,
+ is_mask_from_pts=(point_inputs is not None),
+ )
+ current_out["maskmem_features"] = maskmem_features
+ current_out["maskmem_pos_enc"] = maskmem_pos_enc
+ else:
+ current_out["maskmem_features"] = None
+ current_out["maskmem_pos_enc"] = None
+
+ return current_out
+
+ def _use_multimask(self, is_init_cond_frame, point_inputs):
+ """Whether to use multimask output in the SAM head."""
+ num_pts = 0 if point_inputs is None else point_inputs["point_labels"].size(1)
+ multimask_output = (
+ self.multimask_output_in_sam
+ and (is_init_cond_frame or self.multimask_output_for_tracking)
+ and (self.multimask_min_pt_num <= num_pts <= self.multimask_max_pt_num)
+ )
+ return multimask_output
+
+ def _apply_non_overlapping_constraints(self, pred_masks):
+ """
+ Apply non-overlapping constraints to the object scores in pred_masks. Here we
+ keep only the highest scoring object at each spatial location in pred_masks.
+ """
+ batch_size = pred_masks.size(0)
+ if batch_size == 1:
+ return pred_masks
+
+ device = pred_masks.device
+ # "max_obj_inds": object index of the object with the highest score at each location
+ max_obj_inds = torch.argmax(pred_masks, dim=0, keepdim=True)
+ # "batch_obj_inds": object index of each object slice (along dim 0) in `pred_masks`
+ batch_obj_inds = torch.arange(batch_size, device=device)[:, None, None, None]
+ keep = max_obj_inds == batch_obj_inds
+ # suppress overlapping regions' scores below -10.0 so that the foreground regions
+ # don't overlap (here sigmoid(-10.0)=4.5398e-05)
+ pred_masks = torch.where(keep, pred_masks, torch.clamp(pred_masks, max=-10.0))
+ return pred_masks
diff --git a/py/evf_sam/model/segment_anything_2/sam2/modeling/sam2_utils.py b/py/evf_sam/model/segment_anything_2/sam2/modeling/sam2_utils.py
new file mode 100644
index 0000000..6d97059
--- /dev/null
+++ b/py/evf_sam/model/segment_anything_2/sam2/modeling/sam2_utils.py
@@ -0,0 +1,149 @@
+# Copyright (c) Meta Platforms, Inc. and affiliates.
+# All rights reserved.
+
+# This source code is licensed under the license found in the
+# LICENSE file in the root directory of this source tree.
+
+
+import copy
+
+import torch
+import torch.nn as nn
+import torch.nn.functional as F
+
+
+def select_closest_cond_frames(frame_idx, cond_frame_outputs, max_cond_frame_num):
+ """
+ Select up to `max_cond_frame_num` conditioning frames from `cond_frame_outputs`
+ that are temporally closest to the current frame at `frame_idx`. Here, we take
+ - a) the closest conditioning frame before `frame_idx` (if any);
+ - b) the closest conditioning frame after `frame_idx` (if any);
+ - c) any other temporally closest conditioning frames until reaching a total
+ of `max_cond_frame_num` conditioning frames.
+
+ Outputs:
+ - selected_outputs: selected items (keys & values) from `cond_frame_outputs`.
+ - unselected_outputs: items (keys & values) not selected in `cond_frame_outputs`.
+ """
+ if max_cond_frame_num == -1 or len(cond_frame_outputs) <= max_cond_frame_num:
+ selected_outputs = cond_frame_outputs
+ unselected_outputs = {}
+ else:
+ assert max_cond_frame_num >= 2, "we should allow using 2+ conditioning frames"
+ selected_outputs = {}
+
+ # the closest conditioning frame before `frame_idx` (if any)
+ idx_before = max((t for t in cond_frame_outputs if t < frame_idx), default=None)
+ if idx_before is not None:
+ selected_outputs[idx_before] = cond_frame_outputs[idx_before]
+
+ # the closest conditioning frame after `frame_idx` (if any)
+ idx_after = min((t for t in cond_frame_outputs if t >= frame_idx), default=None)
+ if idx_after is not None:
+ selected_outputs[idx_after] = cond_frame_outputs[idx_after]
+
+ # add other temporally closest conditioning frames until reaching a total
+ # of `max_cond_frame_num` conditioning frames.
+ num_remain = max_cond_frame_num - len(selected_outputs)
+ inds_remain = sorted(
+ (t for t in cond_frame_outputs if t not in selected_outputs),
+ key=lambda x: abs(x - frame_idx),
+ )[:num_remain]
+ selected_outputs.update((t, cond_frame_outputs[t]) for t in inds_remain)
+ unselected_outputs = {
+ t: v for t, v in cond_frame_outputs.items() if t not in selected_outputs
+ }
+
+ return selected_outputs, unselected_outputs
+
+
+def get_1d_sine_pe(pos_inds, dim, temperature=10000):
+ """
+ Get 1D sine positional embedding as in the original Transformer paper.
+ """
+ pe_dim = dim // 2
+ dim_t = torch.arange(pe_dim, dtype=torch.float32, device=pos_inds.device)
+ dim_t = temperature ** (2 * (dim_t // 2) / pe_dim)
+
+ pos_embed = pos_inds.unsqueeze(-1) / dim_t
+ pos_embed = torch.cat([pos_embed.sin(), pos_embed.cos()], dim=-1)
+ return pos_embed
+
+
+def get_activation_fn(activation):
+ """Return an activation function given a string"""
+ if activation == "relu":
+ return F.relu
+ if activation == "gelu":
+ return F.gelu
+ if activation == "glu":
+ return F.glu
+ raise RuntimeError(f"activation should be relu/gelu, not {activation}.")
+
+
+def get_clones(module, N):
+ return nn.ModuleList([copy.deepcopy(module) for i in range(N)])
+
+
+class DropPath(nn.Module):
+ # adapted from https://github.com/huggingface/pytorch-image-models/blob/main/timm/layers/drop.py
+ def __init__(self, drop_prob=0.0, scale_by_keep=True):
+ super(DropPath, self).__init__()
+ self.drop_prob = drop_prob
+ self.scale_by_keep = scale_by_keep
+
+ def forward(self, x):
+ if self.drop_prob == 0.0 or not self.training:
+ return x
+ keep_prob = 1 - self.drop_prob
+ shape = (x.shape[0],) + (1,) * (x.ndim - 1)
+ random_tensor = x.new_empty(shape).bernoulli_(keep_prob)
+ if keep_prob > 0.0 and self.scale_by_keep:
+ random_tensor.div_(keep_prob)
+ return x * random_tensor
+
+
+# Lightly adapted from
+# https://github.com/facebookresearch/MaskFormer/blob/main/mask_former/modeling/transformer/transformer_predictor.py # noqa
+class MLP(nn.Module):
+ def __init__(
+ self,
+ input_dim: int,
+ hidden_dim: int,
+ output_dim: int,
+ num_layers: int,
+ activation: nn.Module = nn.ReLU,
+ sigmoid_output: bool = False,
+ ) -> None:
+ super().__init__()
+ self.num_layers = num_layers
+ h = [hidden_dim] * (num_layers - 1)
+ self.layers = nn.ModuleList(
+ nn.Linear(n, k) for n, k in zip([input_dim] + h, h + [output_dim])
+ )
+ self.sigmoid_output = sigmoid_output
+ self.act = activation()
+
+ def forward(self, x):
+ for i, layer in enumerate(self.layers):
+ x = self.act(layer(x)) if i < self.num_layers - 1 else layer(x)
+ if self.sigmoid_output:
+ x = F.sigmoid(x)
+ return x
+
+
+# From https://github.com/facebookresearch/detectron2/blob/main/detectron2/layers/batch_norm.py # noqa
+# Itself from https://github.com/facebookresearch/ConvNeXt/blob/d1fa8f6fef0a165b27399986cc2bdacc92777e40/models/convnext.py#L119 # noqa
+class LayerNorm2d(nn.Module):
+ def __init__(self, num_channels: int, eps: float = 1e-6) -> None:
+ super().__init__()
+ self.weight = nn.Parameter(torch.ones(num_channels))
+ self.bias = nn.Parameter(torch.zeros(num_channels))
+ self.eps = eps
+
+ def forward(self, x: torch.Tensor) -> torch.Tensor:
+ u = x.mean(1, keepdim=True)
+ s = (x - u).pow(2).mean(1, keepdim=True)
+ x = (x - u) / torch.sqrt(s + self.eps)
+ x = self.weight[:, None, None] * x + self.bias[:, None, None]
+ return x
diff --git a/py/evf_sam/model/segment_anything_2/sam2/sam2_image_predictor.py b/py/evf_sam/model/segment_anything_2/sam2/sam2_image_predictor.py
new file mode 100644
index 0000000..5b7d762
--- /dev/null
+++ b/py/evf_sam/model/segment_anything_2/sam2/sam2_image_predictor.py
@@ -0,0 +1,446 @@
+# Copyright (c) Meta Platforms, Inc. and affiliates.
+# All rights reserved.
+
+# This source code is licensed under the license found in the
+# LICENSE file in the root directory of this source tree.
+
+import logging
+
+from typing import List, Optional, Tuple, Union
+
+import numpy as np
+import torch
+from PIL.Image import Image
+
+from model.segment_anything_2.sam2.modeling.sam2_base import SAM2Base
+
+from model.segment_anything_2.sam2.utils.transforms import SAM2Transforms
+
+
+class SAM2ImagePredictor:
+ def __init__(
+ self,
+ sam_model: SAM2Base,
+ mask_threshold=0.0,
+ max_hole_area=0.0,
+ max_sprinkle_area=0.0,
+ ) -> None:
+ """
+ Uses SAM-2 to calculate the image embedding for an image, and then
+ allow repeated, efficient mask prediction given prompts.
+
+ Arguments:
+ sam_model (Sam-2): The model to use for mask prediction.
+ mask_threshold (float): The threshold to use when converting mask logits
+ to binary masks. Masks are thresholded at 0 by default.
+ fill_hole_area (int): If fill_hole_area > 0, we fill small holes in up to
+ the maximum area of fill_hole_area in low_res_masks.
+ """
+ super().__init__()
+ self.model = sam_model
+ self._transforms = SAM2Transforms(
+ resolution=self.model.image_size,
+ mask_threshold=mask_threshold,
+ max_hole_area=max_hole_area,
+ max_sprinkle_area=max_sprinkle_area,
+ )
+
+ # Predictor state
+ self._is_image_set = False
+ self._features = None
+ self._orig_hw = None
+ # Whether the predictor is set for single image or a batch of images
+ self._is_batch = False
+
+ # Predictor config
+ self.mask_threshold = mask_threshold
+
+ # Spatial dim for backbone feature maps
+ self._bb_feat_sizes = [
+ (256, 256),
+ (128, 128),
+ (64, 64),
+ ]
+
+ @torch.no_grad()
+ def set_image(
+ self,
+ image: Union[np.ndarray, Image],
+ ) -> None:
+ """
+ Calculates the image embeddings for the provided image, allowing
+ masks to be predicted with the 'predict' method.
+
+ Arguments:
+ image (np.ndarray or PIL Image): The input image to embed in RGB format. The image should be in HWC format if np.ndarray, or WHC format if PIL Image
+ with pixel values in [0, 255].
+ image_format (str): The color format of the image, in ['RGB', 'BGR'].
+ """
+ self.reset_predictor()
+ # Transform the image to the form expected by the model
+ if isinstance(image, np.ndarray):
+ logging.info("For numpy array image, we assume (HxWxC) format")
+ self._orig_hw = [image.shape[:2]]
+ elif isinstance(image, Image):
+ w, h = image.size
+ self._orig_hw = [(h, w)]
+ else:
+ raise NotImplementedError("Image format not supported")
+
+ input_image = self._transforms(image)
+ input_image = input_image[None, ...].to(self.device)
+
+ assert (
+ len(input_image.shape) == 4 and input_image.shape[1] == 3
+ ), f"input_image must be of size 1x3xHxW, got {input_image.shape}"
+ logging.info("Computing image embeddings for the provided image...")
+ backbone_out = self.model.forward_image(input_image)
+ _, vision_feats, _, _ = self.model._prepare_backbone_features(backbone_out)
+ # Add no_mem_embed, which is added to the lowest rest feat. map during training on videos
+ if self.model.directly_add_no_mem_embed:
+ vision_feats[-1] = vision_feats[-1] + self.model.no_mem_embed
+
+ feats = [
+ feat.permute(1, 2, 0).view(1, -1, *feat_size)
+ for feat, feat_size in zip(vision_feats[::-1], self._bb_feat_sizes[::-1])
+ ][::-1]
+ self._features = {"image_embed": feats[-1], "high_res_feats": feats[:-1]}
+ self._is_image_set = True
+ logging.info("Image embeddings computed.")
+
+ @torch.no_grad()
+ def set_image_batch(
+ self,
+ image_list: List[Union[np.ndarray]],
+ ) -> None:
+ """
+ Calculates the image embeddings for the provided image batch, allowing
+ masks to be predicted with the 'predict_batch' method.
+
+ Arguments:
+ image_list (List[np.ndarray]): The input images to embed in RGB format. The image should be in HWC format if np.ndarray
+ with pixel values in [0, 255].
+ """
+ self.reset_predictor()
+ assert isinstance(image_list, list)
+ self._orig_hw = []
+ for image in image_list:
+ assert isinstance(
+ image, np.ndarray
+ ), "Images are expected to be an np.ndarray in RGB format, and of shape HWC"
+ self._orig_hw.append(image.shape[:2])
+ # Transform the image to the form expected by the model
+ img_batch = self._transforms.forward_batch(image_list)
+ img_batch = img_batch.to(self.device)
+ batch_size = img_batch.shape[0]
+ assert (
+ len(img_batch.shape) == 4 and img_batch.shape[1] == 3
+ ), f"img_batch must be of size Bx3xHxW, got {img_batch.shape}"
+ logging.info("Computing image embeddings for the provided images...")
+ backbone_out = self.model.forward_image(img_batch)
+ _, vision_feats, _, _ = self.model._prepare_backbone_features(backbone_out)
+ # Add no_mem_embed, which is added to the lowest rest feat. map during training on videos
+ if self.model.directly_add_no_mem_embed:
+ vision_feats[-1] = vision_feats[-1] + self.model.no_mem_embed
+
+ feats = [
+ feat.permute(1, 2, 0).view(batch_size, -1, *feat_size)
+ for feat, feat_size in zip(vision_feats[::-1], self._bb_feat_sizes[::-1])
+ ][::-1]
+ self._features = {"image_embed": feats[-1], "high_res_feats": feats[:-1]}
+ self._is_image_set = True
+ self._is_batch = True
+ logging.info("Image embeddings computed.")
+
+ def predict_batch(
+ self,
+ point_coords_batch: List[np.ndarray] = None,
+ point_labels_batch: List[np.ndarray] = None,
+ box_batch: List[np.ndarray] = None,
+ mask_input_batch: List[np.ndarray] = None,
+ multimask_output: bool = True,
+ return_logits: bool = False,
+ normalize_coords=True,
+ ) -> Tuple[List[np.ndarray], List[np.ndarray], List[np.ndarray]]:
+ """This function is very similar to predict(...), however it is used for batched mode, when the model is expected to generate predictions on multiple images.
+ It returns a tupele of lists of masks, ious, and low_res_masks_logits.
+ """
+ assert self._is_batch, "This function should only be used when in batched mode"
+ if not self._is_image_set:
+ raise RuntimeError(
+ "An image must be set with .set_image_batch(...) before mask prediction."
+ )
+ num_images = len(self._features["image_embed"])
+ all_masks = []
+ all_ious = []
+ all_low_res_masks = []
+ for img_idx in range(num_images):
+ # Transform input prompts
+ point_coords = (
+ point_coords_batch[img_idx] if point_coords_batch is not None else None
+ )
+ point_labels = (
+ point_labels_batch[img_idx] if point_labels_batch is not None else None
+ )
+ box = box_batch[img_idx] if box_batch is not None else None
+ mask_input = (
+ mask_input_batch[img_idx] if mask_input_batch is not None else None
+ )
+ mask_input, unnorm_coords, labels, unnorm_box = self._prep_prompts(
+ point_coords,
+ point_labels,
+ box,
+ mask_input,
+ normalize_coords,
+ img_idx=img_idx,
+ )
+ masks, iou_predictions, low_res_masks = self._predict(
+ unnorm_coords,
+ labels,
+ unnorm_box,
+ mask_input,
+ multimask_output,
+ return_logits=return_logits,
+ img_idx=img_idx,
+ )
+ masks_np = masks.squeeze(0).float().detach().cpu().numpy()
+ iou_predictions_np = (
+ iou_predictions.squeeze(0).float().detach().cpu().numpy()
+ )
+ low_res_masks_np = low_res_masks.squeeze(0).float().detach().cpu().numpy()
+ all_masks.append(masks_np)
+ all_ious.append(iou_predictions_np)
+ all_low_res_masks.append(low_res_masks_np)
+
+ return all_masks, all_ious, all_low_res_masks
+
+ def predict(
+ self,
+ point_coords: Optional[np.ndarray] = None,
+ point_labels: Optional[np.ndarray] = None,
+ box: Optional[np.ndarray] = None,
+ mask_input: Optional[np.ndarray] = None,
+ multimask_output: bool = True,
+ return_logits: bool = False,
+ normalize_coords=True,
+ ) -> Tuple[np.ndarray, np.ndarray, np.ndarray]:
+ """
+ Predict masks for the given input prompts, using the currently set image.
+
+ Arguments:
+ point_coords (np.ndarray or None): A Nx2 array of point prompts to the
+ model. Each point is in (X,Y) in pixels.
+ point_labels (np.ndarray or None): A length N array of labels for the
+ point prompts. 1 indicates a foreground point and 0 indicates a
+ background point.
+ box (np.ndarray or None): A length 4 array given a box prompt to the
+ model, in XYXY format.
+ mask_input (np.ndarray): A low resolution mask input to the model, typically
+ coming from a previous prediction iteration. Has form 1xHxW, where
+ for SAM, H=W=256.
+ multimask_output (bool): If true, the model will return three masks.
+ For ambiguous input prompts (such as a single click), this will often
+ produce better masks than a single prediction. If only a single
+ mask is needed, the model's predicted quality score can be used
+ to select the best mask. For non-ambiguous prompts, such as multiple
+ input prompts, multimask_output=False can give better results.
+ return_logits (bool): If true, returns un-thresholded masks logits
+ instead of a binary mask.
+ normalize_coords (bool): If true, the point coordinates will be normalized to the range [0,1] and point_coords is expected to be wrt. image dimensions.
+
+ Returns:
+ (np.ndarray): The output masks in CxHxW format, where C is the
+ number of masks, and (H, W) is the original image size.
+ (np.ndarray): An array of length C containing the model's
+ predictions for the quality of each mask.
+ (np.ndarray): An array of shape CxHxW, where C is the number
+ of masks and H=W=256. These low resolution logits can be passed to
+ a subsequent iteration as mask input.
+ """
+ if not self._is_image_set:
+ raise RuntimeError(
+ "An image must be set with .set_image(...) before mask prediction."
+ )
+
+ # Transform input prompts
+
+ mask_input, unnorm_coords, labels, unnorm_box = self._prep_prompts(
+ point_coords, point_labels, box, mask_input, normalize_coords
+ )
+
+ masks, iou_predictions, low_res_masks = self._predict(
+ unnorm_coords,
+ labels,
+ unnorm_box,
+ mask_input,
+ multimask_output,
+ return_logits=return_logits,
+ )
+
+ masks_np = masks.squeeze(0).float().detach().cpu().numpy()
+ iou_predictions_np = iou_predictions.squeeze(0).float().detach().cpu().numpy()
+ low_res_masks_np = low_res_masks.squeeze(0).float().detach().cpu().numpy()
+ return masks_np, iou_predictions_np, low_res_masks_np
+
+ def _prep_prompts(
+ self, point_coords, point_labels, box, mask_logits, normalize_coords, img_idx=-1
+ ):
+
+ unnorm_coords, labels, unnorm_box, mask_input = None, None, None, None
+ if point_coords is not None:
+ assert (
+ point_labels is not None
+ ), "point_labels must be supplied if point_coords is supplied."
+ point_coords = torch.as_tensor(
+ point_coords, dtype=torch.float, device=self.device
+ )
+ unnorm_coords = self._transforms.transform_coords(
+ point_coords, normalize=normalize_coords, orig_hw=self._orig_hw[img_idx]
+ )
+ labels = torch.as_tensor(point_labels, dtype=torch.int, device=self.device)
+ if len(unnorm_coords.shape) == 2:
+ unnorm_coords, labels = unnorm_coords[None, ...], labels[None, ...]
+ if box is not None:
+ box = torch.as_tensor(box, dtype=torch.float, device=self.device)
+ unnorm_box = self._transforms.transform_boxes(
+ box, normalize=normalize_coords, orig_hw=self._orig_hw[img_idx]
+ ) # Bx2x2
+ if mask_logits is not None:
+ mask_input = torch.as_tensor(
+ mask_logits, dtype=torch.float, device=self.device
+ )
+ if len(mask_input.shape) == 3:
+ mask_input = mask_input[None, :, :, :]
+ return mask_input, unnorm_coords, labels, unnorm_box
+
+ @torch.no_grad()
+ def _predict(
+ self,
+ point_coords: Optional[torch.Tensor],
+ point_labels: Optional[torch.Tensor],
+ boxes: Optional[torch.Tensor] = None,
+ mask_input: Optional[torch.Tensor] = None,
+ multimask_output: bool = True,
+ return_logits: bool = False,
+ img_idx: int = -1,
+ ) -> Tuple[torch.Tensor, torch.Tensor, torch.Tensor]:
+ """
+ Predict masks for the given input prompts, using the currently set image.
+ Input prompts are batched torch tensors and are expected to already be
+ transformed to the input frame using SAM2Transforms.
+
+ Arguments:
+ point_coords (torch.Tensor or None): A BxNx2 array of point prompts to the
+ model. Each point is in (X,Y) in pixels.
+ point_labels (torch.Tensor or None): A BxN array of labels for the
+ point prompts. 1 indicates a foreground point and 0 indicates a
+ background point.
+ boxes (np.ndarray or None): A Bx4 array given a box prompt to the
+ model, in XYXY format.
+ mask_input (np.ndarray): A low resolution mask input to the model, typically
+ coming from a previous prediction iteration. Has form Bx1xHxW, where
+ for SAM, H=W=256. Masks returned by a previous iteration of the
+ predict method do not need further transformation.
+ multimask_output (bool): If true, the model will return three masks.
+ For ambiguous input prompts (such as a single click), this will often
+ produce better masks than a single prediction. If only a single
+ mask is needed, the model's predicted quality score can be used
+ to select the best mask. For non-ambiguous prompts, such as multiple
+ input prompts, multimask_output=False can give better results.
+ return_logits (bool): If true, returns un-thresholded masks logits
+ instead of a binary mask.
+
+ Returns:
+ (torch.Tensor): The output masks in BxCxHxW format, where C is the
+ number of masks, and (H, W) is the original image size.
+ (torch.Tensor): An array of shape BxC containing the model's
+ predictions for the quality of each mask.
+ (torch.Tensor): An array of shape BxCxHxW, where C is the number
+ of masks and H=W=256. These low res logits can be passed to
+ a subsequent iteration as mask input.
+ """
+ if not self._is_image_set:
+ raise RuntimeError(
+ "An image must be set with .set_image(...) before mask prediction."
+ )
+
+ if point_coords is not None:
+ concat_points = (point_coords, point_labels)
+ else:
+ concat_points = None
+
+ # Embed prompts
+ if boxes is not None:
+ box_coords = boxes.reshape(-1, 2, 2)
+ box_labels = torch.tensor([[2, 3]], dtype=torch.int, device=boxes.device)
+ box_labels = box_labels.repeat(boxes.size(0), 1)
+ # we merge "boxes" and "points" into a single "concat_points" input (where
+ # boxes are added at the beginning) to sam_prompt_encoder
+ if concat_points is not None:
+ concat_coords = torch.cat([box_coords, concat_points[0]], dim=1)
+ concat_labels = torch.cat([box_labels, concat_points[1]], dim=1)
+ concat_points = (concat_coords, concat_labels)
+ else:
+ concat_points = (box_coords, box_labels)
+
+ sparse_embeddings, dense_embeddings = self.model.sam_prompt_encoder(
+ points=concat_points,
+ boxes=None,
+ masks=mask_input,
+ )
+
+ # Predict masks
+ batched_mode = (
+ concat_points is not None and concat_points[0].shape[0] > 1
+ ) # multi object prediction
+ high_res_features = [
+ feat_level[img_idx].unsqueeze(0)
+ for feat_level in self._features["high_res_feats"]
+ ]
+ low_res_masks, iou_predictions, _, _ = self.model.sam_mask_decoder(
+ image_embeddings=self._features["image_embed"][img_idx].unsqueeze(0),
+ image_pe=self.model.sam_prompt_encoder.get_dense_pe(),
+ sparse_prompt_embeddings=sparse_embeddings,
+ dense_prompt_embeddings=dense_embeddings,
+ multimask_output=multimask_output,
+ repeat_image=batched_mode,
+ high_res_features=high_res_features,
+ )
+
+ # Upscale the masks to the original image resolution
+ masks = self._transforms.postprocess_masks(
+ low_res_masks, self._orig_hw[img_idx]
+ )
+ low_res_masks = torch.clamp(low_res_masks, -32.0, 32.0)
+ if not return_logits:
+ masks = masks > self.mask_threshold
+
+ return masks, iou_predictions, low_res_masks
+
+ def get_image_embedding(self) -> torch.Tensor:
+ """
+ Returns the image embeddings for the currently set image, with
+ shape 1xCxHxW, where C is the embedding dimension and (H,W) are
+ the embedding spatial dimension of SAM (typically C=256, H=W=64).
+ """
+ if not self._is_image_set:
+ raise RuntimeError(
+ "An image must be set with .set_image(...) to generate an embedding."
+ )
+ assert (
+ self._features is not None
+ ), "Features must exist if an image has been set."
+ return self._features["image_embed"]
+
+ @property
+ def device(self) -> torch.device:
+ return self.model.device
+
+ def reset_predictor(self) -> None:
+ """
+ Resets the image embeddings and other state variables.
+ """
+ self._is_image_set = False
+ self._features = None
+ self._orig_hw = None
+ self._is_batch = False
diff --git a/py/evf_sam/model/segment_anything_2/sam2/sam2_video_predictor.py b/py/evf_sam/model/segment_anything_2/sam2/sam2_video_predictor.py
new file mode 100644
index 0000000..be925cb
--- /dev/null
+++ b/py/evf_sam/model/segment_anything_2/sam2/sam2_video_predictor.py
@@ -0,0 +1,984 @@
+# Copyright (c) Meta Platforms, Inc. and affiliates.
+# All rights reserved.
+
+# This source code is licensed under the license found in the
+# LICENSE file in the root directory of this source tree.
+
+from collections import OrderedDict
+
+import torch
+
+from tqdm import tqdm
+
+from model.segment_anything_2.sam2.modeling.sam2_base import NO_OBJ_SCORE, SAM2Base
+from model.segment_anything_2.sam2.utils.misc import concat_points, fill_holes_in_mask_scores, load_video_frames
+
+
+class SAM2VideoPredictor(SAM2Base):
+ """The predictor class to handle user interactions and manage inference states."""
+
+ def __init__(
+ self,
+ fill_hole_area=0,
+ # whether to apply non-overlapping constraints on the output object masks
+ non_overlap_masks=False,
+ # whether to clear non-conditioning memory of the surrounding frames (which may contain outdated information) after adding correction clicks;
+ # note that this would only apply to *single-object tracking* unless `clear_non_cond_mem_for_multi_obj` is also set to True)
+ clear_non_cond_mem_around_input=False,
+ # whether to also clear non-conditioning memory of the surrounding frames (only effective when `clear_non_cond_mem_around_input` is True).
+ clear_non_cond_mem_for_multi_obj=False,
+ **kwargs,
+ ):
+ super().__init__(**kwargs)
+ self.fill_hole_area = fill_hole_area
+ self.non_overlap_masks = non_overlap_masks
+ self.clear_non_cond_mem_around_input = clear_non_cond_mem_around_input
+ self.clear_non_cond_mem_for_multi_obj = clear_non_cond_mem_for_multi_obj
+
+ @torch.inference_mode()
+ def init_state(
+ self,
+ video_path,
+ offload_video_to_cpu=False,
+ offload_state_to_cpu=False,
+ async_loading_frames=False,
+ ):
+ """Initialize a inference state."""
+ images, video_height, video_width = load_video_frames(
+ video_path=video_path,
+ image_size=self.image_size,
+ offload_video_to_cpu=offload_video_to_cpu,
+ async_loading_frames=async_loading_frames,
+ )
+ inference_state = {}
+ inference_state["images"] = images
+ inference_state["num_frames"] = len(images)
+ # whether to offload the video frames to CPU memory
+ # turning on this option saves the GPU memory with only a very small overhead
+ inference_state["offload_video_to_cpu"] = offload_video_to_cpu
+ # whether to offload the inference state to CPU memory
+ # turning on this option saves the GPU memory at the cost of a lower tracking fps
+ # (e.g. in a test case of 768x768 model, fps dropped from 27 to 24 when tracking one object
+ # and from 24 to 21 when tracking two objects)
+ inference_state["offload_state_to_cpu"] = offload_state_to_cpu
+ # the original video height and width, used for resizing final output scores
+ inference_state["video_height"] = video_height
+ inference_state["video_width"] = video_width
+ inference_state["device"] = torch.device("cuda")
+ if offload_state_to_cpu:
+ inference_state["storage_device"] = torch.device("cpu")
+ else:
+ inference_state["storage_device"] = torch.device("cuda")
+ # inputs on each frame
+ inference_state["point_inputs_per_obj"] = {}
+ inference_state["mask_inputs_per_obj"] = {}
+ # visual features on a small number of recently visited frames for quick interactions
+ inference_state["cached_features"] = {}
+ # values that don't change across frames (so we only need to hold one copy of them)
+ inference_state["constants"] = {}
+ # mapping between client-side object id and model-side object index
+ inference_state["obj_id_to_idx"] = OrderedDict()
+ inference_state["obj_idx_to_id"] = OrderedDict()
+ inference_state["obj_ids"] = []
+ # A storage to hold the model's tracking results and states on each frame
+ inference_state["output_dict"] = {
+ "cond_frame_outputs": {}, # dict containing {frame_idx: }
+ "non_cond_frame_outputs": {}, # dict containing {frame_idx: }
+ }
+ # Slice (view) of each object tracking results, sharing the same memory with "output_dict"
+ inference_state["output_dict_per_obj"] = {}
+ # A temporary storage to hold new outputs when user interact with a frame
+ # to add clicks or mask (it's merged into "output_dict" before propagation starts)
+ inference_state["temp_output_dict_per_obj"] = {}
+ # Frames that already holds consolidated outputs from click or mask inputs
+ # (we directly use their consolidated outputs during tracking)
+ inference_state["consolidated_frame_inds"] = {
+ "cond_frame_outputs": set(), # set containing frame indices
+ "non_cond_frame_outputs": set(), # set containing frame indices
+ }
+ # metadata for each tracking frame (e.g. which direction it's tracked)
+ inference_state["tracking_has_started"] = False
+ inference_state["frames_already_tracked"] = {}
+ # Warm up the visual backbone and cache the image feature on frame 0
+ self._get_image_feature(inference_state, frame_idx=0, batch_size=1)
+ return inference_state
+
+ def _obj_id_to_idx(self, inference_state, obj_id):
+ """Map client-side object id to model-side object index."""
+ obj_idx = inference_state["obj_id_to_idx"].get(obj_id, None)
+ if obj_idx is not None:
+ return obj_idx
+
+ # This is a new object id not sent to the server before. We only allow adding
+ # new objects *before* the tracking starts.
+ allow_new_object = not inference_state["tracking_has_started"]
+ if allow_new_object:
+ # get the next object slot
+ obj_idx = len(inference_state["obj_id_to_idx"])
+ inference_state["obj_id_to_idx"][obj_id] = obj_idx
+ inference_state["obj_idx_to_id"][obj_idx] = obj_id
+ inference_state["obj_ids"] = list(inference_state["obj_id_to_idx"])
+ # set up input and output structures for this object
+ inference_state["point_inputs_per_obj"][obj_idx] = {}
+ inference_state["mask_inputs_per_obj"][obj_idx] = {}
+ inference_state["output_dict_per_obj"][obj_idx] = {
+ "cond_frame_outputs": {}, # dict containing {frame_idx: }
+ "non_cond_frame_outputs": {}, # dict containing {frame_idx: }
+ }
+ inference_state["temp_output_dict_per_obj"][obj_idx] = {
+ "cond_frame_outputs": {}, # dict containing {frame_idx: }
+ "non_cond_frame_outputs": {}, # dict containing {frame_idx: }
+ }
+ return obj_idx
+ else:
+ raise RuntimeError(
+ f"Cannot add new object id {obj_id} after tracking starts. "
+ f"All existing object ids: {inference_state['obj_ids']}. "
+ f"Please call 'reset_state' to restart from scratch."
+ )
+
+ def _obj_idx_to_id(self, inference_state, obj_idx):
+ """Map model-side object index to client-side object id."""
+ return inference_state["obj_idx_to_id"][obj_idx]
+
+ def _get_obj_num(self, inference_state):
+ """Get the total number of unique object ids received so far in this session."""
+ return len(inference_state["obj_idx_to_id"])
+
+ @torch.inference_mode()
+ def add_new_points(
+ self,
+ inference_state,
+ frame_idx,
+ obj_id,
+ points,
+ labels,
+ clear_old_points=True,
+ normalize_coords=True,
+ ):
+ """Add new points to a frame."""
+ obj_idx = self._obj_id_to_idx(inference_state, obj_id)
+ point_inputs_per_frame = inference_state["point_inputs_per_obj"][obj_idx]
+ mask_inputs_per_frame = inference_state["mask_inputs_per_obj"][obj_idx]
+
+ if not isinstance(points, torch.Tensor):
+ points = torch.tensor(points, dtype=torch.float32)
+ if not isinstance(labels, torch.Tensor):
+ labels = torch.tensor(labels, dtype=torch.int32)
+ if points.dim() == 2:
+ points = points.unsqueeze(0) # add batch dimension
+ if labels.dim() == 1:
+ labels = labels.unsqueeze(0) # add batch dimension
+ if normalize_coords:
+ video_H = inference_state["video_height"]
+ video_W = inference_state["video_width"]
+ points = points / torch.tensor([video_W, video_H]).to(points.device)
+ # scale the (normalized) coordinates by the model's internal image size
+ points = points * self.image_size
+ points = points.to(inference_state["device"])
+ labels = labels.to(inference_state["device"])
+
+ if not clear_old_points:
+ point_inputs = point_inputs_per_frame.get(frame_idx, None)
+ else:
+ point_inputs = None
+ point_inputs = concat_points(point_inputs, points, labels)
+
+ point_inputs_per_frame[frame_idx] = point_inputs
+ mask_inputs_per_frame.pop(frame_idx, None)
+ # If this frame hasn't been tracked before, we treat it as an initial conditioning
+ # frame, meaning that the inputs points are to generate segments on this frame without
+ # using any memory from other frames, like in SAM. Otherwise (if it has been tracked),
+ # the input points will be used to correct the already tracked masks.
+ is_init_cond_frame = frame_idx not in inference_state["frames_already_tracked"]
+ # whether to track in reverse time order
+ if is_init_cond_frame:
+ reverse = False
+ else:
+ reverse = inference_state["frames_already_tracked"][frame_idx]["reverse"]
+ obj_output_dict = inference_state["output_dict_per_obj"][obj_idx]
+ obj_temp_output_dict = inference_state["temp_output_dict_per_obj"][obj_idx]
+ # Add a frame to conditioning output if it's an initial conditioning frame or
+ # if the model sees all frames receiving clicks/mask as conditioning frames.
+ is_cond = is_init_cond_frame or self.add_all_frames_to_correct_as_cond
+ storage_key = "cond_frame_outputs" if is_cond else "non_cond_frame_outputs"
+
+ # Get any previously predicted mask logits on this object and feed it along with
+ # the new clicks into the SAM mask decoder.
+ prev_sam_mask_logits = None
+ # lookup temporary output dict first, which contains the most recent output
+ # (if not found, then lookup conditioning and non-conditioning frame output)
+ prev_out = obj_temp_output_dict[storage_key].get(frame_idx)
+ if prev_out is None:
+ prev_out = obj_output_dict["cond_frame_outputs"].get(frame_idx)
+ if prev_out is None:
+ prev_out = obj_output_dict["non_cond_frame_outputs"].get(frame_idx)
+
+ if prev_out is not None and prev_out["pred_masks"] is not None:
+ prev_sam_mask_logits = prev_out["pred_masks"].cuda(non_blocking=True)
+ # Clamp the scale of prev_sam_mask_logits to avoid rare numerical issues.
+ prev_sam_mask_logits = torch.clamp(prev_sam_mask_logits, -32.0, 32.0)
+ current_out, _ = self._run_single_frame_inference(
+ inference_state=inference_state,
+ output_dict=obj_output_dict, # run on the slice of a single object
+ frame_idx=frame_idx,
+ batch_size=1, # run on the slice of a single object
+ is_init_cond_frame=is_init_cond_frame,
+ point_inputs=point_inputs,
+ mask_inputs=None,
+ reverse=reverse,
+ # Skip the memory encoder when adding clicks or mask. We execute the memory encoder
+ # at the beginning of `propagate_in_video` (after user finalize their clicks). This
+ # allows us to enforce non-overlapping constraints on all objects before encoding
+ # them into memory.
+ run_mem_encoder=False,
+ prev_sam_mask_logits=prev_sam_mask_logits,
+ )
+ # Add the output to the output dict (to be used as future memory)
+ obj_temp_output_dict[storage_key][frame_idx] = current_out
+
+ # Resize the output mask to the original video resolution
+ obj_ids = inference_state["obj_ids"]
+ consolidated_out = self._consolidate_temp_output_across_obj(
+ inference_state,
+ frame_idx,
+ is_cond=is_cond,
+ run_mem_encoder=False,
+ consolidate_at_video_res=True,
+ )
+ _, video_res_masks = self._get_orig_video_res_output(
+ inference_state, consolidated_out["pred_masks_video_res"]
+ )
+ return frame_idx, obj_ids, video_res_masks
+
+ @torch.inference_mode()
+ def add_new_mask(
+ self,
+ inference_state,
+ frame_idx,
+ obj_id,
+ mask,
+ ):
+ """Add new mask to a frame."""
+ obj_idx = self._obj_id_to_idx(inference_state, obj_id)
+ point_inputs_per_frame = inference_state["point_inputs_per_obj"][obj_idx]
+ mask_inputs_per_frame = inference_state["mask_inputs_per_obj"][obj_idx]
+
+ if not isinstance(mask, torch.Tensor):
+ mask = torch.tensor(mask, dtype=torch.bool)
+ assert mask.dim() == 2
+ mask_H, mask_W = mask.shape
+ mask_inputs_orig = mask[None, None] # add batch and channel dimension
+ mask_inputs_orig = mask_inputs_orig.float().to(inference_state["device"])
+
+ # resize the mask if it doesn't match the model's image size
+ if mask_H != self.image_size or mask_W != self.image_size:
+ mask_inputs = torch.nn.functional.interpolate(
+ mask_inputs_orig,
+ size=(self.image_size, self.image_size),
+ align_corners=False,
+ mode="bilinear",
+ antialias=True, # use antialias for downsampling
+ )
+ mask_inputs = (mask_inputs >= 0.5).float()
+ else:
+ mask_inputs = mask_inputs_orig
+
+ mask_inputs_per_frame[frame_idx] = mask_inputs
+ point_inputs_per_frame.pop(frame_idx, None)
+ # If this frame hasn't been tracked before, we treat it as an initial conditioning
+ # frame, meaning that the inputs points are to generate segments on this frame without
+ # using any memory from other frames, like in SAM. Otherwise (if it has been tracked),
+ # the input points will be used to correct the already tracked masks.
+ is_init_cond_frame = frame_idx not in inference_state["frames_already_tracked"]
+ # whether to track in reverse time order
+ if is_init_cond_frame:
+ reverse = False
+ else:
+ reverse = inference_state["frames_already_tracked"][frame_idx]["reverse"]
+ obj_output_dict = inference_state["output_dict_per_obj"][obj_idx]
+ obj_temp_output_dict = inference_state["temp_output_dict_per_obj"][obj_idx]
+ # Add a frame to conditioning output if it's an initial conditioning frame or
+ # if the model sees all frames receiving clicks/mask as conditioning frames.
+ is_cond = is_init_cond_frame or self.add_all_frames_to_correct_as_cond
+ storage_key = "cond_frame_outputs" if is_cond else "non_cond_frame_outputs"
+
+ current_out, _ = self._run_single_frame_inference(
+ inference_state=inference_state,
+ output_dict=obj_output_dict, # run on the slice of a single object
+ frame_idx=frame_idx,
+ batch_size=1, # run on the slice of a single object
+ is_init_cond_frame=is_init_cond_frame,
+ point_inputs=None,
+ mask_inputs=mask_inputs,
+ reverse=reverse,
+ # Skip the memory encoder when adding clicks or mask. We execute the memory encoder
+ # at the beginning of `propagate_in_video` (after user finalize their clicks). This
+ # allows us to enforce non-overlapping constraints on all objects before encoding
+ # them into memory.
+ run_mem_encoder=False,
+ )
+ # Add the output to the output dict (to be used as future memory)
+ obj_temp_output_dict[storage_key][frame_idx] = current_out
+
+ # Resize the output mask to the original video resolution
+ obj_ids = inference_state["obj_ids"]
+ consolidated_out = self._consolidate_temp_output_across_obj(
+ inference_state,
+ frame_idx,
+ is_cond=is_cond,
+ run_mem_encoder=False,
+ consolidate_at_video_res=True,
+ )
+ _, video_res_masks = self._get_orig_video_res_output(
+ inference_state, consolidated_out["pred_masks_video_res"]
+ )
+ return frame_idx, obj_ids, video_res_masks
+
+
+ @torch.inference_mode()
+ def add_new_text(
+ self,
+ inference_state,
+ frame_idx,
+ obj_id,
+ text,
+ clear_old_points=True,
+ normalize_coords=True,
+ ):
+ """Add new text to a frame."""
+ obj_idx = self._obj_id_to_idx(inference_state, obj_id)
+ point_inputs_per_frame = inference_state["point_inputs_per_obj"][obj_idx]
+ mask_inputs_per_frame = inference_state["mask_inputs_per_obj"][obj_idx]
+
+ mask_inputs_per_frame.pop(frame_idx, None)
+ # If this frame hasn't been tracked before, we treat it as an initial conditioning
+ # frame, meaning that the inputs points are to generate segments on this frame without
+ # using any memory from other frames, like in SAM. Otherwise (if it has been tracked),
+ # the input points will be used to correct the already tracked masks.
+ is_init_cond_frame = frame_idx not in inference_state["frames_already_tracked"]
+ # whether to track in reverse time order
+ if is_init_cond_frame:
+ reverse = False
+ else:
+ reverse = inference_state["frames_already_tracked"][frame_idx]["reverse"]
+ obj_output_dict = inference_state["output_dict_per_obj"][obj_idx]
+ obj_temp_output_dict = inference_state["temp_output_dict_per_obj"][obj_idx]
+ # Add a frame to conditioning output if it's an initial conditioning frame or
+ # if the model sees all frames receiving clicks/mask as conditioning frames.
+ is_cond = is_init_cond_frame or self.add_all_frames_to_correct_as_cond
+ storage_key = "cond_frame_outputs" if is_cond else "non_cond_frame_outputs"
+
+ # Get any previously predicted mask logits on this object and feed it along with
+ # the new clicks into the SAM mask decoder.
+ prev_sam_mask_logits = None
+ # lookup temporary output dict first, which contains the most recent output
+ # (if not found, then lookup conditioning and non-conditioning frame output)
+ prev_out = obj_temp_output_dict[storage_key].get(frame_idx)
+ if prev_out is None:
+ prev_out = obj_output_dict["cond_frame_outputs"].get(frame_idx)
+ if prev_out is None:
+ prev_out = obj_output_dict["non_cond_frame_outputs"].get(frame_idx)
+
+ if prev_out is not None and prev_out["pred_masks"] is not None:
+ prev_sam_mask_logits = prev_out["pred_masks"].cuda(non_blocking=True)
+ # Clamp the scale of prev_sam_mask_logits to avoid rare numerical issues.
+ prev_sam_mask_logits = torch.clamp(prev_sam_mask_logits, -32.0, 32.0)
+ current_out, _ = self._run_single_frame_inference(
+ inference_state=inference_state,
+ output_dict=obj_output_dict, # run on the slice of a single object
+ frame_idx=frame_idx,
+ batch_size=1, # run on the slice of a single object
+ is_init_cond_frame=is_init_cond_frame,
+ point_inputs=None,
+ mask_inputs=None,
+ reverse=reverse,
+ # Skip the memory encoder when adding clicks or mask. We execute the memory encoder
+ # at the beginning of `propagate_in_video` (after user finalize their clicks). This
+ # allows us to enforce non-overlapping constraints on all objects before encoding
+ # them into memory.
+ run_mem_encoder=False,
+ prev_sam_mask_logits=prev_sam_mask_logits,
+ text_inputs=text
+ )
+ # Add the output to the output dict (to be used as future memory)
+ obj_temp_output_dict[storage_key][frame_idx] = current_out
+
+ # Resize the output mask to the original video resolution
+ obj_ids = inference_state["obj_ids"]
+ consolidated_out = self._consolidate_temp_output_across_obj(
+ inference_state,
+ frame_idx,
+ is_cond=is_cond,
+ run_mem_encoder=False,
+ consolidate_at_video_res=True,
+ )
+ _, video_res_masks = self._get_orig_video_res_output(
+ inference_state, consolidated_out["pred_masks_video_res"]
+ )
+ return frame_idx, obj_ids, video_res_masks
+
+
+ def _get_orig_video_res_output(self, inference_state, any_res_masks):
+ """
+ Resize the object scores to the original video resolution (video_res_masks)
+ and apply non-overlapping constraints for final output.
+ """
+ device = inference_state["device"]
+ video_H = inference_state["video_height"]
+ video_W = inference_state["video_width"]
+ any_res_masks = any_res_masks.to(device, non_blocking=True)
+ if any_res_masks.shape[-2:] == (video_H, video_W):
+ video_res_masks = any_res_masks
+ else:
+ video_res_masks = torch.nn.functional.interpolate(
+ any_res_masks,
+ size=(video_H, video_W),
+ mode="bilinear",
+ align_corners=False,
+ )
+ if self.non_overlap_masks:
+ video_res_masks = self._apply_non_overlapping_constraints(video_res_masks)
+ return any_res_masks, video_res_masks
+
+ def _consolidate_temp_output_across_obj(
+ self,
+ inference_state,
+ frame_idx,
+ is_cond,
+ run_mem_encoder,
+ consolidate_at_video_res=False,
+ ):
+ """
+ Consolidate the per-object temporary outputs in `temp_output_dict_per_obj` on
+ a frame into a single output for all objects, including
+ 1) fill any missing objects either from `output_dict_per_obj` (if they exist in
+ `output_dict_per_obj` for this frame) or leave them as placeholder values
+ (if they don't exist in `output_dict_per_obj` for this frame);
+ 2) if specified, rerun memory encoder after apply non-overlapping constraints
+ on the object scores.
+ """
+ batch_size = self._get_obj_num(inference_state)
+ storage_key = "cond_frame_outputs" if is_cond else "non_cond_frame_outputs"
+ # Optionally, we allow consolidating the temporary outputs at the original
+ # video resolution (to provide a better editing experience for mask prompts).
+ if consolidate_at_video_res:
+ assert not run_mem_encoder, "memory encoder cannot run at video resolution"
+ consolidated_H = inference_state["video_height"]
+ consolidated_W = inference_state["video_width"]
+ consolidated_mask_key = "pred_masks_video_res"
+ else:
+ consolidated_H = consolidated_W = self.image_size // 4
+ consolidated_mask_key = "pred_masks"
+
+ # Initialize `consolidated_out`. Its "maskmem_features" and "maskmem_pos_enc"
+ # will be added when rerunning the memory encoder after applying non-overlapping
+ # constraints to object scores. Its "pred_masks" are prefilled with a large
+ # negative value (NO_OBJ_SCORE) to represent missing objects.
+ consolidated_out = {
+ "maskmem_features": None,
+ "maskmem_pos_enc": None,
+ consolidated_mask_key: torch.full(
+ size=(batch_size, 1, consolidated_H, consolidated_W),
+ fill_value=NO_OBJ_SCORE,
+ dtype=torch.float32,
+ device=inference_state["storage_device"],
+ ),
+ "obj_ptr": torch.full(
+ size=(batch_size, self.hidden_dim),
+ fill_value=NO_OBJ_SCORE,
+ dtype=torch.float32,
+ device=inference_state["device"],
+ ),
+ }
+ empty_mask_ptr = None
+ for obj_idx in range(batch_size):
+ obj_temp_output_dict = inference_state["temp_output_dict_per_obj"][obj_idx]
+ obj_output_dict = inference_state["output_dict_per_obj"][obj_idx]
+ out = obj_temp_output_dict[storage_key].get(frame_idx, None)
+ # If the object doesn't appear in "temp_output_dict_per_obj" on this frame,
+ # we fall back and look up its previous output in "output_dict_per_obj".
+ # We look up both "cond_frame_outputs" and "non_cond_frame_outputs" in
+ # "output_dict_per_obj" to find a previous output for this object.
+ if out is None:
+ out = obj_output_dict["cond_frame_outputs"].get(frame_idx, None)
+ if out is None:
+ out = obj_output_dict["non_cond_frame_outputs"].get(frame_idx, None)
+ # If the object doesn't appear in "output_dict_per_obj" either, we skip it
+ # and leave its mask scores to the default scores (i.e. the NO_OBJ_SCORE
+ # placeholder above) and set its object pointer to be a dummy pointer.
+ if out is None:
+ # Fill in dummy object pointers for those objects without any inputs or
+ # tracking outcomes on this frame (only do it under `run_mem_encoder=True`,
+ # i.e. when we need to build the memory for tracking).
+ if run_mem_encoder:
+ if empty_mask_ptr is None:
+ empty_mask_ptr = self._get_empty_mask_ptr(
+ inference_state, frame_idx
+ )
+ # fill object pointer with a dummy pointer (based on an empty mask)
+ consolidated_out["obj_ptr"][obj_idx : obj_idx + 1] = empty_mask_ptr
+ continue
+ # Add the temporary object output mask to consolidated output mask
+ obj_mask = out["pred_masks"]
+ consolidated_pred_masks = consolidated_out[consolidated_mask_key]
+ if obj_mask.shape[-2:] == consolidated_pred_masks.shape[-2:]:
+ consolidated_pred_masks[obj_idx : obj_idx + 1] = obj_mask
+ else:
+ # Resize first if temporary object mask has a different resolution
+ resized_obj_mask = torch.nn.functional.interpolate(
+ obj_mask,
+ size=consolidated_pred_masks.shape[-2:],
+ mode="bilinear",
+ align_corners=False,
+ )
+ consolidated_pred_masks[obj_idx : obj_idx + 1] = resized_obj_mask
+ consolidated_out["obj_ptr"][obj_idx : obj_idx + 1] = out["obj_ptr"]
+
+ # Optionally, apply non-overlapping constraints on the consolidated scores
+ # and rerun the memory encoder
+ if run_mem_encoder:
+ device = inference_state["device"]
+ high_res_masks = torch.nn.functional.interpolate(
+ consolidated_out["pred_masks"].to(device, non_blocking=True),
+ size=(self.image_size, self.image_size),
+ mode="bilinear",
+ align_corners=False,
+ )
+ if self.non_overlap_masks_for_mem_enc:
+ high_res_masks = self._apply_non_overlapping_constraints(high_res_masks)
+ maskmem_features, maskmem_pos_enc = self._run_memory_encoder(
+ inference_state=inference_state,
+ frame_idx=frame_idx,
+ batch_size=batch_size,
+ high_res_masks=high_res_masks,
+ is_mask_from_pts=True, # these frames are what the user interacted with
+ )
+ consolidated_out["maskmem_features"] = maskmem_features
+ consolidated_out["maskmem_pos_enc"] = maskmem_pos_enc
+
+ return consolidated_out
+
+ def _get_empty_mask_ptr(self, inference_state, frame_idx):
+ """Get a dummy object pointer based on an empty mask on the current frame."""
+ # A dummy (empty) mask with a single object
+ batch_size = 1
+ mask_inputs = torch.zeros(
+ (batch_size, 1, self.image_size, self.image_size),
+ dtype=torch.float32,
+ device=inference_state["device"],
+ )
+
+ # Retrieve correct image features
+ (
+ _,
+ _,
+ current_vision_feats,
+ current_vision_pos_embeds,
+ feat_sizes,
+ ) = self._get_image_feature(inference_state, frame_idx, batch_size)
+
+ # Feed the empty mask and image feature above to get a dummy object pointer
+ current_out = self.track_step(
+ frame_idx=frame_idx,
+ is_init_cond_frame=True,
+ current_vision_feats=current_vision_feats,
+ current_vision_pos_embeds=current_vision_pos_embeds,
+ feat_sizes=feat_sizes,
+ point_inputs=None,
+ mask_inputs=mask_inputs,
+ output_dict={},
+ num_frames=inference_state["num_frames"],
+ track_in_reverse=False,
+ run_mem_encoder=False,
+ prev_sam_mask_logits=None,
+ )
+ return current_out["obj_ptr"]
+
+ @torch.inference_mode()
+ def propagate_in_video_preflight(self, inference_state):
+ """Prepare inference_state and consolidate temporary outputs before tracking."""
+ # Tracking has started and we don't allow adding new objects until session is reset.
+ inference_state["tracking_has_started"] = True
+ batch_size = self._get_obj_num(inference_state)
+
+ # Consolidate per-object temporary outputs in "temp_output_dict_per_obj" and
+ # add them into "output_dict".
+ temp_output_dict_per_obj = inference_state["temp_output_dict_per_obj"]
+ output_dict = inference_state["output_dict"]
+ # "consolidated_frame_inds" contains indices of those frames where consolidated
+ # temporary outputs have been added (either in this call or any previous calls
+ # to `propagate_in_video_preflight`).
+ consolidated_frame_inds = inference_state["consolidated_frame_inds"]
+ for is_cond in [False, True]:
+ # Separately consolidate conditioning and non-conditioning temp outptus
+ storage_key = "cond_frame_outputs" if is_cond else "non_cond_frame_outputs"
+ # Find all the frames that contain temporary outputs for any objects
+ # (these should be the frames that have just received clicks for mask inputs
+ # via `add_new_points` or `add_new_mask`)
+ temp_frame_inds = set()
+ for obj_temp_output_dict in temp_output_dict_per_obj.values():
+ temp_frame_inds.update(obj_temp_output_dict[storage_key].keys())
+ consolidated_frame_inds[storage_key].update(temp_frame_inds)
+ # consolidate the temprary output across all objects on this frame
+ for frame_idx in temp_frame_inds:
+ consolidated_out = self._consolidate_temp_output_across_obj(
+ inference_state, frame_idx, is_cond=is_cond, run_mem_encoder=True
+ )
+ # merge them into "output_dict" and also create per-object slices
+ output_dict[storage_key][frame_idx] = consolidated_out
+ self._add_output_per_object(
+ inference_state, frame_idx, consolidated_out, storage_key
+ )
+ clear_non_cond_mem = self.clear_non_cond_mem_around_input and (
+ self.clear_non_cond_mem_for_multi_obj or batch_size <= 1
+ )
+ if clear_non_cond_mem:
+ # clear non-conditioning memory of the surrounding frames
+ self._clear_non_cond_mem_around_input(inference_state, frame_idx)
+
+ # clear temporary outputs in `temp_output_dict_per_obj`
+ for obj_temp_output_dict in temp_output_dict_per_obj.values():
+ obj_temp_output_dict[storage_key].clear()
+
+ # edge case: if an output is added to "cond_frame_outputs", we remove any prior
+ # output on the same frame in "non_cond_frame_outputs"
+ for frame_idx in output_dict["cond_frame_outputs"]:
+ output_dict["non_cond_frame_outputs"].pop(frame_idx, None)
+ for obj_output_dict in inference_state["output_dict_per_obj"].values():
+ for frame_idx in obj_output_dict["cond_frame_outputs"]:
+ obj_output_dict["non_cond_frame_outputs"].pop(frame_idx, None)
+ for frame_idx in consolidated_frame_inds["cond_frame_outputs"]:
+ assert frame_idx in output_dict["cond_frame_outputs"]
+ consolidated_frame_inds["non_cond_frame_outputs"].discard(frame_idx)
+
+ # Make sure that the frame indices in "consolidated_frame_inds" are exactly those frames
+ # with either points or mask inputs (which should be true under a correct workflow).
+ # all_consolidated_frame_inds = (
+ # consolidated_frame_inds["cond_frame_outputs"]
+ # | consolidated_frame_inds["non_cond_frame_outputs"]
+ # )
+ # input_frames_inds = set()
+ # for point_inputs_per_frame in inference_state["point_inputs_per_obj"].values():
+ # input_frames_inds.update(point_inputs_per_frame.keys())
+ # for mask_inputs_per_frame in inference_state["mask_inputs_per_obj"].values():
+ # input_frames_inds.update(mask_inputs_per_frame.keys())
+ # assert all_consolidated_frame_inds == input_frames_inds
+
+ @torch.inference_mode()
+ def propagate_in_video(
+ self,
+ inference_state,
+ start_frame_idx=None,
+ max_frame_num_to_track=None,
+ reverse=False,
+ ):
+ """Propagate the input points across frames to track in the entire video."""
+ self.propagate_in_video_preflight(inference_state)
+
+ output_dict = inference_state["output_dict"]
+ consolidated_frame_inds = inference_state["consolidated_frame_inds"]
+ obj_ids = inference_state["obj_ids"]
+ num_frames = inference_state["num_frames"]
+ batch_size = self._get_obj_num(inference_state)
+ if len(output_dict["cond_frame_outputs"]) == 0:
+ raise RuntimeError("No points are provided; please add points first")
+ clear_non_cond_mem = self.clear_non_cond_mem_around_input and (
+ self.clear_non_cond_mem_for_multi_obj or batch_size <= 1
+ )
+
+ # set start index, end index, and processing order
+ if start_frame_idx is None:
+ # default: start from the earliest frame with input points
+ start_frame_idx = min(output_dict["cond_frame_outputs"])
+ if max_frame_num_to_track is None:
+ # default: track all the frames in the video
+ max_frame_num_to_track = num_frames
+ if reverse:
+ end_frame_idx = max(start_frame_idx - max_frame_num_to_track, 0)
+ if start_frame_idx > 0:
+ processing_order = range(start_frame_idx, end_frame_idx - 1, -1)
+ else:
+ processing_order = [] # skip reverse tracking if starting from frame 0
+ else:
+ end_frame_idx = min(
+ start_frame_idx + max_frame_num_to_track, num_frames - 1
+ )
+ processing_order = range(start_frame_idx, end_frame_idx + 1)
+
+ for frame_idx in tqdm(processing_order, desc="propagate in video"):
+ # We skip those frames already in consolidated outputs (these are frames
+ # that received input clicks or mask). Note that we cannot directly run
+ # batched forward on them via `_run_single_frame_inference` because the
+ # number of clicks on each object might be different.
+ if frame_idx in consolidated_frame_inds["cond_frame_outputs"]:
+ storage_key = "cond_frame_outputs"
+ current_out = output_dict[storage_key][frame_idx]
+ pred_masks = current_out["pred_masks"]
+ if clear_non_cond_mem:
+ # clear non-conditioning memory of the surrounding frames
+ self._clear_non_cond_mem_around_input(inference_state, frame_idx)
+ elif frame_idx in consolidated_frame_inds["non_cond_frame_outputs"]:
+ storage_key = "non_cond_frame_outputs"
+ current_out = output_dict[storage_key][frame_idx]
+ pred_masks = current_out["pred_masks"]
+ else:
+ storage_key = "non_cond_frame_outputs"
+ current_out, pred_masks = self._run_single_frame_inference(
+ inference_state=inference_state,
+ output_dict=output_dict,
+ frame_idx=frame_idx,
+ batch_size=batch_size,
+ is_init_cond_frame=False,
+ point_inputs=None,
+ mask_inputs=None,
+ reverse=reverse,
+ run_mem_encoder=True,
+ )
+ output_dict[storage_key][frame_idx] = current_out
+ # Create slices of per-object outputs for subsequent interaction with each
+ # individual object after tracking.
+ self._add_output_per_object(
+ inference_state, frame_idx, current_out, storage_key
+ )
+ inference_state["frames_already_tracked"][frame_idx] = {"reverse": reverse}
+
+ # Resize the output mask to the original video resolution (we directly use
+ # the mask scores on GPU for output to avoid any CPU conversion in between)
+ _, video_res_masks = self._get_orig_video_res_output(
+ inference_state, pred_masks
+ )
+ yield frame_idx, obj_ids, video_res_masks
+
+ def _add_output_per_object(
+ self, inference_state, frame_idx, current_out, storage_key
+ ):
+ """
+ Split a multi-object output into per-object output slices and add them into
+ `output_dict_per_obj`. The resulting slices share the same tensor storage.
+ """
+ maskmem_features = current_out["maskmem_features"]
+ assert maskmem_features is None or isinstance(maskmem_features, torch.Tensor)
+
+ maskmem_pos_enc = current_out["maskmem_pos_enc"]
+ assert maskmem_pos_enc is None or isinstance(maskmem_pos_enc, list)
+
+ output_dict_per_obj = inference_state["output_dict_per_obj"]
+ for obj_idx, obj_output_dict in output_dict_per_obj.items():
+ obj_slice = slice(obj_idx, obj_idx + 1)
+ obj_out = {
+ "maskmem_features": None,
+ "maskmem_pos_enc": None,
+ "pred_masks": current_out["pred_masks"][obj_slice],
+ "obj_ptr": current_out["obj_ptr"][obj_slice],
+ }
+ if maskmem_features is not None:
+ obj_out["maskmem_features"] = maskmem_features[obj_slice]
+ if maskmem_pos_enc is not None:
+ obj_out["maskmem_pos_enc"] = [x[obj_slice] for x in maskmem_pos_enc]
+ obj_output_dict[storage_key][frame_idx] = obj_out
+
+ @torch.inference_mode()
+ def reset_state(self, inference_state):
+ """Remove all input points or mask in all frames throughout the video."""
+ self._reset_tracking_results(inference_state)
+ # Remove all object ids
+ inference_state["obj_id_to_idx"].clear()
+ inference_state["obj_idx_to_id"].clear()
+ inference_state["obj_ids"].clear()
+ inference_state["point_inputs_per_obj"].clear()
+ inference_state["mask_inputs_per_obj"].clear()
+ inference_state["output_dict_per_obj"].clear()
+ inference_state["temp_output_dict_per_obj"].clear()
+
+ def _reset_tracking_results(self, inference_state):
+ """Reset all tracking inputs and results across the videos."""
+ for v in inference_state["point_inputs_per_obj"].values():
+ v.clear()
+ for v in inference_state["mask_inputs_per_obj"].values():
+ v.clear()
+ for v in inference_state["output_dict_per_obj"].values():
+ v["cond_frame_outputs"].clear()
+ v["non_cond_frame_outputs"].clear()
+ for v in inference_state["temp_output_dict_per_obj"].values():
+ v["cond_frame_outputs"].clear()
+ v["non_cond_frame_outputs"].clear()
+ inference_state["output_dict"]["cond_frame_outputs"].clear()
+ inference_state["output_dict"]["non_cond_frame_outputs"].clear()
+ inference_state["consolidated_frame_inds"]["cond_frame_outputs"].clear()
+ inference_state["consolidated_frame_inds"]["non_cond_frame_outputs"].clear()
+ inference_state["tracking_has_started"] = False
+ inference_state["frames_already_tracked"].clear()
+
+ def _get_image_feature(self, inference_state, frame_idx, batch_size):
+ """Compute the image features on a given frame."""
+ # Look up in the cache first
+ image, backbone_out = inference_state["cached_features"].get(
+ frame_idx, (None, None)
+ )
+ if backbone_out is None:
+ # Cache miss -- we will run inference on a single image
+ image = inference_state["images"][frame_idx].cuda().float().unsqueeze(0)
+ backbone_out = self.forward_image(image)
+ # Cache the most recent frame's feature (for repeated interactions with
+ # a frame; we can use an LRU cache for more frames in the future).
+ inference_state["cached_features"] = {frame_idx: (image, backbone_out)}
+
+ # expand the features to have the same dimension as the number of objects
+ expanded_image = image.expand(batch_size, -1, -1, -1)
+ expanded_backbone_out = {
+ "backbone_fpn": backbone_out["backbone_fpn"].copy(),
+ "vision_pos_enc": backbone_out["vision_pos_enc"].copy(),
+ }
+ for i, feat in enumerate(expanded_backbone_out["backbone_fpn"]):
+ expanded_backbone_out["backbone_fpn"][i] = feat.expand(
+ batch_size, -1, -1, -1
+ )
+ for i, pos in enumerate(expanded_backbone_out["vision_pos_enc"]):
+ pos = pos.expand(batch_size, -1, -1, -1)
+ expanded_backbone_out["vision_pos_enc"][i] = pos
+
+ features = self._prepare_backbone_features(expanded_backbone_out)
+ features = (expanded_image,) + features
+ return features
+
+ def _run_single_frame_inference(
+ self,
+ inference_state,
+ output_dict,
+ frame_idx,
+ batch_size,
+ is_init_cond_frame,
+ point_inputs,
+ mask_inputs,
+ reverse,
+ run_mem_encoder,
+ prev_sam_mask_logits=None,
+ text_inputs=None
+ ):
+ """Run tracking on a single frame based on current inputs and previous memory."""
+ # Retrieve correct image features
+ (
+ _,
+ _,
+ current_vision_feats,
+ current_vision_pos_embeds,
+ feat_sizes,
+ ) = self._get_image_feature(inference_state, frame_idx, batch_size)
+
+ # point and mask should not appear as input simultaneously on the same frame
+ assert point_inputs is None or mask_inputs is None
+ current_out = self.track_step(
+ frame_idx=frame_idx,
+ is_init_cond_frame=is_init_cond_frame,
+ current_vision_feats=current_vision_feats,
+ current_vision_pos_embeds=current_vision_pos_embeds,
+ feat_sizes=feat_sizes,
+ point_inputs=point_inputs,
+ mask_inputs=mask_inputs,
+ output_dict=output_dict,
+ num_frames=inference_state["num_frames"],
+ track_in_reverse=reverse,
+ run_mem_encoder=run_mem_encoder,
+ prev_sam_mask_logits=prev_sam_mask_logits,
+ text_inputs=text_inputs
+ )
+
+ # optionally offload the output to CPU memory to save GPU space
+ storage_device = inference_state["storage_device"]
+ maskmem_features = current_out["maskmem_features"]
+ if maskmem_features is not None:
+ maskmem_features = maskmem_features.to(torch.bfloat16)
+ maskmem_features = maskmem_features.to(storage_device, non_blocking=True)
+ pred_masks_gpu = current_out["pred_masks"]
+ # potentially fill holes in the predicted masks
+ if self.fill_hole_area > 0:
+ pred_masks_gpu = fill_holes_in_mask_scores(
+ pred_masks_gpu, self.fill_hole_area
+ )
+ pred_masks = pred_masks_gpu.to(storage_device, non_blocking=True)
+ # "maskmem_pos_enc" is the same across frames, so we only need to store one copy of it
+ maskmem_pos_enc = self._get_maskmem_pos_enc(inference_state, current_out)
+ # object pointer is a small tensor, so we always keep it on GPU memory for fast access
+ obj_ptr = current_out["obj_ptr"]
+ # make a compact version of this frame's output to reduce the state size
+ compact_current_out = {
+ "maskmem_features": maskmem_features,
+ "maskmem_pos_enc": maskmem_pos_enc,
+ "pred_masks": pred_masks,
+ "obj_ptr": obj_ptr,
+ }
+ return compact_current_out, pred_masks_gpu
+
+ def _run_memory_encoder(
+ self, inference_state, frame_idx, batch_size, high_res_masks, is_mask_from_pts
+ ):
+ """
+ Run the memory encoder on `high_res_masks`. This is usually after applying
+ non-overlapping constraints to object scores. Since their scores changed, their
+ memory also need to be computed again with the memory encoder.
+ """
+ # Retrieve correct image features
+ _, _, current_vision_feats, _, feat_sizes = self._get_image_feature(
+ inference_state, frame_idx, batch_size
+ )
+ maskmem_features, maskmem_pos_enc = self._encode_new_memory(
+ current_vision_feats=current_vision_feats,
+ feat_sizes=feat_sizes,
+ pred_masks_high_res=high_res_masks,
+ is_mask_from_pts=is_mask_from_pts,
+ )
+
+ # optionally offload the output to CPU memory to save GPU space
+ storage_device = inference_state["storage_device"]
+ maskmem_features = maskmem_features.to(torch.bfloat16)
+ maskmem_features = maskmem_features.to(storage_device, non_blocking=True)
+ # "maskmem_pos_enc" is the same across frames, so we only need to store one copy of it
+ maskmem_pos_enc = self._get_maskmem_pos_enc(
+ inference_state, {"maskmem_pos_enc": maskmem_pos_enc}
+ )
+ return maskmem_features, maskmem_pos_enc
+
+ def _get_maskmem_pos_enc(self, inference_state, current_out):
+ """
+ `maskmem_pos_enc` is the same across frames and objects, so we cache it as
+ a constant in the inference session to reduce session storage size.
+ """
+ model_constants = inference_state["constants"]
+ # "out_maskmem_pos_enc" should be either a list of tensors or None
+ out_maskmem_pos_enc = current_out["maskmem_pos_enc"]
+ if out_maskmem_pos_enc is not None:
+ if "maskmem_pos_enc" not in model_constants:
+ assert isinstance(out_maskmem_pos_enc, list)
+ # only take the slice for one object, since it's same across objects
+ maskmem_pos_enc = [x[0:1].clone() for x in out_maskmem_pos_enc]
+ model_constants["maskmem_pos_enc"] = maskmem_pos_enc
+ else:
+ maskmem_pos_enc = model_constants["maskmem_pos_enc"]
+ # expand the cached maskmem_pos_enc to the actual batch size
+ batch_size = out_maskmem_pos_enc[0].size(0)
+ expanded_maskmem_pos_enc = [
+ x.expand(batch_size, -1, -1, -1) for x in maskmem_pos_enc
+ ]
+ else:
+ expanded_maskmem_pos_enc = None
+ return expanded_maskmem_pos_enc
+
+ def _clear_non_cond_mem_around_input(self, inference_state, frame_idx):
+ """
+ Remove the non-conditioning memory around the input frame. When users provide
+ correction clicks, the surrounding frames' non-conditioning memories can still
+ contain outdated object appearance information and could confuse the model.
+
+ This method clears those non-conditioning memories surrounding the interacted
+ frame to avoid giving the model both old and new information about the object.
+ """
+ r = self.memory_temporal_stride_for_eval
+ frame_idx_begin = frame_idx - r * self.num_maskmem
+ frame_idx_end = frame_idx + r * self.num_maskmem
+ output_dict = inference_state["output_dict"]
+ non_cond_frame_outputs = output_dict["non_cond_frame_outputs"]
+ for t in range(frame_idx_begin, frame_idx_end + 1):
+ non_cond_frame_outputs.pop(t, None)
+ for obj_output_dict in inference_state["output_dict_per_obj"].values():
+ obj_output_dict["non_cond_frame_outputs"].pop(t, None)
diff --git a/py/evf_sam/model/segment_anything_2/sam2/utils/__init__.py b/py/evf_sam/model/segment_anything_2/sam2/utils/__init__.py
new file mode 100644
index 0000000..5277f46
--- /dev/null
+++ b/py/evf_sam/model/segment_anything_2/sam2/utils/__init__.py
@@ -0,0 +1,5 @@
+# Copyright (c) Meta Platforms, Inc. and affiliates.
+# All rights reserved.
+
+# This source code is licensed under the license found in the
+# LICENSE file in the root directory of this source tree.
diff --git a/py/evf_sam/model/segment_anything_2/sam2/utils/amg.py b/py/evf_sam/model/segment_anything_2/sam2/utils/amg.py
new file mode 100644
index 0000000..9868429
--- /dev/null
+++ b/py/evf_sam/model/segment_anything_2/sam2/utils/amg.py
@@ -0,0 +1,348 @@
+# Copyright (c) Meta Platforms, Inc. and affiliates.
+# All rights reserved.
+
+# This source code is licensed under the license found in the
+# LICENSE file in the root directory of this source tree.
+
+import math
+from copy import deepcopy
+from itertools import product
+from typing import Any, Dict, Generator, ItemsView, List, Tuple
+
+import numpy as np
+import torch
+
+# Very lightly adapted from https://github.com/facebookresearch/segment-anything/blob/main/segment_anything/utils/amg.py
+
+
+class MaskData:
+ """
+ A structure for storing masks and their related data in batched format.
+ Implements basic filtering and concatenation.
+ """
+
+ def __init__(self, **kwargs) -> None:
+ for v in kwargs.values():
+ assert isinstance(
+ v, (list, np.ndarray, torch.Tensor)
+ ), "MaskData only supports list, numpy arrays, and torch tensors."
+ self._stats = dict(**kwargs)
+
+ def __setitem__(self, key: str, item: Any) -> None:
+ assert isinstance(
+ item, (list, np.ndarray, torch.Tensor)
+ ), "MaskData only supports list, numpy arrays, and torch tensors."
+ self._stats[key] = item
+
+ def __delitem__(self, key: str) -> None:
+ del self._stats[key]
+
+ def __getitem__(self, key: str) -> Any:
+ return self._stats[key]
+
+ def items(self) -> ItemsView[str, Any]:
+ return self._stats.items()
+
+ def filter(self, keep: torch.Tensor) -> None:
+ for k, v in self._stats.items():
+ if v is None:
+ self._stats[k] = None
+ elif isinstance(v, torch.Tensor):
+ self._stats[k] = v[torch.as_tensor(keep, device=v.device)]
+ elif isinstance(v, np.ndarray):
+ self._stats[k] = v[keep.detach().cpu().numpy()]
+ elif isinstance(v, list) and keep.dtype == torch.bool:
+ self._stats[k] = [a for i, a in enumerate(v) if keep[i]]
+ elif isinstance(v, list):
+ self._stats[k] = [v[i] for i in keep]
+ else:
+ raise TypeError(f"MaskData key {k} has an unsupported type {type(v)}.")
+
+ def cat(self, new_stats: "MaskData") -> None:
+ for k, v in new_stats.items():
+ if k not in self._stats or self._stats[k] is None:
+ self._stats[k] = deepcopy(v)
+ elif isinstance(v, torch.Tensor):
+ self._stats[k] = torch.cat([self._stats[k], v], dim=0)
+ elif isinstance(v, np.ndarray):
+ self._stats[k] = np.concatenate([self._stats[k], v], axis=0)
+ elif isinstance(v, list):
+ self._stats[k] = self._stats[k] + deepcopy(v)
+ else:
+ raise TypeError(f"MaskData key {k} has an unsupported type {type(v)}.")
+
+ def to_numpy(self) -> None:
+ for k, v in self._stats.items():
+ if isinstance(v, torch.Tensor):
+ self._stats[k] = v.float().detach().cpu().numpy()
+
+
+def is_box_near_crop_edge(
+ boxes: torch.Tensor, crop_box: List[int], orig_box: List[int], atol: float = 20.0
+) -> torch.Tensor:
+ """Filter masks at the edge of a crop, but not at the edge of the original image."""
+ crop_box_torch = torch.as_tensor(crop_box, dtype=torch.float, device=boxes.device)
+ orig_box_torch = torch.as_tensor(orig_box, dtype=torch.float, device=boxes.device)
+ boxes = uncrop_boxes_xyxy(boxes, crop_box).float()
+ near_crop_edge = torch.isclose(boxes, crop_box_torch[None, :], atol=atol, rtol=0)
+ near_image_edge = torch.isclose(boxes, orig_box_torch[None, :], atol=atol, rtol=0)
+ near_crop_edge = torch.logical_and(near_crop_edge, ~near_image_edge)
+ return torch.any(near_crop_edge, dim=1)
+
+
+def box_xyxy_to_xywh(box_xyxy: torch.Tensor) -> torch.Tensor:
+ box_xywh = deepcopy(box_xyxy)
+ box_xywh[2] = box_xywh[2] - box_xywh[0]
+ box_xywh[3] = box_xywh[3] - box_xywh[1]
+ return box_xywh
+
+
+def batch_iterator(batch_size: int, *args) -> Generator[List[Any], None, None]:
+ assert len(args) > 0 and all(
+ len(a) == len(args[0]) for a in args
+ ), "Batched iteration must have inputs of all the same size."
+ n_batches = len(args[0]) // batch_size + int(len(args[0]) % batch_size != 0)
+ for b in range(n_batches):
+ yield [arg[b * batch_size : (b + 1) * batch_size] for arg in args]
+
+
+def mask_to_rle_pytorch(tensor: torch.Tensor) -> List[Dict[str, Any]]:
+ """
+ Encodes masks to an uncompressed RLE, in the format expected by
+ pycoco tools.
+ """
+ # Put in fortran order and flatten h,w
+ b, h, w = tensor.shape
+ tensor = tensor.permute(0, 2, 1).flatten(1)
+
+ # Compute change indices
+ diff = tensor[:, 1:] ^ tensor[:, :-1]
+ change_indices = diff.nonzero()
+
+ # Encode run length
+ out = []
+ for i in range(b):
+ cur_idxs = change_indices[change_indices[:, 0] == i, 1]
+ cur_idxs = torch.cat(
+ [
+ torch.tensor([0], dtype=cur_idxs.dtype, device=cur_idxs.device),
+ cur_idxs + 1,
+ torch.tensor([h * w], dtype=cur_idxs.dtype, device=cur_idxs.device),
+ ]
+ )
+ btw_idxs = cur_idxs[1:] - cur_idxs[:-1]
+ counts = [] if tensor[i, 0] == 0 else [0]
+ counts.extend(btw_idxs.detach().cpu().tolist())
+ out.append({"size": [h, w], "counts": counts})
+ return out
+
+
+def rle_to_mask(rle: Dict[str, Any]) -> np.ndarray:
+ """Compute a binary mask from an uncompressed RLE."""
+ h, w = rle["size"]
+ mask = np.empty(h * w, dtype=bool)
+ idx = 0
+ parity = False
+ for count in rle["counts"]:
+ mask[idx : idx + count] = parity
+ idx += count
+ parity ^= True
+ mask = mask.reshape(w, h)
+ return mask.transpose() # Put in C order
+
+
+def area_from_rle(rle: Dict[str, Any]) -> int:
+ return sum(rle["counts"][1::2])
+
+
+def calculate_stability_score(
+ masks: torch.Tensor, mask_threshold: float, threshold_offset: float
+) -> torch.Tensor:
+ """
+ Computes the stability score for a batch of masks. The stability
+ score is the IoU between the binary masks obtained by thresholding
+ the predicted mask logits at high and low values.
+ """
+ # One mask is always contained inside the other.
+ # Save memory by preventing unnecessary cast to torch.int64
+ intersections = (
+ (masks > (mask_threshold + threshold_offset))
+ .sum(-1, dtype=torch.int16)
+ .sum(-1, dtype=torch.int32)
+ )
+ unions = (
+ (masks > (mask_threshold - threshold_offset))
+ .sum(-1, dtype=torch.int16)
+ .sum(-1, dtype=torch.int32)
+ )
+ return intersections / unions
+
+
+def build_point_grid(n_per_side: int) -> np.ndarray:
+ """Generates a 2D grid of points evenly spaced in [0,1]x[0,1]."""
+ offset = 1 / (2 * n_per_side)
+ points_one_side = np.linspace(offset, 1 - offset, n_per_side)
+ points_x = np.tile(points_one_side[None, :], (n_per_side, 1))
+ points_y = np.tile(points_one_side[:, None], (1, n_per_side))
+ points = np.stack([points_x, points_y], axis=-1).reshape(-1, 2)
+ return points
+
+
+def build_all_layer_point_grids(
+ n_per_side: int, n_layers: int, scale_per_layer: int
+) -> List[np.ndarray]:
+ """Generates point grids for all crop layers."""
+ points_by_layer = []
+ for i in range(n_layers + 1):
+ n_points = int(n_per_side / (scale_per_layer**i))
+ points_by_layer.append(build_point_grid(n_points))
+ return points_by_layer
+
+
+def generate_crop_boxes(
+ im_size: Tuple[int, ...], n_layers: int, overlap_ratio: float
+) -> Tuple[List[List[int]], List[int]]:
+ """
+ Generates a list of crop boxes of different sizes. Each layer
+ has (2**i)**2 boxes for the ith layer.
+ """
+ crop_boxes, layer_idxs = [], []
+ im_h, im_w = im_size
+ short_side = min(im_h, im_w)
+
+ # Original image
+ crop_boxes.append([0, 0, im_w, im_h])
+ layer_idxs.append(0)
+
+ def crop_len(orig_len, n_crops, overlap):
+ return int(math.ceil((overlap * (n_crops - 1) + orig_len) / n_crops))
+
+ for i_layer in range(n_layers):
+ n_crops_per_side = 2 ** (i_layer + 1)
+ overlap = int(overlap_ratio * short_side * (2 / n_crops_per_side))
+
+ crop_w = crop_len(im_w, n_crops_per_side, overlap)
+ crop_h = crop_len(im_h, n_crops_per_side, overlap)
+
+ crop_box_x0 = [int((crop_w - overlap) * i) for i in range(n_crops_per_side)]
+ crop_box_y0 = [int((crop_h - overlap) * i) for i in range(n_crops_per_side)]
+
+ # Crops in XYWH format
+ for x0, y0 in product(crop_box_x0, crop_box_y0):
+ box = [x0, y0, min(x0 + crop_w, im_w), min(y0 + crop_h, im_h)]
+ crop_boxes.append(box)
+ layer_idxs.append(i_layer + 1)
+
+ return crop_boxes, layer_idxs
+
+
+def uncrop_boxes_xyxy(boxes: torch.Tensor, crop_box: List[int]) -> torch.Tensor:
+ x0, y0, _, _ = crop_box
+ offset = torch.tensor([[x0, y0, x0, y0]], device=boxes.device)
+ # Check if boxes has a channel dimension
+ if len(boxes.shape) == 3:
+ offset = offset.unsqueeze(1)
+ return boxes + offset
+
+
+def uncrop_points(points: torch.Tensor, crop_box: List[int]) -> torch.Tensor:
+ x0, y0, _, _ = crop_box
+ offset = torch.tensor([[x0, y0]], device=points.device)
+ # Check if points has a channel dimension
+ if len(points.shape) == 3:
+ offset = offset.unsqueeze(1)
+ return points + offset
+
+
+def uncrop_masks(
+ masks: torch.Tensor, crop_box: List[int], orig_h: int, orig_w: int
+) -> torch.Tensor:
+ x0, y0, x1, y1 = crop_box
+ if x0 == 0 and y0 == 0 and x1 == orig_w and y1 == orig_h:
+ return masks
+ # Coordinate transform masks
+ pad_x, pad_y = orig_w - (x1 - x0), orig_h - (y1 - y0)
+ pad = (x0, pad_x - x0, y0, pad_y - y0)
+ return torch.nn.functional.pad(masks, pad, value=0)
+
+
+def remove_small_regions(
+ mask: np.ndarray, area_thresh: float, mode: str
+) -> Tuple[np.ndarray, bool]:
+ """
+ Removes small disconnected regions and holes in a mask. Returns the
+ mask and an indicator of if the mask has been modified.
+ """
+ import cv2 # type: ignore
+
+ assert mode in ["holes", "islands"]
+ correct_holes = mode == "holes"
+ working_mask = (correct_holes ^ mask).astype(np.uint8)
+ n_labels, regions, stats, _ = cv2.connectedComponentsWithStats(working_mask, 8)
+ sizes = stats[:, -1][1:] # Row 0 is background label
+ small_regions = [i + 1 for i, s in enumerate(sizes) if s < area_thresh]
+ if len(small_regions) == 0:
+ return mask, False
+ fill_labels = [0] + small_regions
+ if not correct_holes:
+ fill_labels = [i for i in range(n_labels) if i not in fill_labels]
+ # If every region is below threshold, keep largest
+ if len(fill_labels) == 0:
+ fill_labels = [int(np.argmax(sizes)) + 1]
+ mask = np.isin(regions, fill_labels)
+ return mask, True
+
+
+def coco_encode_rle(uncompressed_rle: Dict[str, Any]) -> Dict[str, Any]:
+ from pycocotools import mask as mask_utils # type: ignore
+
+ h, w = uncompressed_rle["size"]
+ rle = mask_utils.frPyObjects(uncompressed_rle, h, w)
+ rle["counts"] = rle["counts"].decode("utf-8") # Necessary to serialize with json
+ return rle
+
+
+def batched_mask_to_box(masks: torch.Tensor) -> torch.Tensor:
+ """
+ Calculates boxes in XYXY format around masks. Return [0,0,0,0] for
+ an empty mask. For input shape C1xC2x...xHxW, the output shape is C1xC2x...x4.
+ """
+ # torch.max below raises an error on empty inputs, just skip in this case
+ if torch.numel(masks) == 0:
+ return torch.zeros(*masks.shape[:-2], 4, device=masks.device)
+
+ # Normalize shape to CxHxW
+ shape = masks.shape
+ h, w = shape[-2:]
+ if len(shape) > 2:
+ masks = masks.flatten(0, -3)
+ else:
+ masks = masks.unsqueeze(0)
+
+ # Get top and bottom edges
+ in_height, _ = torch.max(masks, dim=-1)
+ in_height_coords = in_height * torch.arange(h, device=in_height.device)[None, :]
+ bottom_edges, _ = torch.max(in_height_coords, dim=-1)
+ in_height_coords = in_height_coords + h * (~in_height)
+ top_edges, _ = torch.min(in_height_coords, dim=-1)
+
+ # Get left and right edges
+ in_width, _ = torch.max(masks, dim=-2)
+ in_width_coords = in_width * torch.arange(w, device=in_width.device)[None, :]
+ right_edges, _ = torch.max(in_width_coords, dim=-1)
+ in_width_coords = in_width_coords + w * (~in_width)
+ left_edges, _ = torch.min(in_width_coords, dim=-1)
+
+ # If the mask is empty the right edge will be to the left of the left edge.
+ # Replace these boxes with [0, 0, 0, 0]
+ empty_filter = (right_edges < left_edges) | (bottom_edges < top_edges)
+ out = torch.stack([left_edges, top_edges, right_edges, bottom_edges], dim=-1)
+ out = out * (~empty_filter).unsqueeze(-1)
+
+ # Return to original shape
+ if len(shape) > 2:
+ out = out.reshape(*shape[:-2], 4)
+ else:
+ out = out[0]
+
+ return out
diff --git a/py/evf_sam/model/segment_anything_2/sam2/utils/misc.py b/py/evf_sam/model/segment_anything_2/sam2/utils/misc.py
new file mode 100644
index 0000000..9adf6e4
--- /dev/null
+++ b/py/evf_sam/model/segment_anything_2/sam2/utils/misc.py
@@ -0,0 +1,238 @@
+# Copyright (c) Meta Platforms, Inc. and affiliates.
+# All rights reserved.
+
+# This source code is licensed under the license found in the
+# LICENSE file in the root directory of this source tree.
+
+import os
+import warnings
+from threading import Thread
+
+import numpy as np
+import torch
+from PIL import Image
+from tqdm import tqdm
+
+
+def get_sdpa_settings():
+ if torch.cuda.is_available():
+ old_gpu = torch.cuda.get_device_properties(0).major < 7
+ # only use Flash Attention on Ampere (8.0) or newer GPUs
+ use_flash_attn = torch.cuda.get_device_properties(0).major >= 8
+ if not use_flash_attn:
+ warnings.warn(
+ "Flash Attention is disabled as it requires a GPU with Ampere (8.0) CUDA capability.",
+ category=UserWarning,
+ stacklevel=2,
+ )
+ # keep math kernel for PyTorch versions before 2.2 (Flash Attention v2 is only
+ # available on PyTorch 2.2+, while Flash Attention v1 cannot handle all cases)
+ pytorch_version = tuple(int(v) for v in torch.__version__.split(".")[:2])
+ if pytorch_version < (2, 2):
+ warnings.warn(
+ f"You are using PyTorch {torch.__version__} without Flash Attention v2 support. "
+ "Consider upgrading to PyTorch 2.2+ for Flash Attention v2 (which could be faster).",
+ category=UserWarning,
+ stacklevel=2,
+ )
+ math_kernel_on = pytorch_version < (2, 2) or not use_flash_attn
+ else:
+ old_gpu = True
+ use_flash_attn = False
+ math_kernel_on = True
+
+ return old_gpu, use_flash_attn, math_kernel_on
+
+
+def get_connected_components(mask):
+ """
+ Get the connected components (8-connectivity) of binary masks of shape (N, 1, H, W).
+
+ Inputs:
+ - mask: A binary mask tensor of shape (N, 1, H, W), where 1 is foreground and 0 is
+ background.
+
+ Outputs:
+ - labels: A tensor of shape (N, 1, H, W) containing the connected component labels
+ for foreground pixels and 0 for background pixels.
+ - counts: A tensor of shape (N, 1, H, W) containing the area of the connected
+ components for foreground pixels and 0 for background pixels.
+ """
+ from model.segment_anything_2.sam2 import _C
+
+ return _C.get_connected_componnets(mask.to(torch.uint8).contiguous())
+
+
+def mask_to_box(masks: torch.Tensor):
+ """
+ compute bounding box given an input mask
+
+ Inputs:
+ - masks: [B, 1, H, W] boxes, dtype=torch.Tensor
+
+ Returns:
+ - box_coords: [B, 1, 4], contains (x, y) coordinates of top left and bottom right box corners, dtype=torch.Tensor
+ """
+ B, _, h, w = masks.shape
+ device = masks.device
+ xs = torch.arange(w, device=device, dtype=torch.int32)
+ ys = torch.arange(h, device=device, dtype=torch.int32)
+ grid_xs, grid_ys = torch.meshgrid(xs, ys, indexing="xy")
+ grid_xs = grid_xs[None, None, ...].expand(B, 1, h, w)
+ grid_ys = grid_ys[None, None, ...].expand(B, 1, h, w)
+ min_xs, _ = torch.min(torch.where(masks, grid_xs, w).flatten(-2), dim=-1)
+ max_xs, _ = torch.max(torch.where(masks, grid_xs, -1).flatten(-2), dim=-1)
+ min_ys, _ = torch.min(torch.where(masks, grid_ys, h).flatten(-2), dim=-1)
+ max_ys, _ = torch.max(torch.where(masks, grid_ys, -1).flatten(-2), dim=-1)
+ bbox_coords = torch.stack((min_xs, min_ys, max_xs, max_ys), dim=-1)
+
+ return bbox_coords
+
+
+def _load_img_as_tensor(img_path, image_size):
+ img_pil = Image.open(img_path)
+ img_np = np.array(img_pil.convert("RGB").resize((image_size, image_size)))
+ if img_np.dtype == np.uint8: # np.uint8 is expected for JPEG images
+ img_np = img_np / 255.0
+ else:
+ raise RuntimeError(f"Unknown image dtype: {img_np.dtype} on {img_path}")
+ img = torch.from_numpy(img_np).permute(2, 0, 1)
+ video_width, video_height = img_pil.size # the original video size
+ return img, video_height, video_width
+
+
+class AsyncVideoFrameLoader:
+ """
+ A list of video frames to be load asynchronously without blocking session start.
+ """
+
+ def __init__(self, img_paths, image_size, offload_video_to_cpu, img_mean, img_std):
+ self.img_paths = img_paths
+ self.image_size = image_size
+ self.offload_video_to_cpu = offload_video_to_cpu
+ self.img_mean = img_mean
+ self.img_std = img_std
+ # items in `self._images` will be loaded asynchronously
+ self.images = [None] * len(img_paths)
+ # catch and raise any exceptions in the async loading thread
+ self.exception = None
+ # video_height and video_width be filled when loading the first image
+ self.video_height = None
+ self.video_width = None
+
+ # load the first frame to fill video_height and video_width and also
+ # to cache it (since it's most likely where the user will click)
+ self.__getitem__(0)
+
+ # load the rest of frames asynchronously without blocking the session start
+ def _load_frames():
+ try:
+ for n in tqdm(range(len(self.images)), desc="frame loading (JPEG)"):
+ self.__getitem__(n)
+ except Exception as e:
+ self.exception = e
+
+ self.thread = Thread(target=_load_frames, daemon=True)
+ self.thread.start()
+
+ def __getitem__(self, index):
+ if self.exception is not None:
+ raise RuntimeError("Failure in frame loading thread") from self.exception
+
+ img = self.images[index]
+ if img is not None:
+ return img
+
+ img, video_height, video_width = _load_img_as_tensor(
+ self.img_paths[index], self.image_size
+ )
+ self.video_height = video_height
+ self.video_width = video_width
+ # normalize by mean and std
+ img -= self.img_mean
+ img /= self.img_std
+ if not self.offload_video_to_cpu:
+ img = img.cuda(non_blocking=True)
+ self.images[index] = img
+ return img
+
+ def __len__(self):
+ return len(self.images)
+
+
+def load_video_frames(
+ video_path,
+ image_size,
+ offload_video_to_cpu,
+ img_mean=(0.485, 0.456, 0.406),
+ img_std=(0.229, 0.224, 0.225),
+ async_loading_frames=False,
+):
+ """
+ Load the video frames from a directory of JPEG files (".jpg" format).
+
+ The frames are resized to image_size x image_size and are loaded to GPU if
+ `offload_video_to_cpu` is `False` and to CPU if `offload_video_to_cpu` is `True`.
+
+ You can load a frame asynchronously by setting `async_loading_frames` to `True`.
+ """
+ if isinstance(video_path, str) and os.path.isdir(video_path):
+ jpg_folder = video_path
+ else:
+ raise NotImplementedError("Only JPEG frames are supported at this moment")
+
+ frame_names = [
+ p
+ for p in os.listdir(jpg_folder)
+ if os.path.splitext(p)[-1] in [".jpg", ".jpeg", ".JPG", ".JPEG"]
+ ]
+ frame_names.sort(key=lambda p: int(os.path.splitext(p)[0]))
+ num_frames = len(frame_names)
+ if num_frames == 0:
+ raise RuntimeError(f"no images found in {jpg_folder}")
+ img_paths = [os.path.join(jpg_folder, frame_name) for frame_name in frame_names]
+ img_mean = torch.tensor(img_mean, dtype=torch.float32)[:, None, None]
+ img_std = torch.tensor(img_std, dtype=torch.float32)[:, None, None]
+
+ if async_loading_frames:
+ lazy_images = AsyncVideoFrameLoader(
+ img_paths, image_size, offload_video_to_cpu, img_mean, img_std
+ )
+ return lazy_images, lazy_images.video_height, lazy_images.video_width
+
+ images = torch.zeros(num_frames, 3, image_size, image_size, dtype=torch.float32)
+ for n, img_path in enumerate(tqdm(img_paths, desc="frame loading (JPEG)")):
+ images[n], video_height, video_width = _load_img_as_tensor(img_path, image_size)
+ if not offload_video_to_cpu:
+ images = images.cuda()
+ img_mean = img_mean.cuda()
+ img_std = img_std.cuda()
+ # normalize by mean and std
+ images -= img_mean
+ images /= img_std
+ return images, video_height, video_width
+
+
+def fill_holes_in_mask_scores(mask, max_area):
+ """
+ A post processor to fill small holes in mask scores with area under `max_area`.
+ """
+ # Holes are those connected components in background with area <= self.max_area
+ # (background regions are those with mask scores <= 0)
+ assert max_area > 0, "max_area must be positive"
+ labels, areas = get_connected_components(mask <= 0)
+ is_hole = (labels > 0) & (areas <= max_area)
+ # We fill holes with a small positive mask score (0.1) to change them to foreground.
+ mask = torch.where(is_hole, 0.1, mask)
+ return mask
+
+
+def concat_points(old_point_inputs, new_points, new_labels):
+ """Add new points and labels to previous point inputs (add at the end)."""
+ if old_point_inputs is None:
+ points, labels = new_points, new_labels
+ else:
+ points = torch.cat([old_point_inputs["point_coords"], new_points], dim=1)
+ labels = torch.cat([old_point_inputs["point_labels"], new_labels], dim=1)
+
+ return {"point_coords": points, "point_labels": labels}
diff --git a/py/evf_sam/model/segment_anything_2/sam2/utils/transforms.py b/py/evf_sam/model/segment_anything_2/sam2/utils/transforms.py
new file mode 100644
index 0000000..082a351
--- /dev/null
+++ b/py/evf_sam/model/segment_anything_2/sam2/utils/transforms.py
@@ -0,0 +1,99 @@
+# Copyright (c) Meta Platforms, Inc. and affiliates.
+# All rights reserved.
+
+# This source code is licensed under the license found in the
+# LICENSE file in the root directory of this source tree.
+
+import torch
+import torch.nn as nn
+import torch.nn.functional as F
+from torchvision.transforms import Normalize, Resize, ToTensor
+
+
+class SAM2Transforms(nn.Module):
+ def __init__(
+ self, resolution, mask_threshold, max_hole_area=0.0, max_sprinkle_area=0.0
+ ):
+ """
+ Transforms for SAM2.
+ """
+ super().__init__()
+ self.resolution = resolution
+ self.mask_threshold = mask_threshold
+ self.max_hole_area = max_hole_area
+ self.max_sprinkle_area = max_sprinkle_area
+ self.mean = [0.485, 0.456, 0.406]
+ self.std = [0.229, 0.224, 0.225]
+ self.to_tensor = ToTensor()
+ self.transforms = torch.jit.script(
+ nn.Sequential(
+ Resize((self.resolution, self.resolution)),
+ Normalize(self.mean, self.std),
+ )
+ )
+
+ def __call__(self, x):
+ x = self.to_tensor(x)
+ return self.transforms(x)
+
+ def forward_batch(self, img_list):
+ img_batch = [self.transforms(self.to_tensor(img)) for img in img_list]
+ img_batch = torch.stack(img_batch, dim=0)
+ return img_batch
+
+ def transform_coords(
+ self, coords: torch.Tensor, normalize=False, orig_hw=None
+ ) -> torch.Tensor:
+ """
+ Expects a torch tensor with length 2 in the last dimension. The coordinates can be in absolute image or normalized coordinates,
+ If the coords are in absolute image coordinates, normalize should be set to True and original image size is required.
+
+ Returns
+ Un-normalized coordinates in the range of [0, 1] which is expected by the SAM2 model.
+ """
+ if normalize:
+ assert orig_hw is not None
+ h, w = orig_hw
+ coords = coords.clone()
+ coords[..., 0] = coords[..., 0] / w
+ coords[..., 1] = coords[..., 1] / h
+
+ coords = coords * self.resolution # unnormalize coords
+ return coords
+
+ def transform_boxes(
+ self, boxes: torch.Tensor, normalize=False, orig_hw=None
+ ) -> torch.Tensor:
+ """
+ Expects a tensor of shape Bx4. The coordinates can be in absolute image or normalized coordinates,
+ if the coords are in absolute image coordinates, normalize should be set to True and original image size is required.
+ """
+ boxes = self.transform_coords(boxes.reshape(-1, 2, 2), normalize, orig_hw)
+ return boxes
+
+ def postprocess_masks(self, masks: torch.Tensor, orig_hw) -> torch.Tensor:
+ """
+ Perform PostProcessing on output masks.
+ """
+ from model.segment_anything_2.sam2.utils.misc import get_connected_components
+
+ masks = masks.float()
+ if self.max_hole_area > 0:
+ # Holes are those connected components in background with area <= self.fill_hole_area
+ # (background regions are those with mask scores <= self.mask_threshold)
+ mask_flat = masks.flatten(0, 1).unsqueeze(1) # flatten as 1-channel image
+ labels, areas = get_connected_components(mask_flat <= self.mask_threshold)
+ is_hole = (labels > 0) & (areas <= self.max_hole_area)
+ is_hole = is_hole.reshape_as(masks)
+ # We fill holes with a small positive mask score (10.0) to change them to foreground.
+ masks = torch.where(is_hole, self.mask_threshold + 10.0, masks)
+
+ if self.max_sprinkle_area > 0:
+ labels, areas = get_connected_components(mask_flat > self.mask_threshold)
+ is_hole = (labels > 0) & (areas <= self.max_sprinkle_area)
+ is_hole = is_hole.reshape_as(masks)
+ # We fill holes with negative mask score (-10.0) to change them to background.
+ masks = torch.where(is_hole, self.mask_threshold - 10.0, masks)
+
+ masks = F.interpolate(masks, orig_hw, mode="bilinear", align_corners=False)
+ return masks
diff --git a/py/evf_sam/model/segment_anything_2/sam2_configs/__init__.py b/py/evf_sam/model/segment_anything_2/sam2_configs/__init__.py
new file mode 100644
index 0000000..5277f46
--- /dev/null
+++ b/py/evf_sam/model/segment_anything_2/sam2_configs/__init__.py
@@ -0,0 +1,5 @@
+# Copyright (c) Meta Platforms, Inc. and affiliates.
+# All rights reserved.
+
+# This source code is licensed under the license found in the
+# LICENSE file in the root directory of this source tree.
diff --git a/py/evf_sam/model/segment_anything_2/sam2_configs/sam2_hiera_b+.yaml b/py/evf_sam/model/segment_anything_2/sam2_configs/sam2_hiera_b+.yaml
new file mode 100644
index 0000000..5ca7bfc
--- /dev/null
+++ b/py/evf_sam/model/segment_anything_2/sam2_configs/sam2_hiera_b+.yaml
@@ -0,0 +1,113 @@
+# @package _global_
+
+# Model
+model:
+ _target_: model.segment_anything_2.sam2.modeling.sam2_base.SAM2Base
+ image_encoder:
+ _target_: model.segment_anything_2.sam2.modeling.backbones.image_encoder.ImageEncoder
+ scalp: 1
+ trunk:
+ _target_: model.segment_anything_2.sam2.modeling.backbones.hieradet.Hiera
+ embed_dim: 112
+ num_heads: 2
+ neck:
+ _target_: model.segment_anything_2.sam2.modeling.backbones.image_encoder.FpnNeck
+ position_encoding:
+ _target_: model.segment_anything_2.sam2.modeling.position_encoding.PositionEmbeddingSine
+ num_pos_feats: 256
+ normalize: true
+ scale: null
+ temperature: 10000
+ d_model: 256
+ backbone_channel_list: [896, 448, 224, 112]
+ fpn_top_down_levels: [2, 3] # output level 0 and 1 directly use the backbone features
+ fpn_interp_model: nearest
+
+ memory_attention:
+ _target_: model.segment_anything_2.sam2.modeling.memory_attention.MemoryAttention
+ d_model: 256
+ pos_enc_at_input: true
+ layer:
+ _target_: model.segment_anything_2.sam2.modeling.memory_attention.MemoryAttentionLayer
+ activation: relu
+ dim_feedforward: 2048
+ dropout: 0.1
+ pos_enc_at_attn: false
+ self_attention:
+ _target_: model.segment_anything_2.sam2.modeling.sam.transformer.RoPEAttention
+ rope_theta: 10000.0
+ feat_sizes: [32, 32]
+ embedding_dim: 256
+ num_heads: 1
+ downsample_rate: 1
+ dropout: 0.1
+ d_model: 256
+ pos_enc_at_cross_attn_keys: true
+ pos_enc_at_cross_attn_queries: false
+ cross_attention:
+ _target_: model.segment_anything_2.sam2.modeling.sam.transformer.RoPEAttention
+ rope_theta: 10000.0
+ feat_sizes: [32, 32]
+ rope_k_repeat: True
+ embedding_dim: 256
+ num_heads: 1
+ downsample_rate: 1
+ dropout: 0.1
+ kv_in_dim: 64
+ num_layers: 4
+
+ memory_encoder:
+ _target_: model.segment_anything_2.sam2.modeling.memory_encoder.MemoryEncoder
+ out_dim: 64
+ position_encoding:
+ _target_: model.segment_anything_2.sam2.modeling.position_encoding.PositionEmbeddingSine
+ num_pos_feats: 64
+ normalize: true
+ scale: null
+ temperature: 10000
+ mask_downsampler:
+ _target_: model.segment_anything_2.sam2.modeling.memory_encoder.MaskDownSampler
+ kernel_size: 3
+ stride: 2
+ padding: 1
+ fuser:
+ _target_: model.segment_anything_2.sam2.modeling.memory_encoder.Fuser
+ layer:
+ _target_: model.segment_anything_2.sam2.modeling.memory_encoder.CXBlock
+ dim: 256
+ kernel_size: 7
+ padding: 3
+ layer_scale_init_value: 1e-6
+ use_dwconv: True # depth-wise convs
+ num_layers: 2
+
+ num_maskmem: 7
+ image_size: 1024
+ # apply scaled sigmoid on mask logits for memory encoder, and directly feed input mask as output mask
+ sigmoid_scale_for_mem_enc: 20.0
+ sigmoid_bias_for_mem_enc: -10.0
+ use_mask_input_as_output_without_sam: true
+ # Memory
+ directly_add_no_mem_embed: true
+ # use high-resolution feature map in the SAM mask decoder
+ use_high_res_features_in_sam: true
+ # output 3 masks on the first click on initial conditioning frames
+ multimask_output_in_sam: true
+ # SAM heads
+ iou_prediction_use_sigmoid: True
+ # cross-attend to object pointers from other frames (based on SAM output tokens) in the encoder
+ use_obj_ptrs_in_encoder: true
+ add_tpos_enc_to_obj_ptrs: false
+ only_obj_ptrs_in_the_past_for_eval: true
+ # object occlusion prediction
+ pred_obj_scores: true
+ pred_obj_scores_mlp: true
+ fixed_no_obj_ptr: true
+ # multimask tracking settings
+ multimask_output_for_tracking: true
+ use_multimask_token_for_obj_ptr: true
+ multimask_min_pt_num: 0
+ multimask_max_pt_num: 1
+ use_mlp_for_obj_ptr_proj: true
+ # Compilation flag
+ compile_image_encoder: False
diff --git a/py/evf_sam/model/segment_anything_2/sam2_configs/sam2_hiera_l.yaml b/py/evf_sam/model/segment_anything_2/sam2_configs/sam2_hiera_l.yaml
new file mode 100644
index 0000000..293d038
--- /dev/null
+++ b/py/evf_sam/model/segment_anything_2/sam2_configs/sam2_hiera_l.yaml
@@ -0,0 +1,117 @@
+# @package _global_
+
+# Model
+model:
+ _target_: model.segment_anything_2.sam2.modeling.sam2_base.SAM2Base
+ image_encoder:
+ _target_: model.segment_anything_2.sam2.modeling.backbones.image_encoder.ImageEncoder
+ scalp: 1
+ trunk:
+ _target_: model.segment_anything_2.sam2.modeling.backbones.hieradet.Hiera
+ embed_dim: 144
+ num_heads: 2
+ stages: [2, 6, 36, 4]
+ global_att_blocks: [23, 33, 43]
+ window_pos_embed_bkg_spatial_size: [7, 7]
+ window_spec: [8, 4, 16, 8]
+ neck:
+ _target_: model.segment_anything_2.sam2.modeling.backbones.image_encoder.FpnNeck
+ position_encoding:
+ _target_: model.segment_anything_2.sam2.modeling.position_encoding.PositionEmbeddingSine
+ num_pos_feats: 256
+ normalize: true
+ scale: null
+ temperature: 10000
+ d_model: 256
+ backbone_channel_list: [1152, 576, 288, 144]
+ fpn_top_down_levels: [2, 3] # output level 0 and 1 directly use the backbone features
+ fpn_interp_model: nearest
+
+ memory_attention:
+ _target_: model.segment_anything_2.sam2.modeling.memory_attention.MemoryAttention
+ d_model: 256
+ pos_enc_at_input: true
+ layer:
+ _target_: model.segment_anything_2.sam2.modeling.memory_attention.MemoryAttentionLayer
+ activation: relu
+ dim_feedforward: 2048
+ dropout: 0.1
+ pos_enc_at_attn: false
+ self_attention:
+ _target_: model.segment_anything_2.sam2.modeling.sam.transformer.RoPEAttention
+ rope_theta: 10000.0
+ feat_sizes: [32, 32]
+ embedding_dim: 256
+ num_heads: 1
+ downsample_rate: 1
+ dropout: 0.1
+ d_model: 256
+ pos_enc_at_cross_attn_keys: true
+ pos_enc_at_cross_attn_queries: false
+ cross_attention:
+ _target_: model.segment_anything_2.sam2.modeling.sam.transformer.RoPEAttention
+ rope_theta: 10000.0
+ feat_sizes: [32, 32]
+ rope_k_repeat: True
+ embedding_dim: 256
+ num_heads: 1
+ downsample_rate: 1
+ dropout: 0.1
+ kv_in_dim: 64
+ num_layers: 4
+
+ memory_encoder:
+ _target_: model.segment_anything_2.sam2.modeling.memory_encoder.MemoryEncoder
+ out_dim: 64
+ position_encoding:
+ _target_: model.segment_anything_2.sam2.modeling.position_encoding.PositionEmbeddingSine
+ num_pos_feats: 64
+ normalize: true
+ scale: null
+ temperature: 10000
+ mask_downsampler:
+ _target_: model.segment_anything_2.sam2.modeling.memory_encoder.MaskDownSampler
+ kernel_size: 3
+ stride: 2
+ padding: 1
+ fuser:
+ _target_: model.segment_anything_2.sam2.modeling.memory_encoder.Fuser
+ layer:
+ _target_: model.segment_anything_2.sam2.modeling.memory_encoder.CXBlock
+ dim: 256
+ kernel_size: 7
+ padding: 3
+ layer_scale_init_value: 1e-6
+ use_dwconv: True # depth-wise convs
+ num_layers: 2
+
+ num_maskmem: 7
+ image_size: 1024
+ # apply scaled sigmoid on mask logits for memory encoder, and directly feed input mask as output mask
+ sigmoid_scale_for_mem_enc: 20.0
+ sigmoid_bias_for_mem_enc: -10.0
+ use_mask_input_as_output_without_sam: true
+ # Memory
+ directly_add_no_mem_embed: true
+ # use high-resolution feature map in the SAM mask decoder
+ use_high_res_features_in_sam: true
+ # output 3 masks on the first click on initial conditioning frames
+ multimask_output_in_sam: true
+ # SAM heads
+ iou_prediction_use_sigmoid: True
+ # cross-attend to object pointers from other frames (based on SAM output tokens) in the encoder
+ use_obj_ptrs_in_encoder: true
+ add_tpos_enc_to_obj_ptrs: false
+ only_obj_ptrs_in_the_past_for_eval: true
+ # object occlusion prediction
+ pred_obj_scores: true
+ pred_obj_scores_mlp: true
+ fixed_no_obj_ptr: true
+ # multimask tracking settings
+ multimask_output_for_tracking: true
+ use_multimask_token_for_obj_ptr: true
+ multimask_min_pt_num: 0
+ multimask_max_pt_num: 1
+ use_mlp_for_obj_ptr_proj: true
+ # Compilation flag
+ compile_image_encoder: False
diff --git a/py/evf_sam/model/segment_anything_2/sam2_configs/sam2_hiera_s.yaml b/py/evf_sam/model/segment_anything_2/sam2_configs/sam2_hiera_s.yaml
new file mode 100644
index 0000000..8d4627b
--- /dev/null
+++ b/py/evf_sam/model/segment_anything_2/sam2_configs/sam2_hiera_s.yaml
@@ -0,0 +1,116 @@
+# @package _global_
+
+# Model
+model:
+ _target_: model.segment_anything_2.sam2.modeling.sam2_base.SAM2Base
+ image_encoder:
+ _target_: model.segment_anything_2.sam2.modeling.backbones.image_encoder.ImageEncoder
+ scalp: 1
+ trunk:
+ _target_: model.segment_anything_2.sam2.modeling.backbones.hieradet.Hiera
+ embed_dim: 96
+ num_heads: 1
+ stages: [1, 2, 11, 2]
+ global_att_blocks: [7, 10, 13]
+ window_pos_embed_bkg_spatial_size: [7, 7]
+ neck:
+ _target_: model.segment_anything_2.sam2.modeling.backbones.image_encoder.FpnNeck
+ position_encoding:
+ _target_: model.segment_anything_2.sam2.modeling.position_encoding.PositionEmbeddingSine
+ num_pos_feats: 256
+ normalize: true
+ scale: null
+ temperature: 10000
+ d_model: 256
+ backbone_channel_list: [768, 384, 192, 96]
+ fpn_top_down_levels: [2, 3] # output level 0 and 1 directly use the backbone features
+ fpn_interp_model: nearest
+
+ memory_attention:
+ _target_: model.segment_anything_2.sam2.modeling.memory_attention.MemoryAttention
+ d_model: 256
+ pos_enc_at_input: true
+ layer:
+ _target_: model.segment_anything_2.sam2.modeling.memory_attention.MemoryAttentionLayer
+ activation: relu
+ dim_feedforward: 2048
+ dropout: 0.1
+ pos_enc_at_attn: false
+ self_attention:
+ _target_: model.segment_anything_2.sam2.modeling.sam.transformer.RoPEAttention
+ rope_theta: 10000.0
+ feat_sizes: [32, 32]
+ embedding_dim: 256
+ num_heads: 1
+ downsample_rate: 1
+ dropout: 0.1
+ d_model: 256
+ pos_enc_at_cross_attn_keys: true
+ pos_enc_at_cross_attn_queries: false
+ cross_attention:
+ _target_: model.segment_anything_2.sam2.modeling.sam.transformer.RoPEAttention
+ rope_theta: 10000.0
+ feat_sizes: [32, 32]
+ rope_k_repeat: True
+ embedding_dim: 256
+ num_heads: 1
+ downsample_rate: 1
+ dropout: 0.1
+ kv_in_dim: 64
+ num_layers: 4
+
+ memory_encoder:
+ _target_: model.segment_anything_2.sam2.modeling.memory_encoder.MemoryEncoder
+ out_dim: 64
+ position_encoding:
+ _target_: model.segment_anything_2.sam2.modeling.position_encoding.PositionEmbeddingSine
+ num_pos_feats: 64
+ normalize: true
+ scale: null
+ temperature: 10000
+ mask_downsampler:
+ _target_: model.segment_anything_2.sam2.modeling.memory_encoder.MaskDownSampler
+ kernel_size: 3
+ stride: 2
+ padding: 1
+ fuser:
+ _target_: model.segment_anything_2.sam2.modeling.memory_encoder.Fuser
+ layer:
+ _target_: model.segment_anything_2.sam2.modeling.memory_encoder.CXBlock
+ dim: 256
+ kernel_size: 7
+ padding: 3
+ layer_scale_init_value: 1e-6
+ use_dwconv: True # depth-wise convs
+ num_layers: 2
+
+ num_maskmem: 7
+ image_size: 1024
+ # apply scaled sigmoid on mask logits for memory encoder, and directly feed input mask as output mask
+ sigmoid_scale_for_mem_enc: 20.0
+ sigmoid_bias_for_mem_enc: -10.0
+ use_mask_input_as_output_without_sam: true
+ # Memory
+ directly_add_no_mem_embed: true
+ # use high-resolution feature map in the SAM mask decoder
+ use_high_res_features_in_sam: true
+ # output 3 masks on the first click on initial conditioning frames
+ multimask_output_in_sam: true
+ # SAM heads
+ iou_prediction_use_sigmoid: True
+ # cross-attend to object pointers from other frames (based on SAM output tokens) in the encoder
+ use_obj_ptrs_in_encoder: true
+ add_tpos_enc_to_obj_ptrs: false
+ only_obj_ptrs_in_the_past_for_eval: true
+ # object occlusion prediction
+ pred_obj_scores: true
+ pred_obj_scores_mlp: true
+ fixed_no_obj_ptr: true
+ # multimask tracking settings
+ multimask_output_for_tracking: true
+ use_multimask_token_for_obj_ptr: true
+ multimask_min_pt_num: 0
+ multimask_max_pt_num: 1
+ use_mlp_for_obj_ptr_proj: true
+ # Compilation flag
+ compile_image_encoder: False
diff --git a/py/evf_sam/model/segment_anything_2/sam2_configs/sam2_hiera_t.yaml b/py/evf_sam/model/segment_anything_2/sam2_configs/sam2_hiera_t.yaml
new file mode 100644
index 0000000..2664b7c
--- /dev/null
+++ b/py/evf_sam/model/segment_anything_2/sam2_configs/sam2_hiera_t.yaml
@@ -0,0 +1,118 @@
+# @package _global_
+
+# Model
+model:
+ _target_: model.segment_anything_2.sam2.modeling.sam2_base.SAM2Base
+ image_encoder:
+ _target_: model.segment_anything_2.sam2.modeling.backbones.image_encoder.ImageEncoder
+ scalp: 1
+ trunk:
+ _target_: model.segment_anything_2.sam2.modeling.backbones.hieradet.Hiera
+ embed_dim: 96
+ num_heads: 1
+ stages: [1, 2, 7, 2]
+ global_att_blocks: [5, 7, 9]
+ window_pos_embed_bkg_spatial_size: [7, 7]
+ neck:
+ _target_: model.segment_anything_2.sam2.modeling.backbones.image_encoder.FpnNeck
+ position_encoding:
+ _target_: model.segment_anything_2.sam2.modeling.position_encoding.PositionEmbeddingSine
+ num_pos_feats: 256
+ normalize: true
+ scale: null
+ temperature: 10000
+ d_model: 256
+ backbone_channel_list: [768, 384, 192, 96]
+ fpn_top_down_levels: [2, 3] # output level 0 and 1 directly use the backbone features
+ fpn_interp_model: nearest
+
+ memory_attention:
+ _target_: model.segment_anything_2.sam2.modeling.memory_attention.MemoryAttention
+ d_model: 256
+ pos_enc_at_input: true
+ layer:
+ _target_: model.segment_anything_2.sam2.modeling.memory_attention.MemoryAttentionLayer
+ activation: relu
+ dim_feedforward: 2048
+ dropout: 0.1
+ pos_enc_at_attn: false
+ self_attention:
+ _target_: model.segment_anything_2.sam2.modeling.sam.transformer.RoPEAttention
+ rope_theta: 10000.0
+ feat_sizes: [32, 32]
+ embedding_dim: 256
+ num_heads: 1
+ downsample_rate: 1
+ dropout: 0.1
+ d_model: 256
+ pos_enc_at_cross_attn_keys: true
+ pos_enc_at_cross_attn_queries: false
+ cross_attention:
+ _target_: model.segment_anything_2.sam2.modeling.sam.transformer.RoPEAttention
+ rope_theta: 10000.0
+ feat_sizes: [32, 32]
+ rope_k_repeat: True
+ embedding_dim: 256
+ num_heads: 1
+ downsample_rate: 1
+ dropout: 0.1
+ kv_in_dim: 64
+ num_layers: 4
+
+ memory_encoder:
+ _target_: model.segment_anything_2.sam2.modeling.memory_encoder.MemoryEncoder
+ out_dim: 64
+ position_encoding:
+ _target_: model.segment_anything_2.sam2.modeling.position_encoding.PositionEmbeddingSine
+ num_pos_feats: 64
+ normalize: true
+ scale: null
+ temperature: 10000
+ mask_downsampler:
+ _target_: model.segment_anything_2.sam2.modeling.memory_encoder.MaskDownSampler
+ kernel_size: 3
+ stride: 2
+ padding: 1
+ fuser:
+ _target_: model.segment_anything_2.sam2.modeling.memory_encoder.Fuser
+ layer:
+ _target_: model.segment_anything_2.sam2.modeling.memory_encoder.CXBlock
+ dim: 256
+ kernel_size: 7
+ padding: 3
+ layer_scale_init_value: 1e-6
+ use_dwconv: True # depth-wise convs
+ num_layers: 2
+
+ num_maskmem: 7
+ image_size: 1024
+ # apply scaled sigmoid on mask logits for memory encoder, and directly feed input mask as output mask
+ # SAM decoder
+ sigmoid_scale_for_mem_enc: 20.0
+ sigmoid_bias_for_mem_enc: -10.0
+ use_mask_input_as_output_without_sam: true
+ # Memory
+ directly_add_no_mem_embed: true
+ # use high-resolution feature map in the SAM mask decoder
+ use_high_res_features_in_sam: true
+ # output 3 masks on the first click on initial conditioning frames
+ multimask_output_in_sam: true
+ # SAM heads
+ iou_prediction_use_sigmoid: True
+ # cross-attend to object pointers from other frames (based on SAM output tokens) in the encoder
+ use_obj_ptrs_in_encoder: true
+ add_tpos_enc_to_obj_ptrs: false
+ only_obj_ptrs_in_the_past_for_eval: true
+ # object occlusion prediction
+ pred_obj_scores: true
+ pred_obj_scores_mlp: true
+ fixed_no_obj_ptr: true
+ # multimask tracking settings
+ multimask_output_for_tracking: true
+ use_multimask_token_for_obj_ptr: true
+ multimask_min_pt_num: 0
+ multimask_max_pt_num: 1
+ use_mlp_for_obj_ptr_proj: true
+ # Compilation flag
+ # HieraT does not currently support compilation, should always be set to False
+ compile_image_encoder: False
diff --git a/py/evf_sam/model/segment_anything_2/setup.py b/py/evf_sam/model/segment_anything_2/setup.py
new file mode 100644
index 0000000..a228990
--- /dev/null
+++ b/py/evf_sam/model/segment_anything_2/setup.py
@@ -0,0 +1,29 @@
+# Copyright (c) Meta Platforms, Inc. and affiliates.
+# All rights reserved.
+
+# This source code is licensed under the license found in the
+# LICENSE file in the root directory of this source tree.
+
+from setuptools import find_packages, setup
+from torch.utils.cpp_extension import BuildExtension, CUDAExtension
+
+def get_extensions():
+ srcs = ["sam2/csrc/connected_components.cu"]
+ compile_args = {
+ "cxx": [],
+ "nvcc": [
+ "-DCUDA_HAS_FP16=1",
+ "-D__CUDA_NO_HALF_OPERATORS__",
+ "-D__CUDA_NO_HALF_CONVERSIONS__",
+ "-D__CUDA_NO_HALF2_OPERATORS__",
+ ],
+ }
+ ext_modules = [CUDAExtension("sam2._C", srcs, extra_compile_args=compile_args)]
+ return ext_modules
+
+
+# Setup configuration
+setup(
+ ext_modules=get_extensions(),
+ cmdclass={"build_ext": BuildExtension.with_options(no_python_abi_suffix=True)},
+)
diff --git a/py/evf_sam/model/unilm/beit3/README.md b/py/evf_sam/model/unilm/beit3/README.md
new file mode 100644
index 0000000..39c5b2e
--- /dev/null
+++ b/py/evf_sam/model/unilm/beit3/README.md
@@ -0,0 +1,191 @@
+# [(BEiT-3) Image as a Foreign Language: BEiT Pretraining for Vision and Vision-Language Tasks](https://arxiv.org/abs/2208.10442)
+
+Official PyTorch implementation and pretrained models of BEiT-3.
+
+The code and pretrained models of **BEiT** can be found at [here](https://github.com/microsoft/unilm/tree/master/beit).
+
+The code and pretrained models of **BEiT v2** can be found at [here](https://github.com/microsoft/unilm/tree/master/beit2).
+
+- March, 2023: release [the code and pretrained models of **BEiT-3**](https://github.com/microsoft/unilm/tree/master/beit3)
+- March, 2023: [**BEiT-3**](https://arxiv.org/abs/2208.10442) was accepted by **CVPR 2023**.
+- Sept 2022: release [the code and pretrained models of **BEiT v2**](https://github.com/microsoft/unilm/tree/master/beit2)
+- Aug 2022: release preprint [Image as a Foreign Language: BEiT Pretraining for All Vision and Vision-Language Tasks](https://arxiv.org/abs/2208.10442)
+- Aug 2022: release preprint [BEiT v2: Masked Image Modeling with Vector-Quantized Visual Tokenizers](https://arxiv.org/abs/2208.06366)
+- June 2022: release preprint [VL-BEiT: Generative Vision-Language Pretraining](https://arxiv.org/abs/2206.01127)
+- March, 2022: add [linear probe examples](https://github.com/microsoft/unilm/blob/master/beit/get_started_for_image_classification.md#example-linear-probe-on-imagenet)
+- January, 2022: [**BEiT**](https://openreview.net/forum?id=p-BhZSz59o4) was accepted by **ICLR 2022 as Oral presentation** (54 out of 3391).
+- August 2021: [**BEiT**](https://huggingface.co/transformers/master/model_doc/beit.html) is on [HuggingFace](https://github.com/huggingface/transformers)
+- July 2021: BEiT-large achieves **[state-of-the-art results on ADE20K](https://paperswithcode.com/sota/semantic-segmentation-on-ade20k) (a big jump to 57.0 mIoU) for semantic segmentation**.
+- July 2021: BEiT-large achieves **state-of-the-art ImageNet top-1 accuracy (88.6%) under the setting without extra data other than ImageNet-22k**.
+- July 2021: release [the code and pretrained models of **BEiT**](https://github.com/microsoft/unilm/tree/master/beit)
+- June 2021: release preprint [BEiT: BERT Pre-Training of Image Transformers](https://arxiv.org/abs/2106.08254)
+
+## Pretrained models
+
+We provide BEiT-3 weights pretrained on monomodal and multimodal data. Our large-size model outperforms previous large-size models across various vision-language and vision downstream tasks. The models were pretrained with 224x224 resolution.
+
+### Tips
+- For vision-language tasks that require deep fusion, we recommend using `BEiT3-base` and `BEiT3-large`.
+- For image-text retrieval or vision tasks, using `BEiT3-base-itc` and `BEiT3-large-itc` usually achieve better performance.
+
+### Download Checkpoints
+
+1. Models pretrained on ImageNet-21k images, 160 GB text documents, and web-scale image-text pairs (collected from [LAION-400M](https://laion.ai/blog/laion-400-open-dataset/), [English LAION-2B](https://laion.ai/blog/laion-5b/), [COYO-700M](https://github.com/kakaobrain/coyo-dataset), and CC15M).
+ - [`BEiT3-base`](https://github.com/addf400/files/releases/download/beit3/beit3_base_patch16_224.pth): #layer=12; hidden=768; FFN factor=4x; #head=12; patch=16x16; #parameters: 276M
+ - [`BEiT3-large`](https://github.com/addf400/files/releases/download/beit3/beit3_large_patch16_224.pth): #layer=24; hidden=1024; FFN factor=4x; #head=16; patch=16x16; #parameters: 746M
+
+2. Perform image-text contrastive intermediate tuning on `BEiT3-base` and `BEiT3-large`.
+ - [`BEiT3-base-itc`](https://github.com/addf400/files/releases/download/beit3/beit3_base_itc_patch16_224.pth): #layer=12; hidden=768; FFN factor=4x; #head=12; patch=16x16; #parameters: 222M
+ - [`BEiT3-large-itc`](https://github.com/addf400/files/releases/download/beit3/beit3_large_itc_patch16_224.pth): #layer=24; hidden=1024; FFN factor=4x; #head=16; patch=16x16; #parameters: 674M
+
+3. Add indomain image-text pairs (COCO and VG) to continue training `BEiT3-base` and `BEiT3-large` using masked data modeling. The indomain models achieve better performance on VQAv2 and NLVR2 tasks.
+ - [`BEiT3-base-indomain`](https://github.com/addf400/files/releases/download/beit3/beit3_base_indomain_patch16_224.pth): #layer=12; hidden=768; FFN factor=4x; #head=12; patch=16x16; #parameters: 276M
+ - [`BEiT3-large-indomain`](https://github.com/addf400/files/releases/download/beit3/beit3_large_indomain_patch16_224.pth): #layer=24; hidden=1024; FFN factor=4x; #head=16; patch=16x16; #parameters: 746M
+
+### Text Tokenizer
+
+[beit3.spm](https://github.com/addf400/files/releases/download/beit3/beit3.spm) is the sentencepiece model used for tokenizing texts.
+```
+from transformers import XLMRobertaTokenizer
+tokenizer = XLMRobertaTokenizer("/your_beit3_model_path/beit3.spm")
+```
+
+### Architecture
+
+We use [Magneto](https://arxiv.org/abs/2210.06423) with decoupled Multiway Transformer as the backbone architecture. Magneto can have better training stability and obtain better performance across modalities (such as vision, and language). The implementation is based on the [torchscale](https://github.com/microsoft/torchscale/blob/main/torchscale/model/BEiT3.py) package.
+
+
+## Setup
+
+```
+alias=`whoami | cut -d'.' -f2`; docker run -it --rm --runtime=nvidia --ipc=host --privileged -v /home/${alias}:/home/${alias} pytorch/pytorch:1.8.1-cuda11.1-cudnn8-devel bash
+```
+
+Clone the repo and install required packages:
+```
+git clone https://github.com/microsoft/unilm.git
+cd unilm/beit3
+pip install -r requirements.txt
+```
+
+
+## Fine-tuning on ImageNet-1k (Image Classification)
+
+The detailed instructions can be found at [`get_started_for_image_classification.md`](get_started/get_started_for_image_classification.md). We only use vision-related parameters for image classification fine-tuning.
+
+| initialized checkpoint | resolution | acc@1 | acc@5 | #params | weight |
+|:----------------------------------------|:----------:|:-----:|:-----:|:-------:|-------------------|
+| [beit3_base_patch16_224](https://github.com/addf400/files/releases/download/beit3/beit3_base_patch16_224.pth) | 224x224 | 85.4 | 97.6 | 87M | [link](https://github.com/addf400/files/releases/download/beit3/beit3_base_patch16_224_in1k.pth) |
+| [beit3_base_indomain_patch16_224](https://github.com/addf400/files/releases/download/beit3/beit3_base_indomain_patch16_224.pth) | 224x224 | 85.4 | 97.6 | 87M | [link](https://github.com/addf400/files/releases/download/beit3/beit3_base_indomain_patch16_224_in1k.pth) |
+| [beit3_large_patch16_224](https://github.com/addf400/files/releases/download/beit3/beit3_large_patch16_224.pth) | 224x224 | 87.6 | 98.3 | 305M | [link](https://github.com/addf400/files/releases/download/beit3/beit3_large_patch16_224_in1k.pth) |
+| [beit3_large_indomain_patch16_224](https://github.com/addf400/files/releases/download/beit3/beit3_large_indomain_patch16_224.pth) | 224x224 | 87.5 | 98.3 | 305M | [link](https://github.com/addf400/files/releases/download/beit3/beit3_large_indomain_patch16_224_in1k.pth) |
+
+
+## Fine-tuning on VQAv2 (Visual Question Answering)
+
+The detailed instructions can be found at [`get_started_for_vqav2.md`](get_started/get_started_for_vqav2.md).
+
+| initialized checkpoint | resolution | augmented data | test-dev | test-std | #params | weight |
+|:----------------------------------------|:----------:|:-----:|:-----:|:-----:|:-------:|-------------------|
+| [beit3_base_patch16_224](https://github.com/addf400/files/releases/download/beit3/beit3_base_patch16_224.pth) | 480x480 | - | 77.65 | - | 228M | [link](https://github.com/addf400/files/releases/download/beit3/beit3_base_patch16_480_vqa.pth) |
+| [beit3_base_indomain_patch16_224](https://github.com/addf400/files/releases/download/beit3/beit3_base_indomain_patch16_224.pth) | 480x480 | - | 78.46 | - | 228M | [link](https://github.com/addf400/files/releases/download/beit3/beit3_base_indomain_patch16_480_vqa.pth) |
+| [beit3_large_patch16_224](https://github.com/addf400/files/releases/download/beit3/beit3_large_patch16_224.pth) | 480x480 | - | 81.85 | - | 683M | [link](https://github.com/addf400/files/releases/download/beit3/beit3_large_patch16_480_vqa.pth) |
+| [beit3_large_indomain_patch16_224](https://github.com/addf400/files/releases/download/beit3/beit3_large_indomain_patch16_224.pth) | 480x480 | - | 82.53 | - | 683M | [link](https://github.com/addf400/files/releases/download/beit3/beit3_large_indomain_patch16_480_vqa.pth) |
+| [beit3_large_indomain_patch16_224](https://github.com/addf400/files/releases/download/beit3/beit3_large_indomain_patch16_224.pth) | 768x768 | VGQA | 82.97 | 83.03 | 684M | [link](https://github.com/addf400/files/releases/download/beit3/beit3_large_indomain_patch16_768_vgqaaug_vqa.pth) |
+
+
+## Fine-tuning on NLVR2 (Visual Reasoning)
+
+The detailed instructions can be found at [`get_started_for_nlvr2.md`](get_started/get_started_for_nlvr2.md).
+
+| initialized checkpoint | resolution | dev | test-P | #params | weight |
+|:----------------------------------------|:----------:|:-----:|:-----:|:-------:|-------------------|
+| [beit3_base_patch16_224](https://github.com/addf400/files/releases/download/beit3/beit3_base_patch16_224.pth) | 224x224 | 83.6 | 84.4 | 226M | [link](https://github.com/addf400/files/releases/download/beit3/beit3_base_patch16_224_nlvr2.pth) |
+| [beit3_base_indomain_patch16_224](https://github.com/addf400/files/releases/download/beit3/beit3_base_indomain_patch16_224.pth) | 224x224 | 84.6 | 85.3 | 226M | [link](https://github.com/addf400/files/releases/download/beit3/beit3_base_indomain_patch16_224_nlvr2.pth) |
+| [beit3_large_patch16_224](https://github.com/addf400/files/releases/download/beit3/beit3_large_patch16_224.pth) | 224x224 | 88.5 | 89.4 | 681M | [link](https://github.com/addf400/files/releases/download/beit3/beit3_large_patch16_224_nlvr2.pth) |
+| [beit3_large_indomain_patch16_224](https://github.com/addf400/files/releases/download/beit3/beit3_large_indomain_patch16_224.pth) | 224x224 | 89.2 | 90.0 | 681M | [link](https://github.com/addf400/files/releases/download/beit3/beit3_large_indomain_patch16_224_nlvr2.pth) |
+
+
+## Fine-tuning on COCO Captioning and NoCaps (Image Captioning)
+
+The detailed instructions can be found at [`get_started_for_image_captioning.md`](get_started/get_started_for_captioning.md).
+
+### COCO Captioning
+
+| initialized checkpoint | resolution | test CIDEr | #params | weight |
+|:----------------------------------------|:----------:|:-----:|:-------:|-------------------|
+| [beit3_base_patch16_224](https://github.com/addf400/files/releases/download/beit3/beit3_base_patch16_224.pth) | 480x480 | 133.6 | 271M | [link](https://github.com/addf400/files/releases/download/beit3/beit3_base_patch16_480_coco_captioning.pth) |
+| [beit3_base_indomain_patch16_224](https://github.com/addf400/files/releases/download/beit3/beit3_base_indomain_patch16_224.pth) | 480x480 | 135.0 | 271M | [link](https://github.com/addf400/files/releases/download/beit3/beit3_base_indomain_patch16_480_coco_captioning.pth) |
+| [beit3_large_patch16_224](https://github.com/addf400/files/releases/download/beit3/beit3_large_patch16_224.pth) | 480x480 | 143.2 | 739M | [link](https://github.com/addf400/files/releases/download/beit3/beit3_large_patch16_480_coco_captioning.pth) |
+
+### NoCaps
+
+| initialized checkpoint | resolution | val CIDEr | #params | weight |
+|:----------------------------------------|:----------:|:-----:|:-------:|-------------------|
+| [beit3_base_patch16_224](https://github.com/addf400/files/releases/download/beit3/beit3_base_patch16_224.pth) | 480x480 | 104.4 | 271M | [link](https://github.com/addf400/files/releases/download/beit3/beit3_base_patch16_480_nocaps.pth) |
+| [beit3_base_indomain_patch16_224](https://github.com/addf400/files/releases/download/beit3/beit3_base_indomain_patch16_224.pth) | 480x480 | 105.6 | 271M | [link](https://github.com/addf400/files/releases/download/beit3/beit3_base_indomain_patch16_480_nocaps.pth) |
+| [beit3_large_patch16_224](https://github.com/addf400/files/releases/download/beit3/beit3_large_patch16_224.pth) | 480x480 | 120.2 | 739M | [link](https://github.com/addf400/files/releases/download/beit3/beit3_large_patch16_480_nocaps.pth) |
+
+
+## Fine-tuning on COCO and Flickr30k Retrieval (Image-Text Retrieval)
+
+The detailed instructions can be found at [`get_started_for_retrieval.md`](get_started/get_started_for_retrieval.md).
+
+### COCO Retrieval
+
+| initialized checkpoint | resolution | IR@1 | TR@1 | #params | weight |
+|:----------------------------------------|:----------:|:-----:|:-----:|:-------:|-------------------|
+| [beit3_base_itc_patch16_224](https://github.com/addf400/files/releases/download/beit3/beit3_base_itc_patch16_224.pth) | 384x384 | 61.4 | 79.1 | 222M | [link](https://github.com/addf400/files/releases/download/beit3/beit3_base_patch16_384_coco_retrieval.pth) |
+| [beit3_large_itc_patch16_224](https://github.com/addf400/files/releases/download/beit3/beit3_large_itc_patch16_224.pth) | 384x384 | 63.4 | 82.1 | 675M | [link](https://github.com/addf400/files/releases/download/beit3/beit3_large_patch16_384_coco_retrieval.pth) |
+
+### Flickr30k Retrieval
+
+| initialized checkpoint | resolution | IR@1 | TR@1 | #params | weight |
+|:----------------------------------------|:----------:|:-----:|:-----:|:-------:|-------------------|
+| [beit3_base_itc_patch16_224](https://github.com/addf400/files/releases/download/beit3/beit3_base_itc_patch16_224.pth) | 384x384 | 86.2 | 96.3 | 222M | [link](https://github.com/addf400/files/releases/download/beit3/beit3_base_patch16_384_f30k_retrieval.pth) |
+| [beit3_large_itc_patch16_224](https://github.com/addf400/files/releases/download/beit3/beit3_large_itc_patch16_224.pth) | 384x384 | 88.1 | 97.2 | 675M | [link](https://github.com/addf400/files/releases/download/beit3/beit3_large_patch16_384_f30k_retrieval.pth) |
+
+
+## Citation
+
+If you find this repository useful, please consider citing our work:
+```
+@inproceedings{beit3,
+title={Image as a foreign language: {BEiT} pretraining for vision and vision-language tasks},
+author={Wenhui Wang and Hangbo Bao and Li Dong and Johan Bjorck and Zhiliang Peng and Qiang Liu and Kriti Aggarwal and Owais Khan Mohammed and Saksham Singhal and Subhojit Som and Furu Wei},
+booktitle={Proceedings of the IEEE/CVF Conference on Computer Vision and Pattern Recognition},
+year={2023}
+}
+
+@article{beitv2,
+title={{BEiT v2}: Masked Image Modeling with Vector-Quantized Visual Tokenizers},
+author={Zhiliang Peng and Li Dong and Hangbo Bao and Qixiang Ye and Furu Wei},
+year={2022},
+eprint={2208.06366},
+archivePrefix={arXiv},
+primaryClass={cs.CV}
+}
+
+@inproceedings{beit,
+title={{BEiT}: {BERT} Pre-Training of Image Transformers},
+author={Hangbo Bao and Li Dong and Songhao Piao and Furu Wei},
+booktitle={International Conference on Learning Representations},
+year={2022},
+url={https://openreview.net/forum?id=p-BhZSz59o4}
+}
+```
+
+
+## Acknowledgement
+
+This repository is built using the [BEiT](https://github.com/microsoft/unilm/tree/master/beit), the [BEiTv2](https://github.com/microsoft/unilm/tree/master/beit2), the [CLIP](https://github.com/openai/CLIP), the [open_clip](https://github.com/mlfoundations/open_clip), the [Oscar](https://github.com/microsoft/Oscar), the [DeiT](https://github.com/facebookresearch/deit), the [Dino](https://github.com/facebookresearch/dino) repository and the [timm](https://github.com/rwightman/pytorch-image-models) library.
+
+
+## License
+This project is licensed under the license found in the LICENSE file in the root directory of this source tree.
+
+[Microsoft Open Source Code of Conduct](https://opensource.microsoft.com/codeofconduct)
+
+### Contact Information
+
+For help or issues using BEiT-3 models, please submit a GitHub issue.
diff --git a/py/evf_sam/model/unilm/beit3/datasets.py b/py/evf_sam/model/unilm/beit3/datasets.py
new file mode 100644
index 0000000..9f6dab8
--- /dev/null
+++ b/py/evf_sam/model/unilm/beit3/datasets.py
@@ -0,0 +1,847 @@
+# --------------------------------------------------------
+# Image as a Foreign Language: BEiT Pretraining for Vision and Vision-Language Tasks (https://arxiv.org/abs/2208.10442)
+# Github source: https://github.com/microsoft/unilm/tree/master/beit3
+# Copyright (c) 2023 Microsoft
+# Licensed under The MIT License [see LICENSE for details]
+# --------------------------------------------------------'
+
+import os
+import json
+import random
+import torch
+import glob
+from collections import defaultdict, Counter
+from torchvision import transforms
+from torchvision.datasets.folder import default_loader
+from timm.data.constants import IMAGENET_DEFAULT_MEAN, IMAGENET_DEFAULT_STD, IMAGENET_INCEPTION_MEAN, IMAGENET_INCEPTION_STD
+from timm.data.transforms import RandomResizedCropAndInterpolation
+from timm.data import create_transform
+
+import utils
+from glossary import normalize_word
+from randaug import RandomAugment
+
+
+class BaseDataset(torch.utils.data.Dataset):
+ def __init__(
+ self, data_path, split, transform,
+ tokenizer, num_max_bpe_tokens, task=None,
+ ):
+ index_files = self.get_index_files(split, task=task)
+ self.tokenizer = tokenizer
+ self.num_max_bpe_tokens = num_max_bpe_tokens
+ self.data_path = data_path
+ items = []
+ self.index_files = index_files
+
+ offset = 0
+ for _index_file in index_files:
+ index_file = os.path.join(data_path, _index_file)
+ with open(index_file, mode="r", encoding="utf-8") as reader:
+ for line in reader:
+ data = json.loads(line)
+ items.append(data)
+ print("Load %d image-text pairs from %s. " % (len(items) - offset, index_file))
+ offset = len(items)
+ self.items = items
+ self.bos_token_id = tokenizer.bos_token_id
+ self.eos_token_id = tokenizer.eos_token_id
+ self.pad_token_id = tokenizer.pad_token_id
+ self.loader = default_loader
+ self.transform = transform
+ self.split = split
+
+ @staticmethod
+ def get_index_files(split):
+ raise NotImplementedError()
+
+ def _get_image(self, image_path: str):
+ image_path = os.path.join(self.data_path, image_path)
+ image = self.loader(image_path)
+ return self.transform(image)
+
+ def _get_text_segment(self, text_segment, max_len=None):
+ if isinstance(text_segment, str):
+ tokens = self.tokenizer.tokenize(text_segment)
+ else:
+ tokens = text_segment[:]
+ if len(tokens) == 0:
+ raise RuntimeError("The text segment should contains at least one tokens!")
+ if max_len is None:
+ max_len = self.num_max_bpe_tokens
+
+ if len(tokens) > max_len - 2:
+ tokens = tokens[:max_len - 2]
+
+ tokens = [self.bos_token_id] + tokens[:] + [self.eos_token_id]
+ num_tokens = len(tokens)
+ padding_mask = [0] * num_tokens + [1] * (max_len - num_tokens)
+ return tokens + [self.pad_token_id] * (max_len - num_tokens), padding_mask, num_tokens
+
+ def _get_image_text_example(self, index: int, data: dict):
+ item = self.items[index]
+ img_path = item["image_path"]
+ img = self._get_image(img_path)
+ data["image"] = img
+
+ text_segment = item["text_segment"]
+ language_tokens, padding_mask, _ = self._get_text_segment(text_segment)
+ data["language_tokens"] = language_tokens
+ data["padding_mask"] = padding_mask
+
+ def __getitem__(self, index: int):
+ data = dict()
+ self._get_image_text_example(index, data)
+ return data
+
+ def __len__(self) -> int:
+ return len(self.items)
+
+ def __repr__(self) -> str:
+ head = "Dataset " + self.__class__.__name__
+ body = '{' + "\n Number of items: %s," % self.__len__()
+ body += "\n data root = %s," % self.data_path
+ body += "\n split = %s," % self.split
+ body += "\n dataset index files = %s" % str(self.index_files)
+ body += "\n num max bpe tokens = %s" % self.num_max_bpe_tokens
+ body += "\n transforms = ["
+ for t in self.transform.transforms:
+ body += "\n %s" % str(t)
+ body += "\n ]"
+ body += "\n}"
+
+ return head + body
+
+
+def _write_data_into_jsonl(items, jsonl_file):
+ with open(jsonl_file, mode="w", encoding="utf-8") as writer:
+ for data in items:
+ writer.write(json.dumps(data, indent=None))
+ writer.write('\n')
+ print("Write %s with %d items !" % (jsonl_file, len(items)))
+
+
+def _make_retrieval_coco_karpathy_dataset_index(
+ data_path,
+ tokenizer,
+ split=("train", "restval"),
+ split_name="train",
+):
+ coco_karpathy_split_json_file = os.path.join(data_path, "dataset_coco.json")
+ items = []
+ image_counter = set()
+ print("read %s" % coco_karpathy_split_json_file)
+ with open(coco_karpathy_split_json_file, mode="r", encoding="utf-8") as reader:
+ data = json.loads(reader.read())
+ for item in data["images"]:
+ if item["split"] in split:
+ image_path = os.path.join(item["filepath"], item["filename"])
+ for sent in item["sentences"]:
+ tokens = tokenizer.tokenize(sent["raw"])
+ token_ids = tokenizer.convert_tokens_to_ids(tokens)
+ items.append({
+ "image_path": image_path,
+ "text_segment": token_ids,
+ "image_id": len(image_counter),
+ })
+ if image_path not in image_counter:
+ image_counter.add(image_path)
+ print("Find %d images and %d image-text pairs for karpathy dataset %s split !" % \
+ (len(image_counter), len(items), split_name))
+ index_file = os.path.join(data_path, "coco_retrieval.%s.jsonl" % split_name)
+ _write_data_into_jsonl(items, index_file)
+ pass
+
+
+def _make_captioning_coco_karpathy_dataset_index(
+ data_path,
+ tokenizer,
+ split=("train", "restval"),
+ split_name="train",
+):
+ coco_karpathy_split_json_file = os.path.join(data_path, "dataset_coco.json")
+ items = []
+ image_counter = set()
+ print("read %s" % coco_karpathy_split_json_file)
+ with open(coco_karpathy_split_json_file, mode="r", encoding="utf-8") as reader:
+ data = json.loads(reader.read())
+ for item in data["images"]:
+ if item["split"] in split:
+ image_path = os.path.join(item["filepath"], item["filename"])
+ if item["split"] in ["train", "restval"]:
+ for sent in item["sentences"]:
+ tokens = tokenizer.tokenize(sent["raw"])
+ token_ids = tokenizer.convert_tokens_to_ids(tokens)
+ items.append({
+ "image_path": image_path,
+ "text_segment": token_ids,
+ "image_id": item["cocoid"],
+ })
+ else:
+ items.append({
+ "image_path": image_path,
+ "text_segment": None,
+ "image_id": item["cocoid"],
+ })
+ if image_path not in image_counter:
+ image_counter.add(image_path)
+ print("Find %d images and %d image-text pairs for karpathy dataset %s split !" % \
+ (len(image_counter), len(items), split_name))
+ index_file = os.path.join(data_path, "coco_captioning.%s.jsonl" % split_name)
+ _write_data_into_jsonl(items, index_file)
+ pass
+
+
+def _make_nocaps_dataset_index(
+ data_path,
+ split="val",
+):
+ if split == "val":
+ json_file = "nocaps_val_4500_captions.json"
+ elif split == "test":
+ json_file = "nocaps_test_image_info.json"
+ nocaps_split_json_file = os.path.join(data_path, json_file)
+ items = []
+ image_counter = set()
+ print("read %s" % nocaps_split_json_file)
+ with open(nocaps_split_json_file, mode="r", encoding="utf-8") as reader:
+ data = json.loads(reader.read())
+ for item in data["images"]:
+ image_path = os.path.join(split, item["file_name"])
+ items.append({
+ "image_path": image_path,
+ "text_segment": None,
+ "image_id": item["id"],
+ })
+
+ if image_path not in image_counter:
+ image_counter.add(image_path)
+
+ print("Find %d images and %d image-text pairs for nocaps dataset %s split !" % \
+ (len(image_counter), len(items), split))
+ index_file = os.path.join(data_path, "nocaps.%s.jsonl" % split)
+ _write_data_into_jsonl(items, index_file)
+
+
+class NLVR2Dataset(BaseDataset):
+ @staticmethod
+ def get_index_files(split, task=None):
+ if split == "train":
+ return ("nlvr2.train.index.jsonl", )
+ elif split == "val":
+ return ("nlvr2.dev.index.jsonl", )
+ elif split == "test":
+ return ("nlvr2.test-P.index.jsonl", )
+ else:
+ raise RuntimeError("split %s is not found!" % split)
+
+ def __getitem__(self, index: int):
+ data = super().__getitem__(index)
+ item = self.items[index]
+ img_path = item["image2_path"]
+ img = self._get_image(img_path)
+ data["image2"] = img
+ data["label"] = self.items[index]["label"]
+ return data
+
+ @staticmethod
+ def __preprocess_json(preifx, json_file, tokenizer, index_file):
+ items = []
+ with open(json_file, mode="r", encoding="utf-8") as reader:
+ for line in reader:
+ data = json.loads(line)
+ path = os.path.join(preifx, str(data["directory"])) if "directory" in data else preifx
+ path = os.path.join(path, "-".join(data["identifier"].split("-")[:-1]))
+ tokens = tokenizer.tokenize(data["sentence"])
+ token_ids = tokenizer.convert_tokens_to_ids(tokens)
+ items.append({
+ "image_path": path + "-img0.png",
+ "image2_path": path + "-img1.png",
+ "text_segment": token_ids,
+ "label": 1 if data["label"] == "True" else 0,
+ "identifier": data["identifier"],
+ })
+ _write_data_into_jsonl(items, index_file)
+
+ @classmethod
+ def make_dataset_index(cls, data_path, tokenizer, nlvr_repo_path):
+ cls.__preprocess_json(
+ preifx="images/train", json_file=os.path.join(nlvr_repo_path, "nlvr2/data/train.json"),
+ tokenizer=tokenizer, index_file=os.path.join(data_path, cls.get_index_files("train")[0]),
+ )
+ cls.__preprocess_json(
+ preifx="dev", json_file=os.path.join(nlvr_repo_path, "nlvr2/data/dev.json"),
+ tokenizer=tokenizer, index_file=os.path.join(data_path, cls.get_index_files("val")[0]),
+ )
+ cls.__preprocess_json(
+ preifx="test1", json_file=os.path.join(nlvr_repo_path, "nlvr2/data/test1.json"),
+ tokenizer=tokenizer, index_file=os.path.join(data_path, cls.get_index_files("test")[0]),
+ )
+
+
+class ImageNetDataset(BaseDataset):
+ @staticmethod
+ def get_index_files(split, task=None):
+ if split == "train":
+ return ("imagenet.train.index.jsonl", )
+ elif split == "val":
+ return ("imagenet.val.index.jsonl", )
+ elif split == "test":
+ return ("imagenet.val.index.jsonl", )
+ else:
+ raise RuntimeError("split %s is not found!" % split)
+
+ def __getitem__(self, index: int):
+ data = dict()
+ item = self.items[index]
+ img_path = item["image_path"]
+ img = self._get_image(img_path)
+ data["image"] = img
+ data["label"] = item["label"]
+ return data
+
+ @staticmethod
+ def _find_classes(dir):
+ """
+ Finds the class folders in a dataset.
+ Args:
+ dir (string): Root directory path.
+ Returns:
+ tuple: (classes, class_to_idx) where classes are relative to (dir), and class_to_idx is a dictionary.
+ Ensures:
+ No class is a subdirectory of another.
+ """
+ classes = [d.name for d in os.scandir(dir) if d.is_dir()]
+ classes.sort()
+ class_to_idx = {cls_name: i for i, cls_name in enumerate(classes)}
+ return classes, class_to_idx
+
+ @staticmethod
+ def _make_imagenet_index(data_path, index_path, data_path_prefix, class_to_idx, split):
+ items = []
+ index_file = os.path.join(index_path, f"imagenet.{split}.index.jsonl")
+ for target_class in sorted(class_to_idx.keys()):
+ class_index = class_to_idx[target_class]
+ target_dir = os.path.join(data_path, target_class)
+ if not os.path.isdir(target_dir):
+ continue
+ for root, _, fnames in sorted(os.walk(target_dir, followlinks=True)):
+ for fname in sorted(fnames):
+ path = os.path.join(root, fname)
+ path = path.replace(data_path_prefix, "")
+ items.append({
+ "image_path": path,
+ "label": class_index,
+ })
+
+ _write_data_into_jsonl(items, index_file)
+
+ @classmethod
+ def make_dataset_index(cls, train_data_path, val_data_path, index_path):
+ data_path_prefix = train_data_path[:[x[0]==x[1] for x in zip(train_data_path, val_data_path)].index(0)]
+ classes, class_to_idx = cls._find_classes(train_data_path)
+ cls._make_imagenet_index(
+ data_path=train_data_path, index_path=index_path, data_path_prefix=data_path_prefix,
+ class_to_idx=class_to_idx, split="train",
+ )
+ cls._make_imagenet_index(
+ data_path=val_data_path, index_path=index_path, data_path_prefix=data_path_prefix,
+ class_to_idx=class_to_idx, split="val",
+ )
+
+
+class VQAv2Dataset(BaseDataset):
+ def __init__(self, data_path, **kwargs):
+ super().__init__(data_path=data_path, **kwargs)
+ ans2label_file = os.path.join(data_path, "answer2label.txt")
+ ans2label = {}
+ label2ans = []
+ with open(ans2label_file, mode="r", encoding="utf-8") as reader:
+ for i, line in enumerate(reader):
+ data = json.loads(line)
+ ans = data["answer"]
+ label = data["label"]
+ label = int(label)
+ assert label == i
+ ans2label[ans] = i
+ label2ans.append(ans)
+
+ self.ans2label = ans2label
+ self.label2ans = label2ans
+
+ @staticmethod
+ def get_index_files(split, task=None):
+ if split == "train":
+ return ("vqa.train.jsonl", "vqa.trainable_val.jsonl")
+ elif split == "val":
+ return ("vqa.rest_val.jsonl", )
+ elif split == "test":
+ return ("vqa.test.jsonl", )
+ elif split == "test-dev":
+ return ("vqa.test-dev.jsonl", )
+ else:
+ raise RuntimeError("split %s is not found!" % split)
+
+ def __getitem__(self, index: int):
+ data = super().__getitem__(index)
+ if "labels" in self.items[index] and len(self.items[index]["labels"]) > 0:
+ labels = [0.] * len(self.label2ans)
+ for l, s in zip(self.items[index]["labels"], self.items[index]["scores"]):
+ labels[l] = s
+ data["labels"] = torch.FloatTensor(labels)
+ else:
+ data["qid"] = self.items[index]["qid"]
+ return data
+
+ @staticmethod
+ def get_score(occurences):
+ if occurences == 0:
+ return 0.0
+ elif occurences == 1:
+ return 0.3
+ elif occurences == 2:
+ return 0.6
+ elif occurences == 3:
+ return 0.9
+ else:
+ return 1.0
+
+ @classmethod
+ def make_dataset_index(cls, data_path, tokenizer, annotation_data_path):
+ with open(os.path.join(annotation_data_path, "v2_OpenEnded_mscoco_train2014_questions.json"), "r") as fp:
+ questions_train2014 = json.load(fp)["questions"]
+ with open(os.path.join(annotation_data_path, "v2_OpenEnded_mscoco_val2014_questions.json"), "r") as fp:
+ questions_val2014 = json.load(fp)["questions"]
+ with open(os.path.join(annotation_data_path, "v2_OpenEnded_mscoco_test2015_questions.json"), "r") as fp:
+ questions_test2015 = json.load(fp)["questions"]
+ with open(os.path.join(annotation_data_path, "v2_OpenEnded_mscoco_test-dev2015_questions.json"), "r") as fp:
+ questions_test_dev2015 = json.load(fp)["questions"]
+
+ with open(os.path.join(annotation_data_path, "v2_mscoco_train2014_annotations.json"), "r") as fp:
+ annotations_train2014 = json.load(fp)["annotations"]
+ with open(os.path.join(annotation_data_path, "v2_mscoco_val2014_annotations.json"), "r") as fp:
+ annotations_val2014 = json.load(fp)["annotations"]
+
+ annotations = dict()
+
+ for split, questions in zip(
+ ["train", "val", "test", "test-dev"],
+ [questions_train2014, questions_val2014, questions_test2015, questions_test_dev2015],
+ ):
+ _annot = defaultdict(dict)
+ for q in questions:
+ question_text = q["question"]
+ tokens = tokenizer.tokenize(question_text)
+ token_ids = tokenizer.convert_tokens_to_ids(tokens)
+
+ assert q["question_id"] not in _annot[q["image_id"]]
+ _annot[q["image_id"]][q["question_id"]] = {
+ "question": question_text,
+ "token_ids": token_ids,
+ }
+
+ annotations[split] = _annot
+
+ all_major_answers = list()
+
+ for split, annots in zip(
+ ["train", "val"], [annotations_train2014, annotations_val2014],
+ ):
+ # _annot = annotations[split]
+ for q in annots:
+ all_major_answers.append(q["multiple_choice_answer"])
+
+ all_major_answers = [normalize_word(word) for word in all_major_answers]
+ counter = {k: v for k, v in Counter(all_major_answers).items() if v >= 9}
+ ans2label = {k: i for i, k in enumerate(counter.keys())}
+ label2ans = list(counter.keys())
+
+ for split, annots in zip(
+ ["train", "val"], [annotations_train2014, annotations_val2014],
+ ):
+ _annot = annotations[split]
+ for q in annots:
+ answers = q["answers"]
+ answer_count = {}
+ for answer in answers:
+ answer_ = answer["answer"]
+ answer_count[answer_] = answer_count.get(answer_, 0) + 1
+
+ labels = []
+ scores = []
+ for answer in answer_count:
+ if answer not in ans2label:
+ continue
+ labels.append(ans2label[answer])
+ score = cls.get_score(answer_count[answer])
+ scores.append(score)
+
+ assert "labels" not in _annot[q["image_id"]][q["question_id"]]
+ assert "question" in _annot[q["image_id"]][q["question_id"]]
+ _annot[q["image_id"]][q["question_id"]]["labels"] = labels
+ _annot[q["image_id"]][q["question_id"]]["scores"] = scores
+
+ for split in ["train", "val"]:
+ filtered_annot = dict()
+ for ik, iv in annotations[split].items():
+ new_q = dict()
+ for qk, qv in iv.items():
+ if len(qv["labels"]) != 0:
+ new_q[qk] = qv
+ if len(new_q) != 0:
+ filtered_annot[ik] = new_q
+ annotations[split] = filtered_annot
+
+ split2items = {}
+ for split in ["train", "val", "test", "test-dev"]:
+ annot = annotations[split]
+ split_name = {
+ "train": "train2014",
+ "val": "val2014",
+ "test": "test2015",
+ "test-dev": "test2015",
+ }[split]
+ paths = list(glob.glob(f"{data_path}/{split_name}/*.jpg"))
+ random.shuffle(paths)
+ annot_paths = [path for path in paths \
+ if int(path.split("/")[-1].split("_")[-1][:-4]) in annot]
+
+ if len(paths) == len(annot_paths):
+ print("all images have caption annotations")
+ else:
+ print("not all images have caption annotations")
+ print(len(paths), len(annot_paths), len(annot))
+
+ items = []
+ for path in annot_paths:
+ iid = int(path.split("/")[-1].split("_")[-1][:-4])
+ _annot = annotations[split][iid]
+ for qid in _annot:
+ q = _annot[qid]
+ if split in ["train", "val"]:
+ labels = q["labels"]
+ scores = q["scores"]
+ else:
+ labels, scores = [], []
+
+ items.append({
+ "image_path": os.path.join(split_name, path.split('/')[-1]),
+ "text_segment": q["token_ids"],
+ "labels": labels,
+ "scores": scores,
+ "qid": qid,
+ })
+ split2items[split] = items
+
+ _write_data_into_jsonl(items=items, jsonl_file=os.path.join(data_path, "vqa.%s.jsonl" % split))
+
+ # Following ViLT, we use 1000 images of the original val set as the final val set
+ val_image2items = defaultdict(list)
+ for item in split2items["val"]:
+ val_image2items[item["image_path"]].append(item)
+
+ print("Contains %d image and %d pairs for val set!" % (len(val_image2items), len(split2items["val"])))
+
+ val_images = list(val_image2items.keys())
+ random.shuffle(val_images)
+ trainable_val = []
+ rest_val = []
+ for i, image_id in enumerate(val_images):
+ if i < 1000:
+ rest_val += val_image2items[image_id]
+ else:
+ trainable_val += val_image2items[image_id]
+
+ _write_data_into_jsonl(items=trainable_val, jsonl_file=os.path.join(data_path, "vqa.trainable_val.jsonl"))
+ _write_data_into_jsonl(items=rest_val, jsonl_file=os.path.join(data_path, "vqa.rest_val.jsonl"))
+
+ with open(os.path.join(data_path, "answer2label.txt"), mode="w", encoding="utf-8") as writer:
+ for ans in ans2label:
+ to_json = {
+ "answer": ans,
+ "label": ans2label[ans]
+ }
+ writer.write("%s\n" % json.dumps(to_json))
+
+
+class RetrievalDataset(BaseDataset):
+ @staticmethod
+ def get_index_files(split, task=None):
+ if split == "train":
+ return (f"{task}.train.jsonl", )
+ elif split == "val":
+ return (f"{task}.val.jsonl", )
+ elif split == "test":
+ return (f"{task}.test.jsonl", )
+ else:
+ raise RuntimeError("split %s is not found!" % split)
+
+ def __getitem__(self, index: int):
+ data = super().__getitem__(index)
+ data["image_id"] = self.items[index]["image_id"]
+ return data
+
+ @staticmethod
+ def make_flickr30k_dataset_index(data_path, tokenizer, karpathy_path):
+
+ with open(os.path.join(karpathy_path, "dataset_flickr30k.json"), "r") as reader:
+ captions = json.loads(reader.read())
+
+ captions = captions["images"]
+ split2items = defaultdict(list)
+ split2images = defaultdict(set)
+
+ for each_item in captions:
+ image_path = os.path.join("flickr30k-images", each_item["filename"])
+ split = each_item["split"]
+
+ for text_segment in each_item["sentences"]:
+ tokens = tokenizer.tokenize(text_segment["raw"])
+ token_ids = tokenizer.convert_tokens_to_ids(tokens)
+
+ split2items[split].append({
+ "image_path": image_path,
+ "text_segment": token_ids,
+ "image_id": len(split2images[split]),
+ })
+
+ assert each_item["filename"] not in split2images[split]
+ split2images[split].add(each_item["filename"])
+
+ for split in split2items:
+ print("%d images and %d image-text pairs!" % (len(split2images[split]), len(split2items[split])))
+ _write_data_into_jsonl(split2items[split], os.path.join(data_path, "flickr30k.%s.jsonl" % split))
+
+ @staticmethod
+ def make_coco_dataset_index(data_path, tokenizer):
+ _make_retrieval_coco_karpathy_dataset_index(data_path, tokenizer, split=("train", "restval"), split_name="train")
+ _make_retrieval_coco_karpathy_dataset_index(data_path, tokenizer, split=("val", ), split_name="val")
+ _make_retrieval_coco_karpathy_dataset_index(data_path, tokenizer, split=("test", ), split_name="test")
+
+
+class CaptioningDataset(BaseDataset):
+
+ def __init__(self, data_path, split, transform,
+ tokenizer, num_max_bpe_tokens, task, mask_prob):
+ super().__init__(
+ data_path=data_path, split=split,
+ transform=transform, tokenizer=tokenizer,
+ num_max_bpe_tokens=num_max_bpe_tokens, task=task,
+ )
+ self.mask_token_id = tokenizer.mask_token_id
+ self.language_vocab_size = tokenizer.vocab_size
+ self.mask_prob = mask_prob
+
+ @staticmethod
+ def get_index_files(split, task=None):
+ if split == "train":
+ return ("coco_captioning.train.jsonl", )
+ elif split == "val":
+ return (f"{task}.val.jsonl", )
+ elif split == "test":
+ return (f"{task}.test.jsonl", )
+ else:
+ raise RuntimeError("split %s is not found!" % split)
+
+ def _get_mask_token(self, token):
+ p = random.random()
+ if p < 0.8:
+ return self.mask_token_id
+ elif p < 0.9:
+ return token
+ else:
+ return random.randint(3, self.language_vocab_size - 1)
+
+ def _masking_on_text_tokens(self, tokens, num_tokens, mask_prob):
+ bool_masked_pos = [0] * len(tokens)
+ to_mask = min(int(num_tokens * mask_prob + 0.5), num_tokens - 1)
+ to_mask = max(to_mask, 1)
+ num_masked_tokens = 0
+ while num_masked_tokens < to_mask:
+ i = random.randint(1, num_tokens - 1)
+ if bool_masked_pos[i] == 0:
+ bool_masked_pos[i] = 1
+ tokens[i] = self._get_mask_token(tokens[i])
+ num_masked_tokens += 1
+
+ return tokens, bool_masked_pos
+
+ def __getitem__(self, index: int):
+ data = dict()
+ item = self.items[index]
+ img_path = item["image_path"]
+ img = self._get_image(img_path)
+ data["image"] = img
+ data["image_id"] = item["image_id"]
+
+ text_segment = item["text_segment"]
+ if text_segment is not None:
+ language_tokens, padding_mask, num_tokens = self._get_text_segment(text_segment)
+ masked_tokens = language_tokens[:]
+ masked_tokens, language_masked_pos = \
+ self._masking_on_text_tokens(masked_tokens, num_tokens, self.mask_prob)
+ data["language_tokens"] = language_tokens
+ data["masked_tokens"] = masked_tokens
+ data["language_masked_pos"] = language_masked_pos
+ data["padding_mask"] = padding_mask
+ return data
+
+ @staticmethod
+ def make_coco_captioning_dataset_index(data_path, tokenizer):
+ _make_captioning_coco_karpathy_dataset_index(data_path, tokenizer, split=("train", "restval"), split_name="train")
+ _make_captioning_coco_karpathy_dataset_index(data_path, tokenizer, split=("val", ), split_name="val")
+ _make_captioning_coco_karpathy_dataset_index(data_path, tokenizer, split=("test", ), split_name="test")
+
+ @staticmethod
+ def make_nocaps_captioning_dataset_index(data_path):
+ _make_nocaps_dataset_index(data_path, split="val")
+ _make_nocaps_dataset_index(data_path, split="test")
+
+
+task2dataset = {
+ "nlvr2": NLVR2Dataset,
+ "vqav2": VQAv2Dataset,
+ "flickr30k": RetrievalDataset,
+ "coco_retrieval": RetrievalDataset,
+ "coco_captioning": CaptioningDataset,
+ "nocaps": CaptioningDataset,
+ "imagenet": ImageNetDataset,
+}
+
+
+def create_dataloader(dataset, is_train, batch_size, num_workers, pin_mem, dist_eval=False):
+ if is_train or dist_eval:
+ num_tasks = utils.get_world_size()
+ global_rank = utils.get_rank()
+
+ if not is_train and dist_eval and len(dataset) % num_tasks != 0:
+ print('Warning: Enabling distributed evaluation with an eval dataset not divisible by process number. '
+ 'This will slightly alter validation results as extra duplicate entries are added to achieve '
+ 'equal num of samples per-process.')
+
+ sampler = torch.utils.data.DistributedSampler(
+ dataset, num_replicas=num_tasks, rank=global_rank, shuffle=is_train
+ )
+ else:
+ sampler = torch.utils.data.SequentialSampler(dataset)
+
+ return torch.utils.data.DataLoader(
+ dataset, sampler=sampler,
+ batch_size=batch_size,
+ num_workers=num_workers,
+ pin_memory=pin_mem,
+ drop_last=is_train,
+ collate_fn=utils.merge_batch_tensors_by_dict_key,
+ )
+
+
+def build_transform(is_train, args):
+ if args.task in ["imagenet"]:
+ return build_imagenet_transform(is_train, args)
+
+ if is_train:
+ t = [
+ RandomResizedCropAndInterpolation(args.input_size, scale=(0.5, 1.0), interpolation=args.train_interpolation),
+ transforms.RandomHorizontalFlip(),
+ ]
+ if args.randaug:
+ t.append(
+ RandomAugment(
+ 2, 7, isPIL=True,
+ augs=[
+ 'Identity','AutoContrast','Equalize','Brightness','Sharpness',
+ 'ShearX', 'ShearY', 'TranslateX', 'TranslateY', 'Rotate',
+ ]))
+ t += [
+ transforms.ToTensor(),
+ transforms.Normalize(mean=IMAGENET_INCEPTION_MEAN, std=IMAGENET_INCEPTION_STD),
+ ]
+ t = transforms.Compose(t)
+ else:
+ t = transforms.Compose([
+ transforms.Resize((args.input_size, args.input_size), interpolation=3),
+ transforms.ToTensor(),
+ transforms.Normalize(mean=IMAGENET_INCEPTION_MEAN, std=IMAGENET_INCEPTION_STD)
+ ])
+
+ return t
+
+
+def build_imagenet_transform(is_train, args):
+ resize_im = args.input_size > 32
+ if is_train:
+ # this should always dispatch to transforms_imagenet_train
+ transform = create_transform(
+ input_size=args.input_size,
+ is_training=True,
+ color_jitter=args.color_jitter,
+ auto_augment=args.aa,
+ interpolation=args.train_interpolation,
+ re_prob=args.reprob,
+ re_mode=args.remode,
+ re_count=args.recount,
+ mean=IMAGENET_DEFAULT_MEAN,
+ std=IMAGENET_DEFAULT_STD,
+ )
+ if not resize_im:
+ # replace RandomResizedCropAndInterpolation with
+ # RandomCrop
+ transform.transforms[0] = transforms.RandomCrop(
+ args.input_size, padding=4)
+ return transform
+
+ t = []
+ if resize_im:
+ if args.crop_pct is None:
+ args.crop_pct = 1.0
+ size = int(args.input_size / args.crop_pct)
+ t.append(
+ transforms.Resize(size, interpolation=3), # to maintain same ratio w.r.t. 224 images
+ )
+ t.append(transforms.CenterCrop(args.input_size))
+
+ t.append(transforms.ToTensor())
+ t.append(transforms.Normalize(mean=IMAGENET_DEFAULT_MEAN, std=IMAGENET_DEFAULT_STD))
+ return transforms.Compose(t)
+
+
+def get_sentencepiece_model_for_beit3(args):
+ from transformers import XLMRobertaTokenizer
+ return XLMRobertaTokenizer(args.sentencepiece_model)
+
+
+def create_dataset_by_split(args, split, is_train=True):
+ transform = build_transform(is_train=is_train, args=args)
+ dataset_class = task2dataset[args.task]
+ tokenizer = get_sentencepiece_model_for_beit3(args)
+
+ opt_kwargs = {}
+ if args.task in ["coco_captioning", "nocaps"]:
+ opt_kwargs["mask_prob"] = args.captioning_mask_prob
+
+ dataset = dataset_class(
+ data_path=args.data_path, split=split,
+ transform=transform, tokenizer=tokenizer,
+ num_max_bpe_tokens=args.num_max_bpe_tokens,
+ task=args.task, **opt_kwargs,
+ )
+ if is_train:
+ batch_size = args.batch_size
+ elif hasattr(args, "eval_batch_size") and args.eval_batch_size is not None:
+ batch_size = args.eval_batch_size
+ else:
+ batch_size = int(args.batch_size * 1.5)
+
+ return create_dataloader(
+ dataset, is_train=is_train, batch_size=batch_size,
+ num_workers=args.num_workers, pin_mem=args.pin_mem, dist_eval=args.dist_eval,
+ )
+
+
+def create_downstream_dataset(args, is_eval=False):
+ if is_eval:
+ return create_dataset_by_split(args, split="test", is_train=False)
+ else:
+ return \
+ create_dataset_by_split(args, split="train", is_train=True), \
+ create_dataset_by_split(args, split="val", is_train=True)
diff --git a/py/evf_sam/model/unilm/beit3/engine_for_finetuning.py b/py/evf_sam/model/unilm/beit3/engine_for_finetuning.py
new file mode 100644
index 0000000..9be308c
--- /dev/null
+++ b/py/evf_sam/model/unilm/beit3/engine_for_finetuning.py
@@ -0,0 +1,598 @@
+# --------------------------------------------------------
+# Image as a Foreign Language: BEiT Pretraining for Vision and Vision-Language Tasks (https://arxiv.org/abs/2208.10442)
+# Github source: https://github.com/microsoft/unilm/tree/master/beit3
+# Copyright (c) 2023 Microsoft
+# Licensed under The MIT License [see LICENSE for details]
+# --------------------------------------------------------'
+
+import math
+import sys
+import json
+from typing import Iterable, Optional
+
+import torch
+import torch.nn as nn
+import torch.nn.functional as F
+
+from timm.utils import ModelEma
+from timm.utils import accuracy, ModelEma
+from timm.loss import LabelSmoothingCrossEntropy, SoftTargetCrossEntropy
+from datasets import get_sentencepiece_model_for_beit3
+
+import utils
+
+
+class TaskHandler(object):
+ def __init__(self) -> None:
+ self.metric_logger = None
+ self.split = None
+
+ def train_batch(self, model, **kwargs):
+ raise NotImplementedError()
+
+ def eval_batch(self, model, **kwargs):
+ raise NotImplementedError()
+
+ def before_eval(self, metric_logger, data_loader, **kwargs):
+ self.metric_logger = metric_logger
+ self.split = data_loader.dataset.split
+
+ def after_eval(self, **kwargs):
+ raise NotImplementedError()
+
+
+class NLVR2Handler(TaskHandler):
+ def __init__(self) -> None:
+ super().__init__()
+ self.criterion = torch.nn.CrossEntropyLoss()
+
+ def train_batch(self, model, image, image2, language_tokens, padding_mask, label):
+ logits = model(
+ image_a=image, image_b=image2,
+ text_description=language_tokens,
+ padding_mask=padding_mask)
+ acc = (logits.max(-1)[-1] == label).float().mean()
+ return {
+ "loss": self.criterion(input=logits, target=label),
+ "acc": acc,
+ }
+
+ def eval_batch(self, model, image, image2, language_tokens, padding_mask, label):
+ logits = model(
+ image_a=image, image_b=image2,
+ text_description=language_tokens,
+ padding_mask=padding_mask)
+ batch_size = language_tokens.shape[0]
+ acc = (logits.max(-1)[-1] == label).float().sum(0) * 100.0 / batch_size
+ self.metric_logger.meters['acc'].update(acc.item(), n=batch_size)
+
+ def after_eval(self, **kwargs):
+ print('* Acc {acc.global_avg:.3f}'.format(acc=self.metric_logger.acc))
+ return {k: meter.global_avg for k, meter in self.metric_logger.meters.items()}, "acc"
+
+
+class ImageNetHandler(TaskHandler):
+ def __init__(self, args) -> None:
+ super().__init__()
+ mixup_active = args.mixup > 0 or args.cutmix > 0. or args.cutmix_minmax is not None
+ if mixup_active:
+ # smoothing is handled with mixup label transform
+ self.criterion = SoftTargetCrossEntropy()
+ elif args.label_smoothing > 0.:
+ self.criterion = LabelSmoothingCrossEntropy(smoothing=args.label_smoothing)
+ else:
+ self.criterion = torch.nn.CrossEntropyLoss()
+
+ def train_batch(self, model, image, label):
+ logits = model(image=image)
+ return {
+ "loss": self.criterion(logits, label),
+ }
+
+ def eval_batch(self, model, image, label):
+ logits = model(image=image)
+ batch_size = image.shape[0]
+ acc1, acc5 = accuracy(logits, label, topk=(1, 5))
+ self.metric_logger.meters['acc1'].update(acc1.item(), n=batch_size)
+ self.metric_logger.meters['acc5'].update(acc5.item(), n=batch_size)
+
+ def after_eval(self, **kwargs):
+ print('* Acc@1 {top1.global_avg:.3f} Acc@5 {top5.global_avg:.3f}'
+ .format(top1=self.metric_logger.acc1, top5=self.metric_logger.acc5))
+ return {k: meter.global_avg for k, meter in self.metric_logger.meters.items()}, "acc1"
+
+
+class RetrievalHandler(TaskHandler):
+ def __init__(self) -> None:
+ super().__init__()
+ self.image_feats = []
+ self.text_feats = []
+ self.image_ids = []
+ self.metric_logger = None
+
+ def train_batch(self, model, image, language_tokens, padding_mask, image_id):
+ loss, vision_cls, language_cls = model(
+ image=image, text_description=language_tokens, padding_mask=padding_mask)
+ return {
+ "loss": loss,
+ }
+
+ def before_eval(self, metric_logger, **kwargs):
+ self.image_feats.clear()
+ self.text_feats.clear()
+ self.image_ids.clear()
+ self.metric_logger = metric_logger
+
+ def eval_batch(self, model, image, language_tokens, padding_mask, image_id):
+ vision_cls, _ = model(image=image, only_infer=True)
+ _, language_cls = model(
+ text_description=language_tokens, padding_mask=padding_mask, only_infer=True)
+
+ self.image_feats.append(vision_cls.clone())
+ self.text_feats.append(language_cls.clone())
+ self.image_ids.append(image_id.clone())
+
+ def after_eval(self, **kwargs):
+ image_feats = {}
+ for feats, ids in zip(self.image_feats, self.image_ids):
+ for i, _idx in enumerate(ids):
+ idx = _idx.item()
+ if idx not in image_feats:
+ image_feats[idx] = feats[i]
+
+ tiids = torch.cat(self.image_ids, dim=0)
+ iids = []
+ sorted_tensors = []
+ for key in sorted(image_feats.keys()):
+ sorted_tensors.append(image_feats[key].view(1, -1))
+ iids.append(key)
+
+ image_cls_feats = torch.cat(sorted_tensors, dim=0)
+ text_cls_feats = torch.cat(self.text_feats, dim=0)
+
+ scores = image_cls_feats @ text_cls_feats.t()
+ iids = torch.LongTensor(iids).to(scores.device)
+
+ print("scores: {}".format(scores.size()))
+ print("iids: {}".format(iids.size()))
+ print("tiids: {}".format(tiids.size()))
+
+ topk10 = scores.topk(10, dim=1)
+ topk5 = scores.topk(5, dim=1)
+ topk1 = scores.topk(1, dim=1)
+
+ topk10_iids = tiids[topk10.indices]
+ topk5_iids = tiids[topk5.indices]
+ topk1_iids = tiids[topk1.indices]
+
+ tr_r10 = (iids.unsqueeze(1) == topk10_iids).float().max(dim=1)[0].mean()
+ tr_r5 = (iids.unsqueeze(1) == topk5_iids).float().max(dim=1)[0].mean()
+ tr_r1 = (iids.unsqueeze(1) == topk1_iids).float().max(dim=1)[0].mean()
+
+ topk10 = scores.topk(10, dim=0)
+ topk5 = scores.topk(5, dim=0)
+ topk1 = scores.topk(1, dim=0)
+ topk10_iids = iids[topk10.indices]
+ topk5_iids = iids[topk5.indices]
+ topk1_iids = iids[topk1.indices]
+
+ ir_r10 = (tiids.unsqueeze(0) == topk10_iids).float().max(dim=0)[0].mean()
+ ir_r5 = (tiids.unsqueeze(0) == topk5_iids).float().max(dim=0)[0].mean()
+ ir_r1 = (tiids.unsqueeze(0) == topk1_iids).float().max(dim=0)[0].mean()
+
+ eval_result = {
+ "tr_r10": tr_r10.item() * 100.0,
+ "tr_r5": tr_r5.item() * 100.0,
+ "tr_r1": tr_r1.item() * 100.0,
+ "ir_r10": ir_r10.item() * 100.0,
+ "ir_r5": ir_r5.item() * 100.0,
+ "ir_r1": ir_r1.item() * 100.0,
+ "average_score": 100.0 * (tr_r1 + tr_r5 + tr_r10 + ir_r1 + ir_r5 + ir_r10).item() / 6.0,
+ }
+
+ print('* Eval result = %s' % json.dumps(eval_result))
+ return eval_result, "average_score"
+
+
+class VQAHandler(TaskHandler):
+ def __init__(self) -> None:
+ super().__init__()
+ self.predictions = []
+ self.criterion = nn.BCEWithLogitsLoss(reduction='mean')
+ self.label2ans = None
+
+ def train_batch(self, model, image, language_tokens, padding_mask, labels):
+ logits = model(
+ image=image, question=language_tokens,
+ padding_mask=padding_mask)
+ return {
+ "loss": self.criterion(input=logits.float(), target=labels.float()) * labels.shape[1],
+ }
+
+ def before_eval(self, metric_logger, data_loader, **kwargs):
+ self.predictions.clear()
+ self.metric_logger = metric_logger
+ self.label2ans = data_loader.dataset.label2ans
+
+ def eval_batch(self, model, image, language_tokens, padding_mask, labels=None, qid=None):
+ logits = model(
+ image=image, question=language_tokens,
+ padding_mask=padding_mask)
+ batch_size = language_tokens.shape[0]
+ if labels is not None:
+ scores = utils.VQAScore()(logits, labels) * 100.0
+ self.metric_logger.meters['score'].update(scores.item(), n=batch_size)
+ else:
+ _, preds = logits.max(-1)
+ for image_id, pred in zip(qid, preds):
+ self.predictions.append({
+ "question_id": image_id.item(),
+ "answer": self.label2ans[pred.item()],
+ })
+
+ def after_eval(self, **kwargs):
+ if len(self.predictions) == 0:
+ print('* Score {score.global_avg:.3f}'.format(score=self.metric_logger.score))
+ return {k: meter.global_avg for k, meter in self.metric_logger.meters.items()}, "score"
+ else:
+ return self.predictions, "prediction"
+
+
+class CaptioningHandler(TaskHandler):
+ def __init__(self, args) -> None:
+ super().__init__()
+ self.predictions = []
+ self.criterion = utils.BertCaptioningLoss(args.label_smoothing, args.drop_worst_ratio, args.drop_worst_after)
+ self.tokenizer = get_sentencepiece_model_for_beit3(args)
+ self.num_beams = args.num_beams
+ self.max_len = args.num_max_bpe_tokens
+ self.length_penalty = args.length_penalty
+ self.vocab_size = args.vocab_size
+
+ def train_batch(self, model, image, language_tokens, masked_tokens, language_masked_pos, padding_mask, image_id, global_step):
+ logits, _ = model(
+ image=image, text_ids=masked_tokens, padding_mask=padding_mask, language_masked_pos=language_masked_pos, image_id=image_id)
+ masked_labels = language_tokens[language_masked_pos.bool()]
+ score = torch.max(logits, -1)[1].data == masked_labels
+ acc = torch.sum(score.float()) / torch.sum(language_masked_pos)
+ return {
+ "loss": self.criterion(logits, masked_labels, global_step),
+ "acc": acc
+ }
+
+ def before_eval(self, metric_logger, data_loader, **kwargs):
+ self.predictions.clear()
+ self.metric_logger = metric_logger
+
+ def eval_batch(self, model, image, image_id=None):
+ cur_len = 2
+ num_keep_best = 1
+ TOPN_PER_BEAM = 3
+
+ batch_size = image.size(0)
+ mask_id = self.tokenizer.mask_token_id
+ cls_id = self.tokenizer.cls_token_id
+ pad_id = self.tokenizer.pad_token_id
+ sep_id = self.tokenizer.sep_token_id
+ eos_token_ids = [sep_id]
+
+ cls_ids = torch.full(
+ (batch_size, 1), cls_id, dtype=torch.long, device=image.device
+ )
+ mask_ids = torch.full(
+ (batch_size, 1), mask_id, dtype=torch.long, device=image.device
+ )
+ cur_input_ids = torch.cat([cls_ids, mask_ids], dim=1)
+ tmp_ids = torch.full(
+ (batch_size, self.max_len-1), mask_id, dtype=torch.long, device=image.device
+ )
+ decoding_results = torch.cat([cls_ids, tmp_ids], dim=1)
+
+ # Expand input to num beams
+ cur_input_ids = cur_input_ids.unsqueeze(1).expand(batch_size, self.num_beams, cur_len)
+ cur_input_ids = cur_input_ids.contiguous().view(batch_size * self.num_beams, cur_len) # (batch_size * num_beams, cur_len)
+ decoding_results = decoding_results.unsqueeze(1).expand(batch_size, self.num_beams, self.max_len)
+ decoding_results = decoding_results.contiguous().view(batch_size * self.num_beams, self.max_len) # (batch_size * num_beams, cur_len)
+ image = image.unsqueeze(1).expand(batch_size, self.num_beams, image.size(-3), image.size(-2), image.size(-1))
+ image = image.contiguous().view(batch_size * self.num_beams, image.size(-3), image.size(-2), image.size(-1))
+
+ generated_hyps = [
+ utils.BeamHypotheses(
+ num_keep_best, self.max_len, length_penalty=self.length_penalty, early_stopping=False
+ ) for _ in range(batch_size)
+ ]
+ # scores for each sentence in the beam
+ beam_scores = torch.zeros((batch_size, self.num_beams), dtype=torch.float, device=cur_input_ids.device)
+ beam_scores[:, 1:] = -1e9
+ beam_scores = beam_scores.view(-1) # shape (batch_size * num_beams,)
+
+ # done sentences
+ done = [False for _ in range(batch_size)]
+ incremental_state = {}
+
+ while cur_len <= self.max_len:
+ next_token_idx = 1
+ padding_masks = torch.full(
+ cur_input_ids.shape, 0, dtype=torch.long, device=image.device
+ )
+ input_image = image
+ if cur_len != 2:
+ input_image = None
+
+ outputs, incremental_state_next = model(
+ image=input_image, text_ids=cur_input_ids, language_masked_pos=None,
+ padding_mask=padding_masks, text_len=cur_len, incremental_state=incremental_state)
+ incremental_state = incremental_state_next
+
+ # assert outputs.shape[1] == token_len
+ scores = outputs[:, next_token_idx, :] # (batch_size * num_beams, vocab_size)
+ scores = F.log_softmax(scores, dim=-1) # (batch_size * num_beams, vocab_size)
+ assert scores.size() == (batch_size * self.num_beams, self.vocab_size)
+ # Add the log prob of the new beams to the log prob of the beginning of the sequence (sum of logs == log of the product)
+ _scores = scores + beam_scores[:, None].expand_as(scores) # (batch_size * num_beams, vocab_size)
+ # re-organize to group the beam together (we are keeping top hypothesis accross beams)
+ _scores = _scores.view(batch_size, self.num_beams * self.vocab_size) # (batch_size, num_beams * vocab_size)
+ next_scores, next_words = torch.topk(_scores, TOPN_PER_BEAM * self.num_beams, dim=1, largest=True, sorted=True)
+ assert next_scores.size() == next_words.size() == (batch_size, TOPN_PER_BEAM * self.num_beams)
+
+ # next batch beam content
+ # list of (batch_size * num_beams) tuple(next hypothesis score, next word, current position in the batch)
+ next_batch_beam = []
+ # for each sentence
+ for batch_ex in range(batch_size):
+ # if we are done with this sentence
+ done[batch_ex] = done[batch_ex] or generated_hyps[batch_ex].is_done(next_scores[batch_ex].max().item())
+ if done[batch_ex]:
+ next_batch_beam.extend([(0, pad_id, 0)] * self.num_beams) # pad the batch
+ continue
+
+ # next sentence beam content
+ next_sent_beam = []
+ for idx, score in zip(next_words[batch_ex], next_scores[batch_ex]):
+ # get beam and word IDs
+ beam_id = idx // self.vocab_size
+ word_id = idx % self.vocab_size
+ # end of sentence, or next word
+ # if word_id.item() in eos_token_ids or cur_len + 1 == max_len:
+ if (word_id.item() in eos_token_ids and cur_len + 1 <= self.max_len) or (cur_len + 1 == self.max_len):
+ generated_hyps[batch_ex].add(
+ decoding_results[batch_ex * self.num_beams + beam_id, :cur_len].clone(), score.item()
+ )
+ else:
+ next_sent_beam.append((score, word_id, batch_ex * self.num_beams + beam_id))
+ # the beam for next step is full
+ if len(next_sent_beam) == self.num_beams:
+ break
+
+ # update next beam content
+ if cur_len + 1 == self.max_len:
+ assert len(next_sent_beam) == 0
+ else:
+ assert len(next_sent_beam) == self.num_beams
+
+ if len(next_sent_beam) == 0:
+ next_sent_beam = [(0, pad_id, 0)] * self.num_beams # pad the batch
+ next_batch_beam.extend(next_sent_beam)
+ assert len(next_batch_beam) == self.num_beams * (batch_ex + 1)
+
+ # sanity check / prepare next batch
+ assert len(next_batch_beam) == batch_size * self.num_beams
+ beam_scores = beam_scores.new([x[0] for x in next_batch_beam])
+ beam_words = cur_input_ids.new([x[1] for x in next_batch_beam])
+ beam_idx = cur_input_ids.new([x[2] for x in next_batch_beam])
+
+ # re-order batch
+ cur_input_ids = cur_input_ids[beam_idx, :]
+ decoding_results = decoding_results[beam_idx, :]
+ for module in incremental_state:
+ for key in incremental_state[module]:
+ result = incremental_state[module][key].index_select(0, beam_idx)
+ incremental_state[module][key] = result[:,:,:-1,:]
+
+ next_ids = torch.full(
+ (batch_size * self.num_beams, 1), mask_id, dtype=torch.long, device=image.device
+ )
+ cur_input_ids = torch.cat([beam_words.unsqueeze(1), next_ids], dim=1)
+ decoding_results[:, cur_len-1] = beam_words
+ # update current length
+ cur_len = cur_len + 1
+ # stop when we are done with each sentence
+ if all(done):
+ break
+
+ # select the best hypotheses
+ tgt_len = torch.ones(batch_size, num_keep_best, dtype=torch.long)
+ logprobs = torch.zeros(batch_size, num_keep_best,
+ dtype=torch.float).fill_(-1e5).to(cur_input_ids.device)
+ all_best = []
+
+ for i, hypotheses in enumerate(generated_hyps):
+ best = []
+ hyp_scores = torch.tensor([x[0] for x in hypotheses.hyp])
+ _, best_indices = torch.topk(hyp_scores,
+ min(num_keep_best, len(hyp_scores)), largest=True)
+ for best_idx, hyp_idx in enumerate(best_indices):
+ conf, best_hyp = hypotheses.hyp[hyp_idx]
+ best.append(best_hyp)
+ logprobs[i, best_idx] = conf
+ tgt_len[i, best_idx] = len(best_hyp) + 1 # +1 for the symbol
+ all_best.append(best)
+
+ # generate target batch, pad to the same length
+ decoded = cur_input_ids.new(batch_size, num_keep_best, self.max_len).fill_(pad_id)
+ for batch_idx, best in enumerate(all_best):
+ for best_idx, hypo in enumerate(best):
+ decoded[batch_idx, best_idx, : tgt_len[batch_idx, best_idx] - 1] = hypo
+ decoded[batch_idx, best_idx, tgt_len[batch_idx, best_idx] - 1] = eos_token_ids[0]
+
+ captions = self.tokenizer.batch_decode(decoded.squeeze(1), skip_special_tokens=True)
+ for qid, pred in zip(image_id, captions):
+ self.predictions.append({
+ "image_id": qid.item(),
+ "caption": pred,
+ })
+
+ def after_eval(self, **kwargs):
+ return self.predictions, "prediction"
+
+
+def get_handler(args):
+ if args.task == "nlvr2":
+ return NLVR2Handler()
+ elif args.task == "vqav2":
+ return VQAHandler()
+ elif args.task in ("flickr30k", "coco_retrieval"):
+ return RetrievalHandler()
+ elif args.task in ("coco_captioning", "nocaps"):
+ return CaptioningHandler(args)
+ elif args.task in ("imagenet"):
+ return ImageNetHandler(args)
+ else:
+ raise NotImplementedError("Sorry, %s is not support." % args.task)
+
+
+def train_one_epoch(
+ model: torch.nn.Module, data_loader: Iterable,
+ optimizer: torch.optim.Optimizer, device: torch.device,
+ handler: TaskHandler, epoch: int, start_steps: int,
+ lr_schedule_values: list, loss_scaler, max_norm: float = 0,
+ update_freq: int = 1, model_ema: Optional[ModelEma] = None,
+ log_writer: Optional[utils.TensorboardLogger] = None,
+ task = None, mixup_fn=None,
+):
+ model.train(True)
+ metric_logger = utils.MetricLogger(delimiter=" ")
+ metric_logger.add_meter('lr', utils.SmoothedValue(window_size=1, fmt='{value:.6f}'))
+ metric_logger.add_meter('min_lr', utils.SmoothedValue(window_size=1, fmt='{value:.6f}'))
+ header = 'Epoch: [{}]'.format(epoch)
+ print_freq = 10
+
+ if loss_scaler is None:
+ model.zero_grad()
+ model.micro_steps = 0
+ else:
+ optimizer.zero_grad()
+
+ for data_iter_step, data in enumerate(metric_logger.log_every(data_loader, print_freq, header)):
+ step = data_iter_step // update_freq
+ global_step = start_steps + step # global training iteration
+ # Update LR & WD for the first acc
+ if lr_schedule_values is not None and data_iter_step % update_freq == 0:
+ for i, param_group in enumerate(optimizer.param_groups):
+ if lr_schedule_values is not None:
+ param_group["lr"] = lr_schedule_values[global_step] * param_group["lr_scale"]
+ # put input data into cuda
+ for tensor_key in data.keys():
+ data[tensor_key] = data[tensor_key].to(device, non_blocking=True)
+ # print("input %s = %s" % (tensor_key, data[tensor_key]))
+ if loss_scaler is None and tensor_key.startswith("image"):
+ data[tensor_key] = data[tensor_key].half()
+
+ # mixup for imagenet finetuning
+ if mixup_fn is not None:
+ data["image"], data["label"] = mixup_fn(data["image"], data["label"])
+
+ if task in ["coco_captioning", "nocaps"]:
+ data["global_step"] = global_step
+
+ if loss_scaler is None:
+ results = handler.train_batch(model, **data)
+ else:
+ with torch.cuda.amp.autocast():
+ results = handler.train_batch(model, **data)
+
+ loss = results.pop("loss")
+ loss_value = loss.item()
+
+ if not math.isfinite(loss_value):
+ print("Loss is {}, stopping training".format(loss_value))
+ sys.exit(1)
+
+ if loss_scaler is None:
+ loss /= update_freq
+ model.backward(loss)
+ model.step()
+
+ if (data_iter_step + 1) % update_freq == 0:
+ # model.zero_grad()
+ # Deepspeed will call step() & model.zero_grad() automatic
+ if model_ema is not None:
+ model_ema.update(model)
+ grad_norm = None
+ loss_scale_value = utils.get_loss_scale_for_deepspeed(model)
+ else:
+ # this attribute is added by timm on one optimizer (adahessian)
+ is_second_order = hasattr(optimizer, 'is_second_order') and optimizer.is_second_order
+ loss /= update_freq
+ grad_norm = loss_scaler(loss, optimizer, clip_grad=max_norm,
+ parameters=model.parameters(), create_graph=is_second_order,
+ update_grad=(data_iter_step + 1) % update_freq == 0)
+ if (data_iter_step + 1) % update_freq == 0:
+ optimizer.zero_grad()
+ if model_ema is not None:
+ model_ema.update(model)
+ loss_scale_value = loss_scaler.state_dict()["scale"]
+
+ torch.cuda.synchronize()
+
+ metric_logger.update(loss=loss_value)
+ metric_logger.update(loss_scale=loss_scale_value)
+ min_lr = 10.
+ max_lr = 0.
+ for group in optimizer.param_groups:
+ min_lr = min(min_lr, group["lr"])
+ max_lr = max(max_lr, group["lr"])
+
+ metric_logger.update(lr=max_lr)
+ metric_logger.update(min_lr=min_lr)
+ weight_decay_value = None
+ for group in optimizer.param_groups:
+ if group["weight_decay"] > 0:
+ weight_decay_value = group["weight_decay"]
+ metric_logger.update(weight_decay=weight_decay_value)
+ metric_logger.update(grad_norm=grad_norm)
+
+ if log_writer is not None:
+ kwargs = {
+ "loss": loss_value,
+ }
+ for key in results:
+ kwargs[key] = results[key]
+ log_writer.update(head="train", **kwargs)
+
+ kwargs = {
+ "loss_scale": loss_scale_value,
+ "lr": max_lr,
+ "min_lr": min_lr,
+ "weight_decay": weight_decay_value,
+ "grad_norm": grad_norm,
+ }
+ log_writer.update(head="opt", **kwargs)
+ log_writer.set_step()
+
+ # gather the stats from all processes
+ metric_logger.synchronize_between_processes()
+ print("Averaged stats:", metric_logger)
+ return {k: meter.global_avg for k, meter in metric_logger.meters.items()}
+
+
+@torch.no_grad()
+def evaluate(data_loader, model, device, handler):
+ metric_logger = utils.MetricLogger(delimiter=" ")
+ header = 'Test:'
+
+ # switch to evaluation mode
+ model.eval()
+ handler.before_eval(metric_logger=metric_logger, data_loader=data_loader)
+
+ for data in metric_logger.log_every(data_loader, 10, header):
+ for tensor_key in data.keys():
+ data[tensor_key] = data[tensor_key].to(device, non_blocking=True)
+
+ with torch.cuda.amp.autocast():
+ handler.eval_batch(model=model, **data)
+
+ # gather the stats from all processes
+ metric_logger.synchronize_between_processes()
+
+ return handler.after_eval()
diff --git a/py/evf_sam/model/unilm/beit3/get_started/get_started_for_captioning.md b/py/evf_sam/model/unilm/beit3/get_started/get_started_for_captioning.md
new file mode 100644
index 0000000..3b0f6bd
--- /dev/null
+++ b/py/evf_sam/model/unilm/beit3/get_started/get_started_for_captioning.md
@@ -0,0 +1,176 @@
+# Fine-tuning BEiT-3 on Image Captioning
+
+## COCO Captioning Setup
+
+1. [Setup environment](../README.md#setup).
+2. Download [2014 train images](http://images.cocodataset.org/zips/train2014.zip), [2014 val images](http://images.cocodataset.org/zips/val2014.zip) and [karpathy split](https://cs.stanford.edu/people/karpathy/deepimagesent/caption_datasets.zip), then organize the dataset as following structure:
+
+```
+/path/to/your_data/
+ train2014/
+ COCO_train2014_000000000009.jpg
+ ...
+ val2014/
+ COCO_val2014_000000000042.jpg
+ ...
+ dataset_coco.json
+```
+
+We then generate the index json files using the following command. [beit3.spm](https://github.com/addf400/files/releases/download/beit3/beit3.spm) is the sentencepiece model used for tokenizing texts.
+```
+from datasets import CaptioningDataset
+from transformers import XLMRobertaTokenizer
+
+tokenizer = XLMRobertaTokenizer("/your_beit3_model_path/beit3.spm")
+
+CaptioningDataset.make_coco_captioning_dataset_index(
+ data_path="/path/to/your_data",
+ tokenizer=tokenizer,
+)
+```
+
+
+## NoCaps Setup
+
+1. [Setup environment](README.md#setup).
+2. Download [NoCaps val set](https://nocaps.s3.amazonaws.com/nocaps_val_4500_captions.json), [NoCaps test set](https://s3.amazonaws.com/nocaps/nocaps_test_image_info.json) and download imags using the urls in val and test json files, then organize the dataset as following structure:
+
+```
+/path/to/your_data/
+ val/
+ 09c863d76bcf6b00.jpg
+ ...
+ test/
+ 19dc6913830a0a21.jpg
+ ...
+ nocaps_val_4500_captions.json
+ nocaps_test_image_info.json
+```
+
+We then generate the index json files using the following command. [beit3.spm](https://github.com/addf400/files/releases/download/beit3/beit3.spm) is the sentencepiece model used for tokenizing texts.
+```
+from datasets import CaptioningDataset
+from transformers import XLMRobertaTokenizer
+
+tokenizer = XLMRobertaTokenizer("/your_beit3_model_path/beit3.spm")
+
+CaptioningDataset.make_nocaps_captioning_dataset_index(
+ data_path="/path/to/your_data",
+)
+```
+We use COCO captioning training set as the training data of NoCaps.
+
+
+## Example: Fine-tuning BEiT-3 on Captioning
+
+The BEiT-3 **base** model can be fine-tuned on captioning tasks using 8 V100-32GB:
+
+```bash
+python -m torch.distributed.launch --nproc_per_node=8 run_beit3_finetuning.py \
+ --model beit3_base_patch16_480 \
+ --input_size 480 \
+ --task coco_captioning \
+ --batch_size 32 \
+ --layer_decay 1.0 \
+ --lr 4e-5 \
+ --randaug \
+ --epochs 10 \
+ --warmup_epochs 1 \
+ --drop_path 0.1 \
+ --sentencepiece_model /your_beit3_model_path/beit3.spm \
+ --finetune /your_beit3_model_path/beit3_base_patch16_224.pth \
+ --data_path /path/to/your_data \
+ --output_dir /path/to/save/your_model \
+ --log_dir /path/to/save/your_model/log \
+ --weight_decay 0.05 \
+ --seed 42 \
+ --save_ckpt_freq 5 \
+ --num_max_bpe_tokens 32 \
+ --captioning_mask_prob 0.7 \
+ --drop_worst_after 12000 \
+ --dist_eval \
+ --checkpoint_activations \
+ --enable_deepspeed
+```
+- `--batch_size`: batch size per GPU. Effective batch size = `number of GPUs` * `--batch_size` * `--update_freq`. So in the above example, the effective batch size is `8*32 = 256`.
+- `--finetune`: weight path of your pretrained models; please download the pretrained model weights in [README.md](../README.md#pretrained-models).
+- `--task`: **coco_captioning** for COCO captioning and **nocaps** for NoCaps dataset.
+- `lr`: 4e-5 for COCO captioning and 1e-5 for NoCaps.
+- `--enable_deepspeed`: optional. If you use apex, please enable deepspeed.
+- `--checkpoint_activations`: using gradient checkpointing for saving GPU memory.
+
+
+The BEiT-3 **large** model can be fine-tuned on captioning tasks using 8 V100-32GB:
+
+```bash
+python -m torch.distributed.launch --nproc_per_node=8 run_beit3_finetuning.py \
+ --model beit3_large_patch16_480 \
+ --input_size 480 \
+ --task coco_captioning \
+ --batch_size 32 \
+ --layer_decay 1.0 \
+ --lr 8e-6 \
+ --randaug \
+ --epochs 10 \
+ --warmup_epochs 1 \
+ --drop_path 0.1 \
+ --sentencepiece_model /your_beit3_model_path/beit3.spm \
+ --finetune /your_beit3_model_path/beit3_large_patch16_224.pth \
+ --data_path /path/to/your_data \
+ --output_dir /path/to/save/your_model \
+ --log_dir /path/to/save/your_model/log \
+ --weight_decay 0.05 \
+ --seed 42 \
+ --save_ckpt_freq 5 \
+ --num_max_bpe_tokens 32 \
+ --captioning_mask_prob 0.7 \
+ --drop_worst_after 12000 \
+ --dist_eval \
+ --checkpoint_activations \
+ --enable_deepspeed
+```
+- `--batch_size`: batch size per GPU. Effective batch size = `number of GPUs` * `--batch_size` * `--update_freq`. So in the above example, the effective batch size is `8*32 = 256`.
+- `--finetune`: weight path of your pretrained models; please download the pretrained model weights in [README.md](../README.md#pretrained-models).
+- `--task`: **coco_captioning** for COCO captioning and **nocaps** for NoCaps dataset.
+- `lr`: 8e-6 for COCO captioning and NoCaps.
+- `--enable_deepspeed`: optional. If you use apex, please enable deepspeed.
+- `--checkpoint_activations`: using gradient checkpointing for saving GPU memory.
+
+
+## Example: Evaluate BEiT-3 Fine-tuned model on Captioning
+
+- Get the prediction file of the fine-tuned BEiT3-base model on captioning with 8 V100-32GB:
+```bash
+python -m torch.distributed.launch --nproc_per_node=8 run_beit3_finetuning.py \
+ --model beit3_base_patch16_480 \
+ --input_size 480 \
+ --task coco_captioning \
+ --batch_size 16 \
+ --sentencepiece_model /your_beit3_model_path/beit3.spm \
+ --finetune /your_beit3_model_path/beit3_base_patch16_480_coco_captioning.pth \
+ --data_path /path/to/your_data \
+ --output_dir /path/to/save/your_prediction \
+ --eval \
+ --dist_eval
+```
+- `--task`: **coco_captioning** for COCO captioning and **nocaps** for NoCaps dataset.
+- `--finetune`: **beit3_base_patch16_480_coco_captioning.pth** for COCO captioning and **beit3_base_patch16_480_nocaps.pth** for NoCaps dataset.
+
+- Get the prediction file of the fine-tuned BEiT3-large model on captioning with 8 V100-32GB:
+```bash
+python -m torch.distributed.launch --nproc_per_node=8 run_beit3_finetuning.py \
+ --model beit3_large_patch16_480 \
+ --input_size 480 \
+ --task coco_captioning \
+ --batch_size 16 \
+ --sentencepiece_model /your_beit3_model_path/beit3.spm \
+ --finetune /your_beit3_model_path/beit3_large_patch16_480_coco_captioning.pth \
+ --data_path /path/to/your_data \
+ --output_dir /path/to/save/your_prediction \
+ --eval \
+ --dist_eval
+```
+- `--task`: **coco_captioning** for COCO captioning and **nocaps** for NoCaps dataset.
+- `--finetune`: **beit3_large_patch16_480_coco_captioning.pth** for COCO captioning and **beit3_large_patch16_480_nocaps.pth** for NoCaps dataset.
+
+Please then submit the prediction file in the `output_dir` to the [evaluation server](https://eval.ai/web/challenges/challenge-page/355/overview) to obtain the NoCaps val and test results.
diff --git a/py/evf_sam/model/unilm/beit3/get_started/get_started_for_image_classification.md b/py/evf_sam/model/unilm/beit3/get_started/get_started_for_image_classification.md
new file mode 100644
index 0000000..18b0231
--- /dev/null
+++ b/py/evf_sam/model/unilm/beit3/get_started/get_started_for_image_classification.md
@@ -0,0 +1,138 @@
+# Fine-tuning BEiT-3 on ImageNet-1k (Image Classification)
+
+
+## Setup
+
+1. [Setup environment](../README.md#setup).
+2. Download and extract ImageNet-1k from http://image-net.org/.
+
+The directory structure is the standard layout of torchvision's [`datasets.ImageFolder`](https://pytorch.org/docs/stable/torchvision/datasets.html#imagefolder). The training and validation data are expected to be in the `train/` folder and `val/` folder, respectively:
+
+```
+/path/to/imagenet/
+ train/
+ class1/
+ img1.jpeg
+ class2/
+ img2.jpeg
+ val/
+ class1/
+ img3.jpeg
+ class/2
+ img4.jpeg
+```
+
+We then generate the index json files using the following command. [beit3.spm](https://github.com/addf400/files/releases/download/beit3/beit3.spm) is the sentencepiece model used for tokenizing texts.
+```
+from datasets import ImageNetDataset
+
+ImageNetDataset.make_dataset_index(
+ train_data_path = "/path/to/your_data/train",
+ val_data_path = "/path/to/your_data/val",
+ index_path = "/path/to/your_data"
+)
+```
+
+
+## Example: Fine-tuning BEiT-3 on ImageNet-1k (Image Classification)
+
+The BEiT-3 **base** model can be finetuned on ImageNet-1k using 8 V100-32GB:
+
+```bash
+python -m torch.distributed.launch --nproc_per_node=8 run_beit3_finetuning.py \
+ --model beit3_base_patch16_224 \
+ --task imagenet \
+ --batch_size 128 \
+ --layer_decay 0.65 \
+ --lr 7e-4 \
+ --update_freq 1 \
+ --epochs 50 \
+ --warmup_epochs 5 \
+ --drop_path 0.15 \
+ --sentencepiece_model /your_beit3_model_path/beit3.spm \
+ --finetune /your_beit3_model_path/beit3_base_patch16_224.pth \
+ --data_path /path/to/your_data \
+ --output_dir /path/to/save/your_model \
+ --log_dir /path/to/save/your_model/log \
+ --weight_decay 0.05 \
+ --seed 42 \
+ --save_ckpt_freq 5 \
+ --dist_eval \
+ --mixup 0.8 \
+ --cutmix 1.0 \
+ --enable_deepspeed
+```
+- `--batch_size`: batch size per GPU. Effective batch size = `number of GPUs` * `--batch_size` * `--update_freq`. So in the above example, the effective batch size is `8*128*1 = 1024`.
+- `--finetune`: weight path of your pretrained models; please download the pretrained model weights in [README.md](../README.md#pretrained-models)
+- `--enable_deepspeed`: optional. If you use apex, please enable deepspeed.
+
+
+The BEiT-3 **large** model can be finetuned on ImageNet-1k using a DGX box (8 V100-32GB):
+
+```bash
+python -m torch.distributed.launch --nproc_per_node=8 run_beit3_finetuning.py \
+ --model beit3_large_patch16_224 \
+ --task imagenet \
+ --batch_size 128 \
+ --layer_decay 0.8 \
+ --lr 2e-4 \
+ --update_freq 1 \
+ --epochs 50 \
+ --warmup_epochs 5 \
+ --drop_path 0.25 \
+ --sentencepiece_model /your_beit3_model_path/beit3.spm \
+ --finetune /your_beit3_model_path/beit3_large_patch16_224.pth \
+ --data_path /path/to/your_data \
+ --output_dir /path/to/save/your_model \
+ --log_dir /path/to/save/your_model/log \
+ --weight_decay 0.05 \
+ --seed 42 \
+ --save_ckpt_freq 5 \
+ --dist_eval \
+ --mixup 0.8 \
+ --cutmix 1.0 \
+ --enable_deepspeed \
+ --checkpoint_activations
+```
+- `--batch_size`: batch size per GPU. Effective batch size = `number of GPUs` * `--batch_size` * `--update_freq`. So in the above example, the effective batch size is `8*128 = 1024`.
+- `--finetune`: weight path of your pretrained models; please download the pretrained model weights in [README.md](../README.md#pretrained-models)
+- `--enable_deepspeed`: optional. If you use apex, please enable deepspeed.
+- `--checkpoint_activations`: using gradient checkpointing for saving GPU memory
+
+## Example: Evaluate BEiT-3 Finetuned model on ImageNet-1k (Image Classification)
+
+- Evaluate our fine-tuned BEiT3-base model on ImageNet val with a single GPU:
+```bash
+python -m torch.distributed.launch --nproc_per_node=1 run_beit3_finetuning.py \
+ --model beit3_base_patch16_224 \
+ --task imagenet \
+ --batch_size 128 \
+ --sentencepiece_model /your_beit3_model_path/beit3.spm \
+ --finetune /your_beit3_model_path/beit3_base_patch16_224_in1k.pth \
+ --data_path /path/to/your_data \
+ --eval \
+ --dist_eval
+```
+
+Expected results:
+```
+* Acc@1 85.400 Acc@5 97.630
+```
+
+- Evaluate our fine-tuned BEiT3-large model on ImageNet val with a single GPU:
+```bash
+python -m torch.distributed.launch --nproc_per_node=1 run_beit3_finetuning.py \
+ --model beit3_large_patch16_224 \
+ --task imagenet \
+ --batch_size 128 \
+ --sentencepiece_model /your_beit3_model_path/beit3.spm \
+ --finetune /your_beit3_model_path/beit3_large_patch16_224_in1k.pth \
+ --data_path /path/to/your_data \
+ --eval \
+ --dist_eval
+```
+
+Expected results:
+```
+* Acc@1 87.580 Acc@5 98.326
+```
diff --git a/py/evf_sam/model/unilm/beit3/get_started/get_started_for_nlvr2.md b/py/evf_sam/model/unilm/beit3/get_started/get_started_for_nlvr2.md
new file mode 100644
index 0000000..68e482e
--- /dev/null
+++ b/py/evf_sam/model/unilm/beit3/get_started/get_started_for_nlvr2.md
@@ -0,0 +1,136 @@
+# Fine-tuning BEiT-3 on NLVR2 (Visual Reasoning)
+
+
+## Setup
+
+1. [Setup environment](../README.md#setup).
+2. Clone the [repository](https://github.com/lil-lab/nlvr) and sign the [request form](https://goo.gl/forms/yS29stWnFWzrDBFH3) to download the images, then organize the dataset as following structure:
+
+```
+/path/to/your_data/
+ images/train/
+ 0/train-11670-0-img0.png
+ ...
+ dev/
+ dev-269-0-img0.png
+ ...
+ test1/
+ test1-261-0-img0.png
+ ...
+ nlvr/ (nlvr repo)
+ nlvr/
+ nlvr2/
+```
+
+We then generate the index json files using the following command. [beit3.spm](https://github.com/addf400/files/releases/download/beit3/beit3.spm) is the sentencepiece model used for tokenizing texts.
+```
+from datasets import NLVR2Dataset
+from transformers import XLMRobertaTokenizer
+
+tokenizer = XLMRobertaTokenizer("/your_beit3_model_path/beit3.spm")
+
+NLVR2Dataset.make_dataset_index(
+ data_path="/path/to/your_data",
+ tokenizer=tokenizer,
+ nlvr_repo_path="/path/to/your_data/nlvr"
+)
+```
+
+
+## Example: Fine-tuning BEiT-3 on NLVR2 (Visual Reasoning)
+
+The BEiT-3 **base** model can be finetuned on NLVR2 using 8 V100-32GB:
+
+```bash
+python -m torch.distributed.launch --nproc_per_node=8 run_beit3_finetuning.py \
+ --model beit3_base_patch16_224 \
+ --task nlvr2 \
+ --batch_size 32 \
+ --layer_decay 0.65 \
+ --lr 7e-4 \
+ --epochs 20 \
+ --warmup_epochs 5 \
+ --drop_path 0.2 \
+ --sentencepiece_model /your_beit3_model_path/beit3.spm \
+ --finetune /your_beit3_model_path/beit3_base_patch16_224.pth \
+ --data_path /path/to/your_data \
+ --output_dir /path/to/save/your_model \
+ --log_dir /path/to/save/your_model/log \
+ --weight_decay 0.2 \
+ --seed 42 \
+ --save_ckpt_freq 5 \
+ --enable_deepspeed
+```
+- `--batch_size`: batch size per GPU. Effective batch size = `number of GPUs` * `--batch_size` * `--update_freq`. So in the above example, the effective batch size is `8*32 = 256`.
+- `--finetune`: weight path of your pretrained models; please download the pretrained model weights in [README.md](../README.md#pretrained-models).
+- `--enable_deepspeed`: optional. If you use apex, please enable deepspeed.
+- `--lr`: 7e-4 for `BEiT3-base`, 5e-4 for `BEiT3-base-indomain`.
+
+
+The BEiT-3 **large** model can be finetuned on NLVR2 using 8 V100-32GB:
+
+```bash
+python -m torch.distributed.launch --nproc_per_node=8 run_beit3_finetuning.py \
+ --model beit3_large_patch16_224 \
+ --task nlvr2 \
+ --batch_size 32 \
+ --layer_decay 0.85 \
+ --lr 3e-4 \
+ --epochs 20 \
+ --warmup_epochs 5 \
+ --drop_path 0.2 \
+ --sentencepiece_model /your_beit3_model_path/beit3.spm \
+ --finetune /your_beit3_model_path/beit3_large_patch16_224.pth \
+ --data_path /path/to/your_data \
+ --output_dir /path/to/save/your_model \
+ --log_dir /path/to/save/your_model/log \
+ --weight_decay 0.2 \
+ --seed 42 \
+ --save_ckpt_freq 5 \
+ --enable_deepspeed \
+ --checkpoint_activations
+```
+- `--batch_size`: batch size per GPU. Effective batch size = `number of GPUs` * `--batch_size` * `--update_freq`. So in the above example, the effective batch size is `8*32 = 256`.
+- `--finetune`: weight path of your pretrained models; please download the pretrained model weights in [README.md](../README.md#pretrained-models).
+- `--enable_deepspeed`: optional. If you use apex, please enable deepspeed.
+- `--lr`: 3e-4 for `BEiT3-large`, 1e-4 for `BEiT3-large-indomain`.
+- `--checkpoint_activations`: using gradient checkpointing for saving GPU memory.
+
+
+## Example: Evaluate BEiT-3 Finetuned model on NLVR2 (Visual Reasoning)
+
+- Get the result of our fine-tuned BEiT3-base model on NLVR2 test with 8 V100-32GB:
+```bash
+python -m torch.distributed.launch --nproc_per_node=8 run_beit3_finetuning.py \
+ --model beit3_base_patch16_224 \
+ --task nlvr2 \
+ --batch_size 32 \
+ --sentencepiece_model /your_beit3_model_path/beit3.spm \
+ --finetune /your_beit3_model_path/beit3_base_patch16_224_nlvr2.pth \
+ --data_path /path/to/your_data \
+ --eval \
+ --dist_eval
+```
+
+Expected results:
+```
+* Acc 84.386
+```
+
+- Get the result of our fine-tuned BEiT3-large model on NLVR2 test with 8 V100-32GB:
+```bash
+python -m torch.distributed.launch --nproc_per_node=8 run_beit3_finetuning.py \
+ --model beit3_large_patch16_224 \
+ --task nlvr2 \
+ --batch_size 32 \
+ --sentencepiece_model /your_beit3_model_path/beit3.spm \
+ --finetune /your_beit3_model_path/beit3_large_patch16_224_nlvr2.pth \
+ --data_path /path/to/your_data \
+ --eval \
+ --dist_eval
+```
+
+Expected results:
+```
+* Acc 89.437
+```
diff --git a/py/evf_sam/model/unilm/beit3/get_started/get_started_for_retrieval.md b/py/evf_sam/model/unilm/beit3/get_started/get_started_for_retrieval.md
new file mode 100644
index 0000000..a2ef2c1
--- /dev/null
+++ b/py/evf_sam/model/unilm/beit3/get_started/get_started_for_retrieval.md
@@ -0,0 +1,161 @@
+# Fine-tuning BEiT-3 on Image-text Retrieval
+
+## COCO Retrieval Setup
+
+1. [Setup environment](../README.md#setup).
+2. Download [2014 train images](http://images.cocodataset.org/zips/train2014.zip), [2014 val images](http://images.cocodataset.org/zips/val2014.zip) and [karpathy split](https://cs.stanford.edu/people/karpathy/deepimagesent/caption_datasets.zip), then organize the dataset as following structure:
+
+```
+/path/to/your_data/
+ train2014/
+ COCO_train2014_000000000009.jpg
+ ...
+ val2014/
+ COCO_val2014_000000000042.jpg
+ ...
+ dataset_coco.json
+```
+
+We then generate the index json files using the following command. [beit3.spm](https://github.com/addf400/files/releases/download/beit3/beit3.spm) is the sentencepiece model used for tokenizing texts.
+```
+from datasets import RetrievalDataset
+from transformers import XLMRobertaTokenizer
+
+tokenizer = XLMRobertaTokenizer("/your_beit3_model_path/beit3.spm")
+
+RetrievalDataset.make_coco_dataset_index(
+ data_path="/path/to/your_data",
+ tokenizer=tokenizer,
+)
+```
+
+
+## Flickr30k Retrieval Setup
+
+1. [Setup environment](README.md#setup).
+2. Sign [flickr images request form](https://forms.illinois.edu/sec/229675) and download [karpathy split](https://cs.stanford.edu/people/karpathy/deepimagesent/caption_datasets.zip), then organize the dataset as following structure:
+
+```
+/path/to/your_data/
+ flickr30k-images/
+ 2923475135.jpg
+ ...
+ dataset_flickr30k.json
+```
+
+We then generate the index json files using the following command. [beit3.spm](https://github.com/addf400/files/releases/download/beit3/beit3.spm) is the sentencepiece model used for tokenizing texts.
+```
+from datasets import RetrievalDataset
+from transformers import XLMRobertaTokenizer
+
+tokenizer = XLMRobertaTokenizer("/your_beit3_model_path/beit3.spm")
+
+RetrievalDataset.make_flickr30k_dataset_index(
+ data_path="/path/to/your_data",
+ tokenizer=tokenizer,
+ karpathy_path="/path/to/your_data",
+)
+```
+
+
+## Example: Fine-tuning BEiT-3 on Retrieval
+
+The BEiT-3 **base** model can be finetuned on retrieval tasks using 16 V100-32GB:
+
+```bash
+python -m torch.distributed.launch --nproc_per_node=16 run_beit3_finetuning.py \
+ --model beit3_base_patch16_384 \
+ --input_size 384 \
+ --task coco_retrieval \
+ --batch_size 192 \
+ --layer_decay 0.65 \
+ --lr 2e-4 \
+ --epochs 15 \
+ --warmup_epochs 3 \
+ --drop_path 0.2 \
+ --sentencepiece_model /your_beit3_model_path/beit3.spm \
+ --finetune /your_beit3_model_path/beit3_base_itc_patch16_224.pth \
+ --data_path /path/to/your_data \
+ --output_dir /path/to/save/your_model \
+ --log_dir /path/to/save/your_model/log \
+ --weight_decay 0.05 \
+ --seed 42 \
+ --save_ckpt_freq 5 \
+ --enable_deepspeed \
+ --checkpoint_activations
+```
+- `--batch_size`: batch size per GPU. Effective batch size = `number of GPUs` * `--batch_size` * `--update_freq`. So in the above example, the effective batch size is `192*16 = 3072`.
+- `--finetune`: weight path of your pretrained models; please download the pretrained model weights in [README.md](../README.md#pretrained-models)
+- `--task`: **coco_retrieval** for COCO retrieval, **flickr30k** for Flickr30k retrieval
+- `--lr`: 2e-4 for COCO retrieval, 1e-4 for Flickr30k retrieval
+- `--epochs`: 15 for COCO retrieval, 20 for Flickr30k retrieval
+- `--warmup_epochs`: 3 for COCO retrieval, 5 for Flickr30k retrieval
+- `--checkpoint_activations`: using gradient checkpointing for saving GPU memory
+
+
+The BEiT-3 **large** model can be finetuned on retrieval tasks using 2x16 V100-32GB:
+
+```bash
+python -m torch.distributed.launch --nproc_per_node=16 --nnodes=2 --node_rank=$NODE_RANK \
+ --master_addr=$MASTER_ADDR --master_port=$MASTER_PORT run_beit3_finetuning.py \
+ --model beit3_large_patch16_384 \
+ --input_size 384 \
+ --task coco_retrieval \
+ --batch_size 96 \
+ --layer_decay 0.85 \
+ --lr 5e-5 \
+ --epochs 15 \
+ --warmup_epochs 3 \
+ --drop_path 0.2 \
+ --sentencepiece_model /your_beit3_model_path/beit3.spm \
+ --finetune /your_beit3_model_path/beit3_large_itc_patch16_224.pth \
+ --data_path /path/to/your_data \
+ --output_dir /path/to/save/your_model \
+ --log_dir /path/to/save/your_model/log \
+ --weight_decay 0.05 \
+ --seed 42 \
+ --save_ckpt_freq 5 \
+ --enable_deepspeed \
+ --checkpoint_activations
+```
+- `--batch_size`: batch size per GPU. Effective batch size = `number of GPUs` * `--batch_size` * `--update_freq`. So in the above example, the effective batch size is `96*32 = 3072`.
+- `--finetune`: weight path of your pretrained models; please download the pretrained model weights in [README.md](../README.md#pretrained-models)
+- `--task`: **coco_retrieval** for COCO retrieval, **flickr30k** for Flickr30k retrieval
+- `--epochs`: 15 for COCO retrieval, 20 for Flickr30k retrieval
+- `--warmup_epochs`: 3 for COCO retrieval, 5 for Flickr30k retrieval
+- `--checkpoint_activations`: using gradient checkpointing for saving GPU memory
+
+
+## Example: Evaluate BEiT-3 Fine-tuned model on COCO Retrieval and Flickr30k Retrieval
+
+- Get the results of our fine-tuned BEiT3-base model on retrieval tasks using a single GPU:
+```bash
+python -m torch.distributed.launch --nproc_per_node=1 run_beit3_finetuning.py \
+ --model beit3_base_patch16_384 \
+ --input_size 384 \
+ --task coco_retrieval \
+ --batch_size 16 \
+ --sentencepiece_model /your_beit3_model_path/beit3.spm \
+ --finetune /your_beit3_model_path/beit3_base_patch16_384_coco_retrieval.pth \
+ --data_path /path/to/your_data \
+ --eval \
+ --dist_eval
+```
+- `--task`: **coco_retrieval** for COCO retrieval, **flickr30k** for Flickr30k retrieval
+- `--finetune`: **beit3_base_patch16_384_coco_retrieval.pth** for COCO retrieval, **beit3_base_patch16_384_f30k_retrieval.pth** for Flickr30k retrieval
+
+- Get the results of our fine-tuned BEiT3-large model on retrieval tasks using a single GPU:
+```bash
+python -m torch.distributed.launch --nproc_per_node=1 run_beit3_finetuning.py \
+ --model beit3_large_patch16_384 \
+ --input_size 384 \
+ --task coco_retrieval \
+ --batch_size 16 \
+ --sentencepiece_model /your_beit3_model_path/beit3.spm \
+ --finetune /your_beit3_model_path/beit3_large_patch16_384_coco_retrieval.pth \
+ --data_path /path/to/your_data \
+ --eval \
+ --dist_eval
+```
+- `--task`: **coco_retrieval** for COCO retrieval, **flickr30k** for Flickr30k retrieval
+- `--finetune`: **beit3_large_patch16_384_coco_retrieval.pth** for COCO retrieval, **beit3_large_patch16_384_f30k_retrieval.pth** for Flickr30k retrieval
diff --git a/py/evf_sam/model/unilm/beit3/get_started/get_started_for_vqav2.md b/py/evf_sam/model/unilm/beit3/get_started/get_started_for_vqav2.md
new file mode 100644
index 0000000..b834b07
--- /dev/null
+++ b/py/evf_sam/model/unilm/beit3/get_started/get_started_for_vqav2.md
@@ -0,0 +1,144 @@
+# Fine-tuning BEiT-3 on VQAv2 (Visual Question Answering)
+
+
+## Setup
+
+1. [Setup environment](../README.md#setup).
+2. Download COCO [2014 train images](http://images.cocodataset.org/zips/train2014.zip), [2014 val images](http://images.cocodataset.org/zips/val2014.zip), [2015 test images](http://images.cocodataset.org/zips/test2015.zip), annotations ([train](https://s3.amazonaws.com/cvmlp/vqa/mscoco/vqa/v2_Annotations_Train_mscoco.zip), [val](https://s3.amazonaws.com/cvmlp/vqa/mscoco/vqa/v2_Annotations_Val_mscoco.zip)), and questions ([train](https://s3.amazonaws.com/cvmlp/vqa/mscoco/vqa/v2_Questions_Train_mscoco.zip), [val](https://s3.amazonaws.com/cvmlp/vqa/mscoco/vqa/v2_Questions_Val_mscoco.zip), [test](https://s3.amazonaws.com/cvmlp/vqa/mscoco/vqa/v2_Questions_Test_mscoco.zip)), then organize the dataset as following structure:
+
+```
+/path/to/your_data/
+ train2014/
+ COCO_train2014_000000000009.jpg
+ ...
+ val2014/
+ COCO_val2014_000000000042.jpg
+ ...
+ test2015/
+ COCO_test2015_000000000001.jpg
+ ...
+ vqa/
+ v2_OpenEnded_mscoco_train2014_questions.json
+ v2_OpenEnded_mscoco_val2014_questions.json
+ v2_OpenEnded_mscoco_test2015_questions.json
+ v2_OpenEnded_mscoco_test-dev2015_questions.json
+ v2_mscoco_train2014_annotations.json
+ v2_mscoco_val2014_annotations.json
+```
+
+We then generate the index json files using the following command. [beit3.spm](https://github.com/addf400/files/releases/download/beit3/beit3.spm) is the sentencepiece model used for tokenizing texts.
+```
+from datasets import VQAv2Dataset
+from transformers import XLMRobertaTokenizer
+
+tokenizer = XLMRobertaTokenizer("/your_beit3_model_path/beit3.spm")
+
+VQAv2Dataset.make_dataset_index(
+ data_path="/path/to/your_data",
+ tokenizer=tokenizer,
+ annotation_data_path="/path/to/your_data/vqa",
+)
+```
+
+
+## Example: Fine-tuning BEiT-3 on VQAv2 (Visual Question Answering)
+
+The BEiT-3 **base** model can be finetuned on VQAv2 using 8 V100-32GB:
+
+```bash
+python -m torch.distributed.launch --nproc_per_node=8 run_beit3_finetuning.py \
+ --model beit3_base_patch16_480 \
+ --input_size 480 \
+ --task vqav2 \
+ --batch_size 16 \
+ --layer_decay 1.0 \
+ --lr 3e-5 \
+ --update_freq 1 \
+ --randaug \
+ --epochs 10 \
+ --warmup_epochs 1 \
+ --drop_path 0.1 \
+ --sentencepiece_model /your_beit3_model_path/beit3.spm \
+ --finetune /your_beit3_model_path/beit3_base_patch16_224.pth \
+ --data_path /path/to/your_data \
+ --output_dir /path/to/save/your_model \
+ --log_dir /path/to/save/your_model/log \
+ --weight_decay 0.01 \
+ --seed 42 \
+ --save_ckpt_freq 5 \
+ --task_head_lr_weight 20 \
+ --opt_betas 0.9 0.98 \
+ --enable_deepspeed
+```
+- `--batch_size`: batch size per GPU. Effective batch size = `number of GPUs` * `--batch_size` * `--update_freq`. So in the above example, the effective batch size is `8*16 = 128`.
+- `--finetune`: weight path of your pretrained models; please download the pretrained model weights in [README.md](../README.md#pretrained-models)
+- `--enable_deepspeed`: optional. If you use apex, please enable deepspeed.
+
+
+The BEiT-3 **large** model can be finetuned on VQAv2 using 8 V100-32GB:
+
+```bash
+python -m torch.distributed.launch --nproc_per_node=8 run_beit3_finetuning.py \
+ --model beit3_large_patch16_480 \
+ --input_size 480 \
+ --task vqav2 \
+ --batch_size 16 \
+ --layer_decay 1.0 \
+ --lr 2e-5 \
+ --update_freq 1 \
+ --randaug \
+ --epochs 10 \
+ --warmup_epochs 1 \
+ --drop_path 0.15 \
+ --sentencepiece_model /your_beit3_model_path/beit3.spm \
+ --finetune /your_beit3_model_path/beit3_large_patch16_224.pth \
+ --data_path /path/to/your_data \
+ --output_dir /path/to/save/your_model \
+ --log_dir /path/to/save/your_model/log \
+ --weight_decay 0.01 \
+ --seed 42 \
+ --save_ckpt_freq 5 \
+ --task_head_lr_weight 20 \
+ --opt_betas 0.9 0.98 \
+ --enable_deepspeed \
+ --checkpoint_activations
+```
+- `--batch_size`: batch size per GPU. Effective batch size = `number of GPUs` * `--batch_size` * `--update_freq`. So in the above example, the effective batch size is `8*16 = 128`.
+- `--finetune`: weight path of your pretrained models; please download the pretrained model weights in [README.md](../README.md#pretrained-models)
+- `--enable_deepspeed`: optional. If you use apex, please enable deepspeed.
+- `--checkpoint_activations`: using gradient checkpointing for saving GPU memory
+
+
+## Example: Evaluate BEiT-3 Finetuned model on VQAv2 (Visual Question Answering)
+
+- Get the prediction file of the fine-tuned BEiT3-base model on VQAv2 test with 8 V100-32GB:
+```bash
+python -m torch.distributed.launch --nproc_per_node=8 run_beit3_finetuning.py \
+ --model beit3_base_patch16_480 \
+ --input_size 480 \
+ --task vqav2 \
+ --batch_size 16 \
+ --sentencepiece_model /your_beit3_model_path/beit3.spm \
+ --finetune /your_beit3_model_path/beit3_base_patch16_480_vqa.pth \
+ --data_path /path/to/your_data \
+ --output_dir /path/to/save/your_prediction \
+ --eval \
+ --dist_eval
+```
+
+- Get the prediction file of the fine-tuned BEiT3-large model on VQAv2 test with 8 V100-32GB:
+```bash
+python -m torch.distributed.launch --nproc_per_node=8 run_beit3_finetuning.py \
+ --model beit3_large_patch16_480 \
+ --input_size 480 \
+ --task vqav2 \
+ --batch_size 16 \
+ --sentencepiece_model /your_beit3_model_path/beit3.spm \
+ --finetune /your_beit3_model_path/beit3_large_patch16_480_vqa.pth \
+ --data_path /path/to/your_data \
+ --output_dir /path/to/save/your_prediction \
+ --eval \
+ --dist_eval
+```
+
+Please then submit the prediction file in the `output_dir` to the [evaluation server](https://eval.ai/web/challenges/challenge-page/830/overview) to obtain the VQAv2 test-dev and test-std results.
diff --git a/py/evf_sam/model/unilm/beit3/glossary.py b/py/evf_sam/model/unilm/beit3/glossary.py
new file mode 100644
index 0000000..81fddf3
--- /dev/null
+++ b/py/evf_sam/model/unilm/beit3/glossary.py
@@ -0,0 +1,190 @@
+import re
+
+contractions = {
+ "aint": "ain't",
+ "arent": "aren't",
+ "cant": "can't",
+ "couldve": "could've",
+ "couldnt": "couldn't",
+ "couldn'tve": "couldn't've",
+ "couldnt've": "couldn't've",
+ "didnt": "didn't",
+ "doesnt": "doesn't",
+ "dont": "don't",
+ "hadnt": "hadn't",
+ "hadnt've": "hadn't've",
+ "hadn'tve": "hadn't've",
+ "hasnt": "hasn't",
+ "havent": "haven't",
+ "hed": "he'd",
+ "hed've": "he'd've",
+ "he'dve": "he'd've",
+ "hes": "he's",
+ "howd": "how'd",
+ "howll": "how'll",
+ "hows": "how's",
+ "Id've": "I'd've",
+ "I'dve": "I'd've",
+ "Im": "I'm",
+ "Ive": "I've",
+ "isnt": "isn't",
+ "itd": "it'd",
+ "itd've": "it'd've",
+ "it'dve": "it'd've",
+ "itll": "it'll",
+ "let's": "let's",
+ "maam": "ma'am",
+ "mightnt": "mightn't",
+ "mightnt've": "mightn't've",
+ "mightn'tve": "mightn't've",
+ "mightve": "might've",
+ "mustnt": "mustn't",
+ "mustve": "must've",
+ "neednt": "needn't",
+ "notve": "not've",
+ "oclock": "o'clock",
+ "oughtnt": "oughtn't",
+ "ow's'at": "'ow's'at",
+ "'ows'at": "'ow's'at",
+ "'ow'sat": "'ow's'at",
+ "shant": "shan't",
+ "shed've": "she'd've",
+ "she'dve": "she'd've",
+ "she's": "she's",
+ "shouldve": "should've",
+ "shouldnt": "shouldn't",
+ "shouldnt've": "shouldn't've",
+ "shouldn'tve": "shouldn't've",
+ "somebody'd": "somebodyd",
+ "somebodyd've": "somebody'd've",
+ "somebody'dve": "somebody'd've",
+ "somebodyll": "somebody'll",
+ "somebodys": "somebody's",
+ "someoned": "someone'd",
+ "someoned've": "someone'd've",
+ "someone'dve": "someone'd've",
+ "someonell": "someone'll",
+ "someones": "someone's",
+ "somethingd": "something'd",
+ "somethingd've": "something'd've",
+ "something'dve": "something'd've",
+ "somethingll": "something'll",
+ "thats": "that's",
+ "thered": "there'd",
+ "thered've": "there'd've",
+ "there'dve": "there'd've",
+ "therere": "there're",
+ "theres": "there's",
+ "theyd": "they'd",
+ "theyd've": "they'd've",
+ "they'dve": "they'd've",
+ "theyll": "they'll",
+ "theyre": "they're",
+ "theyve": "they've",
+ "twas": "'twas",
+ "wasnt": "wasn't",
+ "wed've": "we'd've",
+ "we'dve": "we'd've",
+ "weve": "we've",
+ "werent": "weren't",
+ "whatll": "what'll",
+ "whatre": "what're",
+ "whats": "what's",
+ "whatve": "what've",
+ "whens": "when's",
+ "whered": "where'd",
+ "wheres": "where's",
+ "whereve": "where've",
+ "whod": "who'd",
+ "whod've": "who'd've",
+ "who'dve": "who'd've",
+ "wholl": "who'll",
+ "whos": "who's",
+ "whove": "who've",
+ "whyll": "why'll",
+ "whyre": "why're",
+ "whys": "why's",
+ "wont": "won't",
+ "wouldve": "would've",
+ "wouldnt": "wouldn't",
+ "wouldnt've": "wouldn't've",
+ "wouldn'tve": "wouldn't've",
+ "yall": "y'all",
+ "yall'll": "y'all'll",
+ "y'allll": "y'all'll",
+ "yall'd've": "y'all'd've",
+ "y'alld've": "y'all'd've",
+ "y'all'dve": "y'all'd've",
+ "youd": "you'd",
+ "youd've": "you'd've",
+ "you'dve": "you'd've",
+ "youll": "you'll",
+ "youre": "you're",
+ "youve": "you've",
+}
+
+manual_map = {
+ "none": "0",
+ "zero": "0",
+ "one": "1",
+ "two": "2",
+ "three": "3",
+ "four": "4",
+ "five": "5",
+ "six": "6",
+ "seven": "7",
+ "eight": "8",
+ "nine": "9",
+ "ten": "10",
+}
+articles = ["a", "an", "the"]
+period_strip = re.compile("(?!<=\d)(\.)(?!\d)")
+comma_strip = re.compile("(\d)(\,)(\d)")
+punct = [
+ ";",
+ r"/",
+ "[",
+ "]",
+ '"',
+ "{",
+ "}",
+ "(",
+ ")",
+ "=",
+ "+",
+ "\\",
+ "_",
+ "-",
+ ">",
+ "<",
+ "@",
+ "`",
+ ",",
+ "?",
+ "!",
+]
+
+
+def normalize_word(token):
+ _token = token
+ for p in punct:
+ if (p + " " in token or " " + p in token) or (
+ re.search(comma_strip, token) != None
+ ):
+ _token = _token.replace(p, "")
+ else:
+ _token = _token.replace(p, " ")
+ token = period_strip.sub("", _token, re.UNICODE)
+
+ _token = []
+ temp = token.lower().split()
+ for word in temp:
+ word = manual_map.setdefault(word, word)
+ if word not in articles:
+ _token.append(word)
+ for i, word in enumerate(_token):
+ if word in contractions:
+ _token[i] = contractions[word]
+ token = " ".join(_token)
+ token = token.replace(",", "")
+ return token
diff --git a/py/evf_sam/model/unilm/beit3/modeling_finetune.py b/py/evf_sam/model/unilm/beit3/modeling_finetune.py
new file mode 100644
index 0000000..dc5ea0a
--- /dev/null
+++ b/py/evf_sam/model/unilm/beit3/modeling_finetune.py
@@ -0,0 +1,386 @@
+# --------------------------------------------------------
+# Image as a Foreign Language: BEiT Pretraining for Vision and Vision-Language Tasks (https://arxiv.org/abs/2208.10442)
+# Github source: https://github.com/microsoft/unilm/tree/master/beit3
+# Copyright (c) 2023 Microsoft
+# Licensed under The MIT License [see LICENSE for details]
+# --------------------------------------------------------'
+
+import torch
+import torch.nn as nn
+import torch.nn.functional as F
+from timm.models import register_model
+import numpy as np
+
+import utils
+from modeling_utils import BEiT3Wrapper, _get_base_config, _get_large_config
+
+
+class TwoLayerMLP(nn.Module):
+ def __init__(
+ self,
+ in_features,
+ hidden_features,
+ out_features,
+ norm_layer,
+ norm_input=True,
+ ):
+ super().__init__()
+ self.norm1 = norm_layer(in_features) if norm_input else nn.Identity()
+ self.dense1 = nn.Linear(in_features, hidden_features)
+ self.norm2 = norm_layer(hidden_features)
+ self.act = nn.GELU()
+ self.dense2 = nn.Linear(hidden_features, out_features)
+
+ def forward(self, x):
+ x = self.norm1(x)
+ x = self.dense1(x)
+ x = self.norm2(x)
+ x = self.act(x)
+ return self.dense2(x)
+
+
+class Pooler(nn.Module):
+ def __init__(self, input_features, output_features, norm_layer):
+ super().__init__()
+ self.norm = norm_layer(input_features)
+ self.dense = nn.Linear(input_features, output_features)
+ self.activation = nn.Tanh()
+
+ def forward(self, x):
+ cls_rep = x[:, 0, :]
+ cls_rep = self.norm(cls_rep)
+ pooled_output = self.dense(cls_rep)
+ pooled_output = self.activation(pooled_output)
+ return pooled_output
+
+
+class BEiT3ForVisualReasoning(BEiT3Wrapper):
+ def __init__(
+ self,
+ args,
+ num_classes,
+ norm_layer=nn.LayerNorm,
+ **kwargs
+ ):
+ super(BEiT3ForVisualReasoning, self).__init__(args=args)
+ embed_dim = args.encoder_embed_dim
+ self.head = TwoLayerMLP(
+ in_features=embed_dim * 4,
+ hidden_features=embed_dim * 2,
+ out_features=num_classes,
+ norm_layer=norm_layer,
+ )
+ init_scale = 0.001
+ self.head.apply(self._init_weights)
+ if isinstance(self.head.dense1, nn.Linear):
+ self.head.dense1.weight.data.mul_(init_scale)
+ self.head.dense1.bias.data.mul_(init_scale)
+
+ if isinstance(self.head.dense2, nn.Linear):
+ self.head.dense2.weight.data.mul_(init_scale)
+ self.head.dense2.bias.data.mul_(init_scale)
+
+ def forward(self, image_a, image_b, text_description, padding_mask, **kwargs):
+ bsz, _ = text_description.size()
+
+ vision_input = torch.cat((image_a, image_b), dim=0)
+ language_input = torch.cat((text_description, text_description), dim=0)
+ padding_mask = torch.cat((padding_mask, padding_mask), dim=0)
+
+ outputs = self.beit3(
+ textual_tokens=language_input,
+ visual_tokens=vision_input,
+ text_padding_position=padding_mask,
+ )
+ x = outputs["encoder_out"]
+ multiway_split_position = outputs["multiway_split_position"]
+
+ vision_cls = x[:, 0, :]
+ language_cls = x[:, multiway_split_position, :]
+ cls_rep = torch.cat((vision_cls, language_cls), dim=-1)
+ a, b = torch.split(cls_rep, split_size_or_sections=[bsz, bsz], dim=0)
+ cls_rep = torch.cat((a, b), dim=-1)
+ return self.head(cls_rep)
+
+
+class BEiT3ForImageClassification(BEiT3Wrapper):
+ def __init__(
+ self,
+ args,
+ num_classes,
+ norm_layer=nn.LayerNorm,
+ **kwargs
+ ):
+ super(BEiT3ForImageClassification, self).__init__(args=args)
+ embed_dim = args.encoder_embed_dim
+ self.fc_norm = norm_layer(embed_dim)
+ self.head = nn.Linear(embed_dim, num_classes) if num_classes > 0 else nn.Identity()
+
+ self.fc_norm.apply(self._init_weights)
+ self.head.apply(self._init_weights)
+ init_scale = 0.001
+ if isinstance(self.head, nn.Linear):
+ self.head.weight.data.mul_(init_scale)
+ self.head.bias.data.mul_(init_scale)
+
+ def forward(self, image, **kwargs):
+ x = self.beit3(textual_tokens=None, visual_tokens=image)["encoder_out"]
+ t = x[:, 1:, :]
+ cls_x = self.fc_norm(t.mean(1))
+ return self.head(cls_x)
+
+
+class BEiT3ForCaptioning(BEiT3Wrapper):
+ def __init__(
+ self,
+ args,
+ **kwargs
+ ):
+ super(BEiT3ForCaptioning, self).__init__(args=args)
+ embed_dim = args.encoder_embed_dim
+ self.mlm_head = nn.Linear(embed_dim, args.vocab_size)
+ self.mlm_head.apply(self._init_weights)
+
+ def forward(self, image, text_ids, padding_mask, language_masked_pos, text_len=None, incremental_state=None, **kwargs):
+ text_len = text_len if text_len is not None else text_ids.size(1)
+ image_len = self.beit3.vision_embed.num_position_embeddings()
+ max_len = text_len + image_len
+ uni_mask = torch.zeros((max_len, max_len), dtype=torch.long, device=text_ids.device)
+ i_start, i_end = 0, image_len
+ t_start, t_end = image_len, max_len
+ # triangle mask for caption to caption
+ uni_mask[t_start:t_end, t_start:t_end] = torch.tril(torch.ones(text_len, text_len, dtype=torch.long, device=text_ids.device))
+ # full attention for caption to image
+ uni_mask[t_start:t_end, i_start:i_end] = 1
+ # full attention for image to image
+ uni_mask[i_start:i_end, i_start:i_end] = 1
+ uni_mask = 1-uni_mask
+
+ if incremental_state is not None:
+ for idx in range(self.get_num_layers()):
+ if idx not in incremental_state:
+ incremental_state[idx] = {}
+
+ # for incremental decoding
+ positions = None
+ if image is None:
+ uni_mask = uni_mask[-2:]
+ padding_mask = None
+ # start position (2 (fairseq starts at 2) + cur_position) is equal to text_len
+ positions = torch.arange(text_len, text_ids.size(1) + text_len, device=text_ids.device).long().unsqueeze(0)
+
+ outputs = self.beit3(
+ textual_tokens=text_ids,
+ visual_tokens=image,
+ text_padding_position=padding_mask,
+ attn_mask=uni_mask,
+ incremental_state=incremental_state,
+ positions=positions,
+ )
+ if image is not None:
+ text_feats = outputs["encoder_out"][:, image_len:]
+ else:
+ text_feats = outputs["encoder_out"]
+
+ if language_masked_pos is not None:
+ text_feats = text_feats[language_masked_pos.bool()]
+
+ return self.mlm_head(text_feats), incremental_state
+
+
+class BEiT3ForVisualQuestionAnswering(BEiT3Wrapper):
+ def __init__(
+ self,
+ args,
+ num_classes,
+ norm_layer=nn.LayerNorm,
+ **kwargs
+ ):
+ super(BEiT3ForVisualQuestionAnswering, self).__init__(args=args)
+ embed_dim = args.encoder_embed_dim
+ self.pooler = Pooler(
+ input_features=embed_dim,
+ output_features=embed_dim,
+ norm_layer=norm_layer,
+ )
+ self.pooler.apply(self._init_weights)
+ self.head = nn.Sequential(
+ nn.Linear(embed_dim, embed_dim * 2),
+ norm_layer(embed_dim * 2),
+ nn.GELU(),
+ nn.Linear(embed_dim * 2, num_classes),
+ )
+ self.head.apply(self._init_weights)
+
+ def forward(self, image, question, padding_mask, **kwargs):
+ outputs = self.beit3(
+ textual_tokens=question,
+ visual_tokens=image,
+ text_padding_position=padding_mask,
+ )
+ x = outputs["encoder_out"]
+ cls_rep = self.pooler(x)
+ return self.head(cls_rep)
+
+
+class BEiT3ForRetrieval(BEiT3Wrapper):
+ def __init__(
+ self,
+ args,
+ **kwargs
+ ):
+ super(BEiT3ForRetrieval, self).__init__(args=args)
+ embed_dim = args.encoder_embed_dim
+ self.language_head = nn.Linear(embed_dim, embed_dim, bias=False)
+ self.vision_head = nn.Linear(embed_dim, embed_dim, bias=False)
+ self.language_head.apply(self._init_weights)
+ self.vision_head.apply(self._init_weights)
+ self.criterion = utils.ClipLoss(
+ rank=utils.get_rank(),
+ world_size=utils.get_world_size(),
+ )
+ self.logit_scale = nn.Parameter(torch.ones([]) * np.log(1 / 0.07))
+
+ def forward(self, image=None, text_description=None, padding_mask=None, only_infer=False, **kwargs):
+ if image is not None:
+ outputs = self.beit3(
+ textual_tokens=None,
+ visual_tokens=image,
+ text_padding_position=None,
+ )
+ x = outputs["encoder_out"]
+ vision_cls = self.vision_head(x[:, 0, :])
+ vision_cls = F.normalize(vision_cls, dim=-1)
+ else:
+ vision_cls = None
+
+ if text_description is not None:
+ outputs = self.beit3(
+ textual_tokens=text_description,
+ visual_tokens=None,
+ text_padding_position=padding_mask,
+ )
+ x = outputs["encoder_out"]
+ language_cls = self.language_head(x[:, 0, :])
+ language_cls = F.normalize(language_cls, dim=-1)
+ else:
+ language_cls = None
+
+ if only_infer:
+ return vision_cls, language_cls
+ else:
+ loss, logits_per_image, logits_per_text = self.criterion(
+ vision_cls, language_cls, self.logit_scale.exp())
+ return loss, vision_cls, language_cls
+
+
+@register_model
+def beit3_base_patch16_224_imageclassification(pretrained=False, **kwargs):
+ args = _get_base_config(**kwargs)
+ args.normalize_output = False
+ model = BEiT3ForImageClassification(args, num_classes=1000, **kwargs)
+ return model
+
+
+@register_model
+def beit3_large_patch16_224_imageclassification(pretrained=False, **kwargs):
+ args = _get_large_config(**kwargs)
+ args.normalize_output = False
+ model = BEiT3ForImageClassification(args, num_classes=1000, **kwargs)
+ return model
+
+
+@register_model
+def beit3_base_patch16_224_nlvr2(pretrained=False, **kwargs):
+ args = _get_base_config(**kwargs)
+ model = BEiT3ForVisualReasoning(args, num_classes=2, **kwargs)
+ return model
+
+
+@register_model
+def beit3_large_patch16_224_nlvr2(pretrained=False, **kwargs):
+ args = _get_large_config(**kwargs)
+ model = BEiT3ForVisualReasoning(args, num_classes=2, **kwargs)
+ return model
+
+
+@register_model
+def beit3_base_patch16_384_vqav2(pretrained=False, **kwargs):
+ args = _get_base_config(img_size=384, **kwargs)
+ args.normalize_output = False
+ model = BEiT3ForVisualQuestionAnswering(args, num_classes=3129, **kwargs)
+ return model
+
+
+@register_model
+def beit3_base_patch16_480_vqav2(pretrained=False, **kwargs):
+ args = _get_base_config(img_size=480, **kwargs)
+ args.normalize_output = False
+ model = BEiT3ForVisualQuestionAnswering(args, num_classes=3129, **kwargs)
+ return model
+
+
+@register_model
+def beit3_large_patch16_384_vqav2(pretrained=False, **kwargs):
+ args = _get_large_config(img_size=384, **kwargs)
+ args.normalize_output = False
+ model = BEiT3ForVisualQuestionAnswering(args, num_classes=3129, **kwargs)
+ return model
+
+
+@register_model
+def beit3_large_patch16_480_vqav2(pretrained=False, **kwargs):
+ args = _get_large_config(img_size=480, **kwargs)
+ args.normalize_output = False
+ model = BEiT3ForVisualQuestionAnswering(args, num_classes=3129, **kwargs)
+ return model
+
+
+@register_model
+def beit3_large_patch16_768_vqav2(pretrained=False, **kwargs):
+ args = _get_large_config(img_size=768, **kwargs)
+ args.normalize_output = False
+ model = BEiT3ForVisualQuestionAnswering(args, num_classes=3129, **kwargs)
+ return model
+
+
+@register_model
+def beit3_base_patch16_224_captioning(pretrained=False, **kwargs):
+ args = _get_base_config(**kwargs)
+ model = BEiT3ForCaptioning(args, **kwargs)
+ return model
+
+
+@register_model
+def beit3_base_patch16_480_captioning(pretrained=False, **kwargs):
+ args = _get_base_config(img_size=480, **kwargs)
+ model = BEiT3ForCaptioning(args, **kwargs)
+ return model
+
+
+@register_model
+def beit3_large_patch16_480_captioning(pretrained=False, **kwargs):
+ args = _get_large_config(img_size=480, **kwargs)
+ model = BEiT3ForCaptioning(args, **kwargs)
+ return model
+
+
+@register_model
+def beit3_base_patch16_224_retrieval(pretrained=False, **kwargs):
+ args = _get_base_config(**kwargs)
+ model = BEiT3ForRetrieval(args, **kwargs)
+ return model
+
+
+@register_model
+def beit3_base_patch16_384_retrieval(pretrained=False, **kwargs):
+ args = _get_base_config(img_size=384, **kwargs)
+ model = BEiT3ForRetrieval(args, **kwargs)
+ return model
+
+
+@register_model
+def beit3_large_patch16_384_retrieval(pretrained=False, **kwargs):
+ args = _get_large_config(img_size=384, **kwargs)
+ model = BEiT3ForRetrieval(args, **kwargs)
+ return model
diff --git a/py/evf_sam/model/unilm/beit3/modeling_utils.py b/py/evf_sam/model/unilm/beit3/modeling_utils.py
new file mode 100644
index 0000000..6581159
--- /dev/null
+++ b/py/evf_sam/model/unilm/beit3/modeling_utils.py
@@ -0,0 +1,76 @@
+# --------------------------------------------------------
+# Image as a Foreign Language: BEiT Pretraining for Vision and Vision-Language Tasks (https://arxiv.org/abs/2208.10442)
+# Github source: https://github.com/microsoft/unilm/tree/master/beit3
+# Copyright (c) 2023 Microsoft
+# Licensed under The MIT License [see LICENSE for details]
+# --------------------------------------------------------'
+
+import math
+import torch
+import torch.nn as nn
+from timm.models.layers import trunc_normal_ as __call_trunc_normal_
+
+from torchscale.model.BEiT3 import BEiT3
+from torchscale.architecture.config import EncoderConfig
+
+
+def trunc_normal_(tensor, mean=0., std=1.):
+ __call_trunc_normal_(tensor, mean=mean, std=std, a=-std, b=std)
+
+
+def _get_base_config(
+ img_size=224, patch_size=16, drop_path_rate=0,
+ checkpoint_activations=None, mlp_ratio=4, vocab_size=64010, **kwargs
+):
+ return EncoderConfig(
+ img_size=img_size, patch_size=patch_size, vocab_size=vocab_size, multiway=True,
+ layernorm_embedding=False, normalize_output=True, no_output_layer=True,
+ drop_path_rate=drop_path_rate, encoder_embed_dim=768, encoder_attention_heads=12,
+ encoder_ffn_embed_dim=int(768 * mlp_ratio), encoder_layers=12,
+ checkpoint_activations=checkpoint_activations,
+ )
+
+
+def _get_large_config(
+ img_size=224, patch_size=16, drop_path_rate=0,
+ checkpoint_activations=None, mlp_ratio=4, vocab_size=64010, **kwargs
+):
+ return EncoderConfig(
+ img_size=img_size, patch_size=patch_size, vocab_size=vocab_size, multiway=True,
+ layernorm_embedding=False, normalize_output=True, no_output_layer=True,
+ drop_path_rate=drop_path_rate, encoder_embed_dim=1024, encoder_attention_heads=16,
+ encoder_ffn_embed_dim=int(1024 * mlp_ratio), encoder_layers=24,
+ checkpoint_activations=checkpoint_activations,
+ )
+
+
+class BEiT3Wrapper(nn.Module):
+ def __init__(self, args, **kwargs):
+ super().__init__()
+ self.args = args
+ self.beit3 = BEiT3(args)
+ self.apply(self._init_weights)
+
+ def fix_init_weight(self):
+ def rescale(param, layer_id):
+ param.div_(math.sqrt(2.0 * layer_id))
+
+ for layer_id, layer in enumerate(self.blocks):
+ rescale(layer.attn.proj.weight.data, layer_id + 1)
+ rescale(layer.mlp.fc2.weight.data, layer_id + 1)
+
+ def get_num_layers(self):
+ return self.beit3.encoder.num_layers
+
+ @torch.jit.ignore
+ def no_weight_decay(self):
+ return {'pos_embed', 'cls_token', 'beit3.encoder.embed_positions.A.weight', 'beit3.vision_embed.cls_token', 'logit_scale'}
+
+ def _init_weights(self, m):
+ if isinstance(m, nn.Linear):
+ trunc_normal_(m.weight, std=.02)
+ if isinstance(m, nn.Linear) and m.bias is not None:
+ nn.init.constant_(m.bias, 0)
+ elif isinstance(m, nn.LayerNorm):
+ nn.init.constant_(m.bias, 0)
+ nn.init.constant_(m.weight, 1.0)
diff --git a/py/evf_sam/model/unilm/beit3/optim_factory.py b/py/evf_sam/model/unilm/beit3/optim_factory.py
new file mode 100644
index 0000000..a3d32bb
--- /dev/null
+++ b/py/evf_sam/model/unilm/beit3/optim_factory.py
@@ -0,0 +1,128 @@
+# --------------------------------------------------------
+# Image as a Foreign Language: BEiT Pretraining for Vision and Vision-Language Tasks (https://arxiv.org/abs/2208.10442)
+# Github source: https://github.com/microsoft/unilm/tree/master/beit3
+# Copyright (c) 2023 Microsoft
+# Licensed under The MIT License [see LICENSE for details]
+# --------------------------------------------------------'
+
+from torch import optim as optim
+from timm.optim.lookahead import Lookahead
+
+import json
+
+
+def get_num_layer_for_vit(var_name, num_max_layer):
+ if "embed" in var_name:
+ return 0
+ elif var_name in (
+ "cls_token", "mask_token", "pos_embed", "language_pos_embed",
+ "word_embeddings.weight", "vision_cls_token", "vision_pos_embed"
+ ):
+ return 0
+ elif var_name.startswith("patch_embed"):
+ return 0
+ elif var_name.startswith("rel_pos_bias"):
+ return num_max_layer - 1
+ elif "layers." in var_name:
+ layer_id = int(var_name.split('layers.')[1].split('.')[0])
+ return layer_id + 1
+ else:
+ return num_max_layer - 1
+
+
+def get_is_head_flag_for_vit(var_name, num_max_layer):
+ if var_name.startswith("head"):
+ return 1
+ # elif var_name.startswith("pooler"):
+ # return 1
+ else:
+ return 0
+
+
+class LayerDecayValueAssigner(object):
+ def __init__(self, values, scale_handler=None):
+ self.scale_handler = scale_handler or get_num_layer_for_vit
+ self.values = values
+
+ def get_scale(self, layer_id):
+ return self.values[layer_id]
+
+ def get_layer_id(self, var_name):
+ return self.scale_handler(var_name, len(self.values))
+
+
+# The implementation code is modified from Timm (https://github.com/huggingface/pytorch-image-models/tree/main/timm
+def get_parameter_groups(model, weight_decay=1e-5, skip_list=(), get_num_layer=None, get_layer_scale=None):
+ parameter_group_names = {}
+ parameter_group_vars = {}
+
+ for name, param in model.named_parameters():
+ if not param.requires_grad:
+ continue # frozen weights
+ if len(param.shape) == 1 or name.endswith(".bias") or name in skip_list:
+ group_name = "no_decay"
+ this_weight_decay = 0.
+ else:
+ group_name = "decay"
+ this_weight_decay = weight_decay
+ if get_num_layer is not None:
+ layer_id = get_num_layer(name)
+ group_name = "layer_%d_%s" % (layer_id, group_name)
+ else:
+ layer_id = None
+
+ if group_name not in parameter_group_names:
+ if get_layer_scale is not None:
+ scale = get_layer_scale(layer_id)
+ else:
+ scale = 1.
+
+ parameter_group_names[group_name] = {
+ "weight_decay": this_weight_decay,
+ "params": [],
+ "lr_scale": scale
+ }
+ parameter_group_vars[group_name] = {
+ "weight_decay": this_weight_decay,
+ "params": [],
+ "lr_scale": scale
+ }
+
+ parameter_group_vars[group_name]["params"].append(param)
+ parameter_group_names[group_name]["params"].append(name)
+ print("Param groups = %s" % json.dumps(parameter_group_names, indent=2))
+ return list(parameter_group_vars.values())
+
+
+def create_optimizer(args, model, get_num_layer=None, get_layer_scale=None, filter_bias_and_bn=True, skip_list=None):
+ opt_lower = args.opt.lower()
+ weight_decay = args.weight_decay
+ if weight_decay and filter_bias_and_bn:
+ skip = {}
+ if skip_list is not None:
+ skip = skip_list
+ elif hasattr(model, 'no_weight_decay'):
+ skip = model.no_weight_decay()
+ parameters = get_parameter_groups(model, weight_decay, skip, get_num_layer, get_layer_scale)
+ weight_decay = 0.
+ else:
+ parameters = model.parameters()
+
+ opt_args = dict(lr=args.lr, weight_decay=weight_decay)
+ if hasattr(args, 'opt_eps') and args.opt_eps is not None:
+ opt_args['eps'] = args.opt_eps
+ if hasattr(args, 'opt_betas') and args.opt_betas is not None:
+ opt_args['betas'] = args.opt_betas
+
+ opt_split = opt_lower.split('_')
+ opt_lower = opt_split[-1]
+ if opt_lower == 'adamw':
+ optimizer = optim.AdamW(parameters, **opt_args)
+ else:
+ raise ValueError("Invalid optimizer")
+
+ if len(opt_split) > 1:
+ if opt_split[0] == 'lookahead':
+ optimizer = Lookahead(optimizer)
+
+ return optimizer
diff --git a/py/evf_sam/model/unilm/beit3/randaug.py b/py/evf_sam/model/unilm/beit3/randaug.py
new file mode 100644
index 0000000..359e97d
--- /dev/null
+++ b/py/evf_sam/model/unilm/beit3/randaug.py
@@ -0,0 +1,340 @@
+import cv2
+import numpy as np
+
+
+## aug functions
+def identity_func(img):
+ return img
+
+
+def autocontrast_func(img, cutoff=0):
+ '''
+ same output as PIL.ImageOps.autocontrast
+ '''
+ n_bins = 256
+
+ def tune_channel(ch):
+ n = ch.size
+ cut = cutoff * n // 100
+ if cut == 0:
+ high, low = ch.max(), ch.min()
+ else:
+ hist = cv2.calcHist([ch], [0], None, [n_bins], [0, n_bins])
+ low = np.argwhere(np.cumsum(hist) > cut)
+ low = 0 if low.shape[0] == 0 else low[0]
+ high = np.argwhere(np.cumsum(hist[::-1]) > cut)
+ high = n_bins - 1 if high.shape[0] == 0 else n_bins - 1 - high[0]
+ if high <= low:
+ table = np.arange(n_bins)
+ else:
+ scale = (n_bins - 1) / (high - low)
+ offset = -low * scale
+ table = np.arange(n_bins) * scale + offset
+ table[table < 0] = 0
+ table[table > n_bins - 1] = n_bins - 1
+ table = table.clip(0, 255).astype(np.uint8)
+ return table[ch]
+
+ channels = [tune_channel(ch) for ch in cv2.split(img)]
+ out = cv2.merge(channels)
+ return out
+
+
+def equalize_func(img):
+ '''
+ same output as PIL.ImageOps.equalize
+ PIL's implementation is different from cv2.equalize
+ '''
+ n_bins = 256
+
+ def tune_channel(ch):
+ hist = cv2.calcHist([ch], [0], None, [n_bins], [0, n_bins])
+ non_zero_hist = hist[hist != 0].reshape(-1)
+ step = np.sum(non_zero_hist[:-1]) // (n_bins - 1)
+ if step == 0: return ch
+ n = np.empty_like(hist)
+ n[0] = step // 2
+ n[1:] = hist[:-1]
+ table = (np.cumsum(n) // step).clip(0, 255).astype(np.uint8)
+ return table[ch]
+
+ channels = [tune_channel(ch) for ch in cv2.split(img)]
+ out = cv2.merge(channels)
+ return out
+
+
+def rotate_func(img, degree, fill=(0, 0, 0)):
+ '''
+ like PIL, rotate by degree, not radians
+ '''
+ H, W = img.shape[0], img.shape[1]
+ center = W / 2, H / 2
+ M = cv2.getRotationMatrix2D(center, degree, 1)
+ out = cv2.warpAffine(img, M, (W, H), borderValue=fill)
+ return out
+
+
+def solarize_func(img, thresh=128):
+ '''
+ same output as PIL.ImageOps.posterize
+ '''
+ table = np.array([el if el < thresh else 255 - el for el in range(256)])
+ table = table.clip(0, 255).astype(np.uint8)
+ out = table[img]
+ return out
+
+
+def color_func(img, factor):
+ '''
+ same output as PIL.ImageEnhance.Color
+ '''
+ ## implementation according to PIL definition, quite slow
+ # degenerate = cv2.cvtColor(img, cv2.COLOR_BGR2GRAY)[:, :, np.newaxis]
+ # out = blend(degenerate, img, factor)
+ # M = (
+ # np.eye(3) * factor
+ # + np.float32([0.114, 0.587, 0.299]).reshape(3, 1) * (1. - factor)
+ # )[np.newaxis, np.newaxis, :]
+ M = (
+ np.float32([
+ [0.886, -0.114, -0.114],
+ [-0.587, 0.413, -0.587],
+ [-0.299, -0.299, 0.701]]) * factor
+ + np.float32([[0.114], [0.587], [0.299]])
+ )
+ out = np.matmul(img, M).clip(0, 255).astype(np.uint8)
+ return out
+
+
+def contrast_func(img, factor):
+ """
+ same output as PIL.ImageEnhance.Contrast
+ """
+ mean = np.sum(np.mean(img, axis=(0, 1)) * np.array([0.114, 0.587, 0.299]))
+ table = np.array([(
+ el - mean) * factor + mean
+ for el in range(256)
+ ]).clip(0, 255).astype(np.uint8)
+ out = table[img]
+ return out
+
+
+def brightness_func(img, factor):
+ '''
+ same output as PIL.ImageEnhance.Contrast
+ '''
+ table = (np.arange(256, dtype=np.float32) * factor).clip(0, 255).astype(np.uint8)
+ out = table[img]
+ return out
+
+
+def sharpness_func(img, factor):
+ '''
+ The differences the this result and PIL are all on the 4 boundaries, the center
+ areas are same
+ '''
+ kernel = np.ones((3, 3), dtype=np.float32)
+ kernel[1][1] = 5
+ kernel /= 13
+ degenerate = cv2.filter2D(img, -1, kernel)
+ if factor == 0.0:
+ out = degenerate
+ elif factor == 1.0:
+ out = img
+ else:
+ out = img.astype(np.float32)
+ degenerate = degenerate.astype(np.float32)[1:-1, 1:-1, :]
+ out[1:-1, 1:-1, :] = degenerate + factor * (out[1:-1, 1:-1, :] - degenerate)
+ out = out.astype(np.uint8)
+ return out
+
+
+def shear_x_func(img, factor, fill=(0, 0, 0)):
+ H, W = img.shape[0], img.shape[1]
+ M = np.float32([[1, factor, 0], [0, 1, 0]])
+ out = cv2.warpAffine(img, M, (W, H), borderValue=fill, flags=cv2.INTER_LINEAR).astype(np.uint8)
+ return out
+
+
+def translate_x_func(img, offset, fill=(0, 0, 0)):
+ '''
+ same output as PIL.Image.transform
+ '''
+ H, W = img.shape[0], img.shape[1]
+ M = np.float32([[1, 0, -offset], [0, 1, 0]])
+ out = cv2.warpAffine(img, M, (W, H), borderValue=fill, flags=cv2.INTER_LINEAR).astype(np.uint8)
+ return out
+
+
+def translate_y_func(img, offset, fill=(0, 0, 0)):
+ '''
+ same output as PIL.Image.transform
+ '''
+ H, W = img.shape[0], img.shape[1]
+ M = np.float32([[1, 0, 0], [0, 1, -offset]])
+ out = cv2.warpAffine(img, M, (W, H), borderValue=fill, flags=cv2.INTER_LINEAR).astype(np.uint8)
+ return out
+
+
+def posterize_func(img, bits):
+ '''
+ same output as PIL.ImageOps.posterize
+ '''
+ out = np.bitwise_and(img, np.uint8(255 << (8 - bits)))
+ return out
+
+
+def shear_y_func(img, factor, fill=(0, 0, 0)):
+ H, W = img.shape[0], img.shape[1]
+ M = np.float32([[1, 0, 0], [factor, 1, 0]])
+ out = cv2.warpAffine(img, M, (W, H), borderValue=fill, flags=cv2.INTER_LINEAR).astype(np.uint8)
+ return out
+
+
+def cutout_func(img, pad_size, replace=(0, 0, 0)):
+ replace = np.array(replace, dtype=np.uint8)
+ H, W = img.shape[0], img.shape[1]
+ rh, rw = np.random.random(2)
+ pad_size = pad_size // 2
+ ch, cw = int(rh * H), int(rw * W)
+ x1, x2 = max(ch - pad_size, 0), min(ch + pad_size, H)
+ y1, y2 = max(cw - pad_size, 0), min(cw + pad_size, W)
+ out = img.copy()
+ out[x1:x2, y1:y2, :] = replace
+ return out
+
+
+### level to args
+def enhance_level_to_args(MAX_LEVEL):
+ def level_to_args(level):
+ return ((level / MAX_LEVEL) * 1.8 + 0.1,)
+ return level_to_args
+
+
+def shear_level_to_args(MAX_LEVEL, replace_value):
+ def level_to_args(level):
+ level = (level / MAX_LEVEL) * 0.3
+ if np.random.random() > 0.5: level = -level
+ return (level, replace_value)
+
+ return level_to_args
+
+
+def translate_level_to_args(translate_const, MAX_LEVEL, replace_value):
+ def level_to_args(level):
+ level = (level / MAX_LEVEL) * float(translate_const)
+ if np.random.random() > 0.5: level = -level
+ return (level, replace_value)
+
+ return level_to_args
+
+
+def cutout_level_to_args(cutout_const, MAX_LEVEL, replace_value):
+ def level_to_args(level):
+ level = int((level / MAX_LEVEL) * cutout_const)
+ return (level, replace_value)
+
+ return level_to_args
+
+
+def solarize_level_to_args(MAX_LEVEL):
+ def level_to_args(level):
+ level = int((level / MAX_LEVEL) * 256)
+ return (level, )
+ return level_to_args
+
+
+def none_level_to_args(level):
+ return ()
+
+
+def posterize_level_to_args(MAX_LEVEL):
+ def level_to_args(level):
+ level = int((level / MAX_LEVEL) * 4)
+ return (level, )
+ return level_to_args
+
+
+def rotate_level_to_args(MAX_LEVEL, replace_value):
+ def level_to_args(level):
+ level = (level / MAX_LEVEL) * 30
+ if np.random.random() < 0.5:
+ level = -level
+ return (level, replace_value)
+
+ return level_to_args
+
+
+func_dict = {
+ 'Identity': identity_func,
+ 'AutoContrast': autocontrast_func,
+ 'Equalize': equalize_func,
+ 'Rotate': rotate_func,
+ 'Solarize': solarize_func,
+ 'Color': color_func,
+ 'Contrast': contrast_func,
+ 'Brightness': brightness_func,
+ 'Sharpness': sharpness_func,
+ 'ShearX': shear_x_func,
+ 'TranslateX': translate_x_func,
+ 'TranslateY': translate_y_func,
+ 'Posterize': posterize_func,
+ 'ShearY': shear_y_func,
+}
+
+translate_const = 10
+MAX_LEVEL = 10
+replace_value = (128, 128, 128)
+arg_dict = {
+ 'Identity': none_level_to_args,
+ 'AutoContrast': none_level_to_args,
+ 'Equalize': none_level_to_args,
+ 'Rotate': rotate_level_to_args(MAX_LEVEL, replace_value),
+ 'Solarize': solarize_level_to_args(MAX_LEVEL),
+ 'Color': enhance_level_to_args(MAX_LEVEL),
+ 'Contrast': enhance_level_to_args(MAX_LEVEL),
+ 'Brightness': enhance_level_to_args(MAX_LEVEL),
+ 'Sharpness': enhance_level_to_args(MAX_LEVEL),
+ 'ShearX': shear_level_to_args(MAX_LEVEL, replace_value),
+ 'TranslateX': translate_level_to_args(
+ translate_const, MAX_LEVEL, replace_value
+ ),
+ 'TranslateY': translate_level_to_args(
+ translate_const, MAX_LEVEL, replace_value
+ ),
+ 'Posterize': posterize_level_to_args(MAX_LEVEL),
+ 'ShearY': shear_level_to_args(MAX_LEVEL, replace_value),
+}
+
+
+class RandomAugment(object):
+
+ def __init__(self, N=2, M=10, isPIL=False, augs=[]):
+ self.N = N
+ self.M = M
+ self.isPIL = isPIL
+ if augs:
+ self.augs = augs
+ else:
+ self.augs = list(arg_dict.keys())
+
+ def get_random_ops(self):
+ sampled_ops = np.random.choice(self.augs, self.N)
+ return [(op, 0.5, self.M) for op in sampled_ops]
+
+ def __call__(self, img):
+ if self.isPIL:
+ img = np.array(img)
+ ops = self.get_random_ops()
+ for name, prob, level in ops:
+ if np.random.random() > prob:
+ continue
+ args = arg_dict[name](level)
+ img = func_dict[name](img, *args)
+ return img
+
+
+if __name__ == '__main__':
+ a = RandomAugment()
+ img = np.random.randn(32, 32, 3)
+ a(img)
diff --git a/py/evf_sam/model/unilm/beit3/requirements.txt b/py/evf_sam/model/unilm/beit3/requirements.txt
new file mode 100644
index 0000000..4d8ffa5
--- /dev/null
+++ b/py/evf_sam/model/unilm/beit3/requirements.txt
@@ -0,0 +1,22 @@
+torch
+torchvision
+timm==0.4.12
+Pillow
+blobfile
+mypy
+numpy
+pytest
+requests
+einops
+tensorboardX
+scipy
+ftfy
+opencv-python
+sentencepiece
+pyarrow
+torchmetrics==0.7.3
+transformers
+deepspeed==0.4.0
+pycocotools
+pycocoevalcap
+torchscale==0.2.0
diff --git a/py/evf_sam/model/unilm/beit3/run_beit3_finetuning.py b/py/evf_sam/model/unilm/beit3/run_beit3_finetuning.py
new file mode 100644
index 0000000..758cd69
--- /dev/null
+++ b/py/evf_sam/model/unilm/beit3/run_beit3_finetuning.py
@@ -0,0 +1,448 @@
+# --------------------------------------------------------
+# Image as a Foreign Language: BEiT Pretraining for Vision and Vision-Language Tasks (https://arxiv.org/abs/2208.10442)
+# Github source: https://github.com/microsoft/unilm/tree/master/beit3
+# Copyright (c) 2023 Microsoft
+# Licensed under The MIT License [see LICENSE for details]
+# --------------------------------------------------------'
+
+import argparse
+import datetime
+import numpy as np
+import time
+import torch
+import torch.backends.cudnn as cudnn
+import json
+import os
+
+from pathlib import Path
+
+from timm.data.mixup import Mixup
+from timm.models import create_model
+from timm.utils import ModelEma
+from optim_factory import create_optimizer, get_parameter_groups, \
+ LayerDecayValueAssigner, get_is_head_flag_for_vit
+
+from engine_for_finetuning import train_one_epoch, get_handler, evaluate
+from datasets import create_downstream_dataset
+from utils import NativeScalerWithGradNormCount as NativeScaler
+import utils
+import modeling_finetune
+
+
+def get_args():
+ parser = argparse.ArgumentParser('BEiT fine-tuning and evaluation script for image classification', add_help=False)
+
+ # Model parameters
+ parser.add_argument('--model', default='beit_base_patch16_224', type=str, metavar='MODEL',
+ help='Name of model to train')
+ parser.add_argument('--task', type=str, required=True,
+ choices=['nlvr2', 'vqav2', 'flickr30k', 'coco_retrieval', 'coco_captioning', 'nocaps', 'imagenet'],
+ help='Name of task to fine-tuning')
+
+ parser.add_argument('--input_size', default=224, type=int,
+ help='images input size')
+ parser.add_argument('--drop_path', type=float, default=0.1, metavar='PCT',
+ help='Drop path rate (default: 0.1)')
+
+ parser.add_argument('--checkpoint_activations', action='store_true', default=None,
+ help='Enable checkpointing to save your memory.')
+ parser.add_argument('--sentencepiece_model', type=str, required=True,
+ help='Sentencepiece model path for the pretrained model.')
+ parser.add_argument('--vocab_size', type=int, default=64010)
+ parser.add_argument('--num_max_bpe_tokens', type=int, default=64)
+
+ parser.add_argument('--model_ema', action='store_true', default=False)
+ parser.add_argument('--model_ema_decay', type=float, default=0.9999, help='')
+ parser.add_argument('--model_ema_force_cpu', action='store_true', default=False, help='')
+
+ # Optimizer parameters
+ parser.add_argument('--opt', default='adamw', type=str, metavar='OPTIMIZER',
+ help='Optimizer (default: "adamw"')
+ parser.add_argument('--opt_eps', default=1e-8, type=float, metavar='EPSILON',
+ help='Optimizer Epsilon (default: 1e-8)')
+ parser.add_argument('--opt_betas', default=[0.9, 0.999], type=float, nargs='+', metavar='BETA',
+ help='Optimizer Betas (default: 0.9, 0.999, use opt default)')
+ parser.add_argument('--clip_grad', type=float, default=None, metavar='NORM',
+ help='Clip gradient norm (default: None, no clipping)')
+ parser.add_argument('--momentum', type=float, default=0.9, metavar='M',
+ help='SGD momentum (default: 0.9)')
+ parser.add_argument('--weight_decay', type=float, default=0.05,
+ help='weight decay (default: 0.05)')
+
+ parser.add_argument('--lr', type=float, default=5e-4, metavar='LR',
+ help='learning rate (default: 5e-4)')
+ parser.add_argument('--layer_decay', type=float, default=0.9)
+ parser.add_argument('--task_head_lr_weight', type=float, default=0)
+
+ parser.add_argument('--warmup_lr', type=float, default=1e-6, metavar='LR',
+ help='warmup learning rate (default: 1e-6)')
+ parser.add_argument('--min_lr', type=float, default=1e-6, metavar='LR',
+ help='lower lr bound for cyclic schedulers that hit 0 (1e-6)')
+ parser.add_argument('--warmup_epochs', type=int, default=5, metavar='N',
+ help='epochs to warmup LR, if scheduler supports')
+ parser.add_argument('--warmup_steps', type=int, default=-1, metavar='N',
+ help='num of steps to warmup LR, will overload warmup_epochs if set > 0')
+
+ parser.add_argument('--batch_size', default=64, type=int)
+ parser.add_argument('--eval_batch_size', default=None, type=int)
+ parser.add_argument('--epochs', default=20, type=int)
+ parser.add_argument('--update_freq', default=1, type=int)
+ parser.add_argument('--save_ckpt_freq', default=5, type=int)
+
+ # Augmentation parameters
+ parser.add_argument('--randaug', action='store_true', default=False)
+ parser.add_argument('--train_interpolation', type=str, default='bicubic',
+ help='Training interpolation (random, bilinear, bicubic default: "bicubic")')
+
+ # Finetuning params
+ parser.add_argument('--finetune', default='',
+ help='finetune from checkpoint')
+ parser.add_argument('--model_key', default='model|module', type=str)
+ parser.add_argument('--model_prefix', default='', type=str)
+
+ # Dataset parameters
+ parser.add_argument('--data_path', default='/datasets01/imagenet_full_size/061417/', type=str,
+ help='dataset path')
+
+ parser.add_argument('--output_dir', default='',
+ help='path where to save, empty for no saving')
+ parser.add_argument('--log_dir', default=None,
+ help='path where to tensorboard log')
+ parser.add_argument('--device', default='cuda',
+ help='device to use for training / testing')
+ parser.add_argument('--seed', default=0, type=int)
+ parser.add_argument('--resume', default='',
+ help='resume from checkpoint')
+ parser.add_argument('--auto_resume', action='store_true')
+ parser.add_argument('--no_auto_resume', action='store_false', dest='auto_resume')
+ parser.set_defaults(auto_resume=True)
+
+ parser.add_argument('--save_ckpt', action='store_true')
+ parser.add_argument('--no_save_ckpt', action='store_false', dest='save_ckpt')
+ parser.set_defaults(save_ckpt=True)
+
+ parser.add_argument('--start_epoch', default=0, type=int, metavar='N',
+ help='start epoch')
+ parser.add_argument('--eval', action='store_true',
+ help='Perform evaluation only')
+ parser.add_argument('--dist_eval', action='store_true', default=False,
+ help='Enabling distributed evaluation')
+ parser.add_argument('--num_workers', default=10, type=int)
+ parser.add_argument('--pin_mem', action='store_true',
+ help='Pin CPU memory in DataLoader for more efficient (sometimes) transfer to GPU.')
+ parser.add_argument('--no_pin_mem', action='store_false', dest='pin_mem')
+ parser.set_defaults(pin_mem=True)
+
+ # distributed training parameters
+ parser.add_argument('--world_size', default=1, type=int,
+ help='number of distributed processes')
+ parser.add_argument('--local_rank', default=-1, type=int)
+ parser.add_argument('--dist_on_itp', action='store_true')
+ parser.add_argument('--dist_url', default='env://',
+ help='url used to set up distributed training')
+
+ # parameter for dump predictions (VQA, COCO captioning, NoCaps)
+ parser.add_argument('--task_cache_path', default=None, type=str)
+
+ # parameter for imagenet finetuning
+ parser.add_argument('--nb_classes', default=1000, type=int,
+ help='number of the classification types')
+ parser.add_argument('--mixup', type=float, default=0,
+ help='mixup alpha, mixup enabled if > 0.')
+ parser.add_argument('--cutmix', type=float, default=0,
+ help='cutmix alpha, cutmix enabled if > 0.')
+ parser.add_argument('--cutmix_minmax', type=float, nargs='+', default=None,
+ help='cutmix min/max ratio, overrides alpha and enables cutmix if set (default: None)')
+ parser.add_argument('--mixup_prob', type=float, default=1.0,
+ help='Probability of performing mixup or cutmix when either/both is enabled')
+ parser.add_argument('--mixup_switch_prob', type=float, default=0.5,
+ help='Probability of switching to cutmix when both mixup and cutmix enabled')
+ parser.add_argument('--mixup_mode', type=str, default='batch',
+ help='How to apply mixup/cutmix params. Per "batch", "pair", or "elem"')
+
+ # augmentation parameters for imagenet finetuning
+ parser.add_argument('--color_jitter', type=float, default=0.4, metavar='PCT',
+ help='Color jitter factor (default: 0.4)')
+ parser.add_argument('--aa', type=str, default='rand-m9-mstd0.5-inc1', metavar='NAME',
+ help='Use AutoAugment policy. "v0" or "original". " + "(default: rand-m9-mstd0.5-inc1)')
+ parser.add_argument('--smoothing', type=float, default=0.1,
+ help='Label smoothing (default: 0.1)')
+
+ # evaluation parameters for imagenet
+ parser.add_argument('--crop_pct', type=float, default=None)
+
+ # random Erase params for imagenet finetuning
+ parser.add_argument('--reprob', type=float, default=0.25, metavar='PCT',
+ help='Random erase prob (default: 0.25)')
+ parser.add_argument('--remode', type=str, default='pixel',
+ help='Random erase mode (default: "pixel")')
+ parser.add_argument('--recount', type=int, default=1,
+ help='Random erase count (default: 1)')
+ parser.add_argument('--resplit', action='store_true', default=False,
+ help='Do not random erase first (clean) augmentation split')
+
+ # parameter for captioning finetuning
+ parser.add_argument('--captioning_mask_prob', type=float, default=0.6)
+ parser.add_argument('--drop_worst_ratio', type=float, default=0.2)
+ parser.add_argument('--drop_worst_after', type=int, default=12000)
+ parser.add_argument('--num_beams', type=int, default=3)
+ parser.add_argument('--length_penalty', type=float, default=0.6)
+
+ # label smoothing for imagenet and captioning
+ parser.add_argument('--label_smoothing', type=float, default=0.1)
+
+ # deepspeed parameters
+ parser.add_argument('--enable_deepspeed', action='store_true', default=False)
+ parser.add_argument('--initial_scale_power', type=int, default=16)
+ parser.add_argument('--zero_stage', default=0, type=int,
+ help='ZeRO optimizer stage (default: 0)')
+
+ known_args, _ = parser.parse_known_args()
+
+ if known_args.enable_deepspeed:
+ try:
+ import deepspeed
+ from deepspeed import DeepSpeedConfig
+ parser = deepspeed.add_config_arguments(parser)
+ ds_init = deepspeed.initialize
+ except:
+ print("Please 'pip install deepspeed==0.4.0'")
+ exit(0)
+ else:
+ ds_init = None
+
+ return parser.parse_args(), ds_init
+
+
+def main(args, ds_init):
+ utils.init_distributed_mode(args)
+
+ if ds_init is not None:
+ utils.create_ds_config(args)
+
+ if args.task_cache_path is None:
+ args.task_cache_path = args.output_dir
+
+ print(args)
+
+ device = torch.device(args.device)
+
+ # fix the seed for reproducibility
+ seed = args.seed + utils.get_rank()
+ torch.manual_seed(seed)
+ np.random.seed(seed)
+ # random.seed(seed)
+
+ cudnn.benchmark = True
+
+ if utils.get_rank() == 0 and args.log_dir is not None:
+ os.makedirs(args.log_dir, exist_ok=True)
+ log_writer = utils.TensorboardLogger(log_dir=args.log_dir)
+ else:
+ log_writer = None
+
+ data_loader_train, data_loader_val = create_downstream_dataset(args)
+
+ if not args.model.endswith(args.task):
+ if args.task in ("flickr30k", "coco_retrieval"):
+ model_config = "%s_retrieval" % args.model
+ elif args.task in ("coco_captioning", "nocaps"):
+ model_config = "%s_captioning" % args.model
+ elif args.task in ("imagenet"):
+ model_config = "%s_imageclassification" % args.model
+ else:
+ model_config = "%s_%s" % (args.model, args.task)
+ else:
+ model_config = args.model
+ print("model_config = %s" % model_config)
+ model = create_model(
+ model_config,
+ pretrained=False,
+ drop_path_rate=args.drop_path,
+ vocab_size=args.vocab_size,
+ checkpoint_activations=args.checkpoint_activations,
+ )
+
+ if args.finetune:
+ utils.load_model_and_may_interpolate(args.finetune, model, args.model_key, args.model_prefix)
+
+ model.to(device)
+
+ model_ema = None
+ if args.model_ema:
+ # Important to create EMA model after cuda(), DP wrapper, and AMP but before SyncBN and DDP wrapper
+ model_ema = ModelEma(
+ model,
+ decay=args.model_ema_decay,
+ device='cpu' if args.model_ema_force_cpu else '',
+ resume='')
+ print("Using EMA with decay = %.8f" % args.model_ema_decay)
+
+ model_without_ddp = model
+ n_parameters = sum(p.numel() for p in model.parameters() if p.requires_grad)
+
+ print("Model = %s" % str(model_without_ddp))
+ print('number of params:', n_parameters)
+
+ total_batch_size = args.batch_size * args.update_freq * utils.get_world_size()
+ num_training_steps_per_epoch = len(data_loader_train.dataset) // total_batch_size
+ print("LR = %.8f" % args.lr)
+ print("Batch size = %d" % total_batch_size)
+ print("Update frequent = %d" % args.update_freq)
+ print("Number of training examples = %d" % len(data_loader_train.dataset))
+ print("Number of training training per epoch = %d" % num_training_steps_per_epoch)
+
+ num_layers = model_without_ddp.get_num_layers()
+ if args.layer_decay < 1.0:
+ lrs = list(args.layer_decay ** (num_layers + 1 - i) for i in range(num_layers + 2))
+ assigner = LayerDecayValueAssigner(lrs)
+ elif args.task_head_lr_weight > 1:
+ assigner = LayerDecayValueAssigner([1.0, args.task_head_lr_weight], scale_handler=get_is_head_flag_for_vit)
+ else:
+ assigner = None
+
+ if assigner is not None:
+ print("Assigned values = %s" % str(assigner.values))
+
+ skip_weight_decay_list = model.no_weight_decay()
+
+ if args.distributed:
+ torch.distributed.barrier()
+ if args.enable_deepspeed:
+ loss_scaler = None
+ optimizer_params = get_parameter_groups(
+ model, args.weight_decay, skip_weight_decay_list,
+ assigner.get_layer_id if assigner is not None else None,
+ assigner.get_scale if assigner is not None else None)
+ model, optimizer, _, _ = ds_init(
+ args=args, model=model, model_parameters=optimizer_params,
+ dist_init_required=not args.distributed,
+ )
+
+ print("model.gradient_accumulation_steps() = %d" % model.gradient_accumulation_steps())
+ assert model.gradient_accumulation_steps() == args.update_freq
+ else:
+ if args.distributed:
+ model = torch.nn.parallel.DistributedDataParallel(model, device_ids=[args.gpu], find_unused_parameters=True)
+ model_without_ddp = model.module
+
+ optimizer = create_optimizer(
+ args, model_without_ddp, skip_list=skip_weight_decay_list,
+ get_num_layer=assigner.get_layer_id if assigner is not None else None,
+ get_layer_scale=assigner.get_scale if assigner is not None else None)
+ loss_scaler = NativeScaler()
+
+ lr_schedule_values = utils.cosine_scheduler(
+ args.lr, args.min_lr, args.epochs, num_training_steps_per_epoch,
+ warmup_epochs=args.warmup_epochs, warmup_steps=args.warmup_steps,
+ )
+
+ utils.auto_load_model(
+ args=args, model=model, model_without_ddp=model_without_ddp,
+ optimizer=optimizer, loss_scaler=loss_scaler, model_ema=model_ema)
+
+ task_handler = get_handler(args)
+
+ # mixup for imagenet
+ mixup_fn = None
+ if args.task in ["imagenet", "in1k"]:
+ mixup_active = args.mixup > 0 or args.cutmix > 0. or args.cutmix_minmax is not None
+ if mixup_active:
+ print("Mixup is activated!")
+ mixup_fn = Mixup(
+ mixup_alpha=args.mixup, cutmix_alpha=args.cutmix, cutmix_minmax=args.cutmix_minmax,
+ prob=args.mixup_prob, switch_prob=args.mixup_switch_prob, mode=args.mixup_mode,
+ label_smoothing=args.label_smoothing, num_classes=args.nb_classes)
+
+ if args.eval:
+ data_loader_test = create_downstream_dataset(args, is_eval=True)
+ if args.task in ["nlvr2", "flickr30k", "coco_retrieval", "imagenet"]:
+ ext_test_stats, task_key = evaluate(data_loader_test, model, device, task_handler)
+ print(f"Accuracy of the network on the {len(data_loader_test.dataset)} test images: {ext_test_stats[task_key]:.3f}%")
+ exit(0)
+ elif args.task == "vqav2":
+ result, _ = evaluate(data_loader_test, model, device, task_handler)
+ utils.dump_predictions(args, result, "vqav2_test")
+ exit(0)
+ elif args.task in ["coco_captioning", "nocaps"]:
+ predictions, _ = evaluate(data_loader_test, model, device, task_handler)
+ prediction_file = utils.dump_predictions(args, predictions, "{}_test".format(args.task))
+ if utils.is_main_process() and args.task == "coco_captioning":
+ captioning_result = utils.coco_caption_eval(args.output_dir, prediction_file, "{}_test".format(args.task))
+ result_file = os.path.join(args.output_dir, f"{args.task}_result.json")
+ print(json.dumps(captioning_result))
+ utils.write_result_to_jsonl(captioning_result, result_file)
+ exit(0)
+
+ print(f"Start training for {args.epochs} epochs")
+ start_time = time.time()
+
+ max_accuracy = 0.0
+ for epoch in range(args.start_epoch, args.epochs):
+ if args.distributed:
+ data_loader_train.sampler.set_epoch(epoch)
+ if log_writer is not None:
+ log_writer.set_step(epoch * num_training_steps_per_epoch * args.update_freq)
+ train_stats = train_one_epoch(
+ model, data_loader_train, optimizer, device, task_handler, epoch,
+ epoch * num_training_steps_per_epoch, lr_schedule_values, loss_scaler,
+ args.clip_grad, args.update_freq, model_ema, log_writer, args.task, mixup_fn,
+ )
+ if args.output_dir and args.save_ckpt:
+ if (epoch + 1) % args.save_ckpt_freq == 0 or epoch + 1 == args.epochs:
+ utils.save_model(
+ args=args, model=model, model_without_ddp=model_without_ddp, optimizer=optimizer,
+ loss_scaler=loss_scaler, epoch=epoch, model_ema=model_ema)
+ if data_loader_val is not None:
+ if args.task not in ["coco_captioning", "nocaps"]:
+ test_stats, task_key = evaluate(data_loader_val, model, device, task_handler)
+ else:
+ predictions, _ = evaluate(data_loader_val, model, device, task_handler)
+ prediction_file = utils.dump_predictions(args, predictions, f"{args.task}_val_e{epoch}")
+ result_file = os.path.join(args.output_dir, f"{args.task}_result_val_e{epoch}.json")
+ task_key = "CIDEr"
+ if utils.is_main_process():
+ test_stats = utils.coco_caption_eval(args.output_dir, prediction_file, "{}_val".format(args.task))
+ utils.write_result_to_jsonl(test_stats, result_file)
+ torch.distributed.barrier()
+ if not utils.is_main_process():
+ test_stats = utils.read_result_from_jsonl(result_file)
+
+ print(f"Performance of the network on the {len(data_loader_val.dataset)} val images: {test_stats[task_key]:.1f}%")
+ if max_accuracy < test_stats[task_key]:
+ max_accuracy = test_stats[task_key]
+ if args.output_dir and args.save_ckpt:
+ utils.save_model(
+ args=args, model=model, model_without_ddp=model_without_ddp, optimizer=optimizer,
+ loss_scaler=loss_scaler, epoch="best", model_ema=model_ema)
+
+ print(f'Max performance: {max_accuracy:.2f}%')
+ if log_writer is not None:
+ log_writer.update(acc=test_stats[task_key], head="perf", step=epoch)
+
+ log_stats = {**{f'train_{k}': v for k, v in train_stats.items()},
+ **{f'val_{k}': v for k, v in test_stats.items()},
+ 'epoch': epoch,
+ 'n_parameters': n_parameters}
+ else:
+ log_stats = {**{f'train_{k}': v for k, v in train_stats.items()},
+ # **{f'test_{k}': v for k, v in test_stats.items()},
+ 'epoch': epoch,
+ 'n_parameters': n_parameters}
+
+ if args.output_dir and utils.is_main_process():
+ if log_writer is not None:
+ log_writer.flush()
+ with open(os.path.join(args.output_dir, "log.txt"), mode="a", encoding="utf-8") as f:
+ f.write(json.dumps(log_stats) + "\n")
+
+ total_time = time.time() - start_time
+ total_time_str = str(datetime.timedelta(seconds=int(total_time)))
+ print('Training time {}'.format(total_time_str))
+
+
+if __name__ == '__main__':
+ opts, ds_init = get_args()
+ if opts.output_dir:
+ Path(opts.output_dir).mkdir(parents=True, exist_ok=True)
+ main(opts, ds_init)
diff --git a/py/evf_sam/model/unilm/beit3/utils.py b/py/evf_sam/model/unilm/beit3/utils.py
new file mode 100644
index 0000000..ca052f0
--- /dev/null
+++ b/py/evf_sam/model/unilm/beit3/utils.py
@@ -0,0 +1,913 @@
+# --------------------------------------------------------
+# Image as a Foreign Language: BEiT Pretraining for Vision and Vision-Language Tasks (https://arxiv.org/abs/2208.10442)
+# Github source: https://github.com/microsoft/unilm/tree/master/beit3
+# Copyright (c) 2023 Microsoft
+# Licensed under The MIT License [see LICENSE for details]
+# --------------------------------------------------------'
+
+import datetime
+import io
+import os
+import math
+import time
+import json
+import argparse
+import numpy as np
+from pathlib import Path
+from collections import defaultdict, deque
+from timm.utils import get_state_dict
+
+import torch
+import torch.distributed as dist
+import torch.nn as nn
+import torch.nn.functional as F
+from torch._six import inf
+from torchmetrics import Metric
+from tensorboardX import SummaryWriter
+
+
+def bool_flag(s):
+ """
+ Parse boolean arguments from the command line.
+ """
+ FALSY_STRINGS = {"off", "false", "0"}
+ TRUTHY_STRINGS = {"on", "true", "1"}
+ if s.lower() in FALSY_STRINGS:
+ return False
+ elif s.lower() in TRUTHY_STRINGS:
+ return True
+ else:
+ raise argparse.ArgumentTypeError("invalid value for a boolean flag")
+
+
+class SmoothedValue(object):
+ """Track a series of values and provide access to smoothed values over a
+ window or the global series average.
+ """
+
+ def __init__(self, window_size=20, fmt=None):
+ if fmt is None:
+ fmt = "{median:.4f} ({global_avg:.4f})"
+ self.deque = deque(maxlen=window_size)
+ self.total = 0.0
+ self.count = 0
+ self.fmt = fmt
+
+ def update(self, value, n=1):
+ self.deque.append(value)
+ self.count += n
+ self.total += value * n
+
+ def synchronize_between_processes(self):
+ """
+ Warning: does not synchronize the deque!
+ """
+ if not is_dist_avail_and_initialized():
+ return
+ t = torch.tensor([self.count, self.total], dtype=torch.float64, device='cuda')
+ dist.barrier()
+ dist.all_reduce(t)
+ t = t.tolist()
+ self.count = int(t[0])
+ self.total = t[1]
+
+ @property
+ def median(self):
+ d = torch.tensor(list(self.deque))
+ return d.median().item()
+
+ @property
+ def avg(self):
+ d = torch.tensor(list(self.deque), dtype=torch.float32)
+ return d.mean().item()
+
+ @property
+ def global_avg(self):
+ return self.total / self.count
+
+ @property
+ def max(self):
+ return max(self.deque)
+
+ @property
+ def value(self):
+ return self.deque[-1]
+
+ def __str__(self):
+ return self.fmt.format(
+ median=self.median,
+ avg=self.avg,
+ global_avg=self.global_avg,
+ max=self.max,
+ value=self.value)
+
+
+class MetricLogger(object):
+ def __init__(self, delimiter="\t"):
+ self.meters = defaultdict(SmoothedValue)
+ self.delimiter = delimiter
+
+ def update(self, **kwargs):
+ for k, v in kwargs.items():
+ if v is None:
+ continue
+ if isinstance(v, torch.Tensor):
+ v = v.item()
+ assert isinstance(v, (float, int))
+ self.meters[k].update(v)
+
+ def __getattr__(self, attr):
+ if attr in self.meters:
+ return self.meters[attr]
+ if attr in self.__dict__:
+ return self.__dict__[attr]
+ raise AttributeError("'{}' object has no attribute '{}'".format(
+ type(self).__name__, attr))
+
+ def __str__(self):
+ loss_str = []
+ for name, meter in self.meters.items():
+ loss_str.append(
+ "{}: {}".format(name, str(meter))
+ )
+ return self.delimiter.join(loss_str)
+
+ def synchronize_between_processes(self):
+ for meter in self.meters.values():
+ meter.synchronize_between_processes()
+
+ def add_meter(self, name, meter):
+ self.meters[name] = meter
+
+ def log_every(self, iterable, print_freq, header=None):
+ i = 0
+ if not header:
+ header = ''
+ start_time = time.time()
+ end = time.time()
+ iter_time = SmoothedValue(fmt='{avg:.4f}')
+ data_time = SmoothedValue(fmt='{avg:.4f}')
+ space_fmt = ':' + str(len(str(len(iterable)))) + 'd'
+ log_msg = [
+ header,
+ '[{0' + space_fmt + '}/{1}]',
+ 'eta: {eta}',
+ '{meters}',
+ 'time: {time}',
+ 'data: {data}'
+ ]
+ if torch.cuda.is_available():
+ log_msg.append('max mem: {memory:.0f}')
+ log_msg = self.delimiter.join(log_msg)
+ MB = 1024.0 * 1024.0
+ for obj in iterable:
+ data_time.update(time.time() - end)
+ yield obj
+ iter_time.update(time.time() - end)
+ if i % print_freq == 0 or i == len(iterable) - 1:
+ eta_seconds = iter_time.global_avg * (len(iterable) - i)
+ eta_string = str(datetime.timedelta(seconds=int(eta_seconds)))
+ if torch.cuda.is_available():
+ print(log_msg.format(
+ i, len(iterable), eta=eta_string,
+ meters=str(self),
+ time=str(iter_time), data=str(data_time),
+ memory=torch.cuda.max_memory_allocated() / MB))
+ else:
+ print(log_msg.format(
+ i, len(iterable), eta=eta_string,
+ meters=str(self),
+ time=str(iter_time), data=str(data_time)))
+ i += 1
+ end = time.time()
+ total_time = time.time() - start_time
+ total_time_str = str(datetime.timedelta(seconds=int(total_time)))
+ print('{} Total time: {} ({:.4f} s / it)'.format(
+ header, total_time_str, total_time / len(iterable)))
+
+
+class TensorboardLogger(object):
+ def __init__(self, log_dir):
+ self.writer = SummaryWriter(logdir=log_dir)
+ self.step = 0
+
+ def set_step(self, step=None):
+ if step is not None:
+ self.step = step
+ else:
+ self.step += 1
+
+ def update(self, head='scalar', step=None, **kwargs):
+ for k, v in kwargs.items():
+ if v is None:
+ continue
+ if isinstance(v, torch.Tensor):
+ v = v.item()
+ assert isinstance(v, (float, int))
+ self.writer.add_scalar(head + "/" + k, v, self.step if step is None else step)
+
+ def flush(self):
+ self.writer.flush()
+
+
+def _load_checkpoint_for_ema(model_ema, checkpoint):
+ """
+ Workaround for ModelEma._load_checkpoint to accept an already-loaded object
+ """
+ mem_file = io.BytesIO()
+ torch.save(checkpoint, mem_file)
+ mem_file.seek(0)
+ model_ema._load_checkpoint(mem_file)
+
+
+def setup_for_distributed(is_master):
+ """
+ This function disables printing when not in master process
+ """
+ import builtins as __builtin__
+ builtin_print = __builtin__.print
+
+ def print(*args, **kwargs):
+ force = kwargs.pop('force', False)
+ if is_master or force:
+ builtin_print(*args, **kwargs)
+
+ __builtin__.print = print
+
+
+def is_dist_avail_and_initialized():
+ if not dist.is_available():
+ return False
+ if not dist.is_initialized():
+ return False
+ return True
+
+
+def get_world_size():
+ if not is_dist_avail_and_initialized():
+ return 1
+ return dist.get_world_size()
+
+
+def get_rank():
+ if not is_dist_avail_and_initialized():
+ return 0
+ return dist.get_rank()
+
+
+def is_main_process():
+ return get_rank() == 0
+
+
+def save_on_master(*args, **kwargs):
+ if is_main_process():
+ torch.save(*args, **kwargs)
+
+
+def _get_rank_env():
+ if "RANK" in os.environ:
+ return int(os.environ["RANK"])
+ else:
+ return int(os.environ['OMPI_COMM_WORLD_RANK'])
+
+
+def _get_local_rank_env():
+ if "LOCAL_RANK" in os.environ:
+ return int(os.environ["LOCAL_RANK"])
+ else:
+ return int(os.environ['OMPI_COMM_WORLD_LOCAL_RANK'])
+
+
+def _get_world_size_env():
+ if "WORLD_SIZE" in os.environ:
+ return int(os.environ["WORLD_SIZE"])
+ else:
+ return int(os.environ['OMPI_COMM_WORLD_SIZE'])
+
+
+# The implementation code is modified from DeiT (https://github.com/facebookresearch/deit.git)
+def init_distributed_mode(args):
+ if args.dist_on_itp:
+ args.rank = _get_rank_env()
+ args.world_size = _get_world_size_env() # int(os.environ['OMPI_COMM_WORLD_SIZE'])
+ args.gpu = _get_local_rank_env()
+ args.dist_url = "tcp://%s:%s" % (os.environ['MASTER_ADDR'], os.environ['MASTER_PORT'])
+ os.environ['LOCAL_RANK'] = str(args.gpu)
+ os.environ['RANK'] = str(args.rank)
+ os.environ['WORLD_SIZE'] = str(args.world_size)
+ # ["RANK", "WORLD_SIZE", "MASTER_ADDR", "MASTER_PORT", "LOCAL_RANK"]
+ elif 'RANK' in os.environ and 'WORLD_SIZE' in os.environ:
+ args.rank = int(os.environ["RANK"])
+ args.world_size = int(os.environ['WORLD_SIZE'])
+ args.gpu = int(os.environ['LOCAL_RANK'])
+ elif 'SLURM_PROCID' in os.environ:
+ args.rank = int(os.environ['SLURM_PROCID'])
+ args.gpu = args.rank % torch.cuda.device_count()
+ else:
+ print('Not using distributed mode')
+ args.distributed = False
+ return
+
+ args.distributed = True
+
+ torch.cuda.set_device(args.gpu)
+ args.dist_backend = 'nccl'
+ print('| distributed init (rank {}): {}, gpu {}'.format(
+ args.rank, args.dist_url, args.gpu), flush=True)
+ torch.distributed.init_process_group(
+ backend=args.dist_backend, init_method=args.dist_url,
+ world_size=args.world_size, rank=args.rank,
+ timeout=datetime.timedelta(0, 7200)
+ )
+ torch.distributed.barrier()
+ setup_for_distributed(args.rank == 0)
+
+
+def load_state_dict(model, state_dict, prefix='', ignore_missing="relative_position_index"):
+ missing_keys = []
+ unexpected_keys = []
+ error_msgs = []
+ # copy state_dict so _load_from_state_dict can modify it
+ metadata = getattr(state_dict, '_metadata', None)
+ state_dict = state_dict.copy()
+ if metadata is not None:
+ state_dict._metadata = metadata
+
+ def load(module, prefix=''):
+ local_metadata = {} if metadata is None else metadata.get(
+ prefix[:-1], {})
+ module._load_from_state_dict(
+ state_dict, prefix, local_metadata, True, missing_keys, unexpected_keys, error_msgs)
+ for name, child in module._modules.items():
+ if child is not None:
+ load(child, prefix + name + '.')
+
+ load(model, prefix=prefix)
+
+ warn_missing_keys = []
+ ignore_missing_keys = []
+ for key in missing_keys:
+ keep_flag = True
+ for ignore_key in ignore_missing.split('|'):
+ if ignore_key in key:
+ keep_flag = False
+ break
+ if keep_flag:
+ warn_missing_keys.append(key)
+ else:
+ ignore_missing_keys.append(key)
+
+ missing_keys = warn_missing_keys
+
+ if len(missing_keys) > 0:
+ print("Weights of {} not initialized from pretrained model: {}".format(
+ model.__class__.__name__, missing_keys))
+ if len(unexpected_keys) > 0:
+ print("Weights from pretrained model not used in {}: {}".format(
+ model.__class__.__name__, unexpected_keys))
+ if len(ignore_missing_keys) > 0:
+ print("Ignored weights of {} not initialized from pretrained model: {}".format(
+ model.__class__.__name__, ignore_missing_keys))
+ if len(error_msgs) > 0:
+ print('\n'.join(error_msgs))
+
+
+class NativeScalerWithGradNormCount:
+ state_dict_key = "amp_scaler"
+
+ def __init__(self):
+ self._scaler = torch.cuda.amp.GradScaler()
+
+ def __call__(self, loss, optimizer, clip_grad=None, parameters=None, create_graph=False, update_grad=True):
+ self._scaler.scale(loss).backward(create_graph=create_graph)
+ if update_grad:
+ if clip_grad is not None:
+ assert parameters is not None
+ self._scaler.unscale_(optimizer) # unscale the gradients of optimizer's assigned params in-place
+ norm = torch.nn.utils.clip_grad_norm_(parameters, clip_grad)
+ else:
+ self._scaler.unscale_(optimizer)
+ norm = get_grad_norm_(parameters)
+ self._scaler.step(optimizer)
+ self._scaler.update()
+ else:
+ norm = None
+ return norm
+
+ def state_dict(self):
+ return self._scaler.state_dict()
+
+ def load_state_dict(self, state_dict):
+ self._scaler.load_state_dict(state_dict)
+
+
+def get_grad_norm_(parameters, norm_type: float = 2.0) -> torch.Tensor:
+ if isinstance(parameters, torch.Tensor):
+ parameters = [parameters]
+ parameters = [p for p in parameters if p.grad is not None]
+ norm_type = float(norm_type)
+ if len(parameters) == 0:
+ return torch.tensor(0.)
+ device = parameters[0].grad.device
+ if norm_type == inf:
+ total_norm = max(p.grad.detach().abs().max().to(device) for p in parameters)
+ else:
+ total_norm = torch.norm(torch.stack([torch.norm(p.grad.detach(), norm_type).to(device) for p in parameters]), norm_type)
+ return total_norm
+
+
+def cosine_scheduler(base_value, final_value, epochs, niter_per_ep, warmup_epochs=0,
+ start_warmup_value=0, warmup_steps=-1, sched_type="cos"):
+ warmup_schedule = np.array([])
+ warmup_iters = warmup_epochs * niter_per_ep
+ if warmup_steps > 0:
+ warmup_iters = warmup_steps
+ print("Set warmup steps = %d" % warmup_iters)
+ if warmup_epochs > 0:
+ warmup_schedule = np.linspace(start_warmup_value, base_value, warmup_iters)
+
+ if sched_type == "cos":
+ iters = np.arange(epochs * niter_per_ep - warmup_iters)
+ schedule = np.array([
+ final_value + 0.5 * (base_value - final_value) * (1 + math.cos(math.pi * i / (len(iters)))) for i in iters])
+ elif sched_type == "linear":
+ schedule = np.linspace(base_value, final_value, epochs * niter_per_ep - warmup_iters)
+ else:
+ raise NotImplementedError()
+
+ schedule = np.concatenate((warmup_schedule, schedule))
+
+ assert len(schedule) == epochs * niter_per_ep
+ return schedule
+
+
+def save_model(args, epoch, model, model_without_ddp, optimizer, loss_scaler, model_ema=None):
+ output_dir = Path(args.output_dir)
+ if loss_scaler is not None:
+ checkpoint_paths = [output_dir / ('checkpoint-%s.pth' % epoch)]
+ for checkpoint_path in checkpoint_paths:
+ to_save = {
+ 'model': model_without_ddp.state_dict(),
+ 'optimizer': optimizer.state_dict(),
+ 'epoch': epoch,
+ 'scaler': loss_scaler.state_dict(),
+ 'args': args,
+ }
+
+ if model_ema is not None:
+ to_save['model_ema'] = get_state_dict(model_ema)
+
+ save_on_master(to_save, checkpoint_path)
+ else:
+ client_state = {'epoch': epoch, "args": args}
+ if model_ema is not None:
+ client_state['model_ema'] = get_state_dict(model_ema)
+ model.save_checkpoint(save_dir=args.output_dir, tag="checkpoint-%s" % epoch, client_state=client_state)
+
+
+def auto_load_model(args, model, model_without_ddp, optimizer, loss_scaler, model_ema=None):
+ output_dir = Path(args.output_dir)
+ if loss_scaler is not None:
+ # torch.amp
+ if args.auto_resume and len(args.resume) == 0:
+ import glob
+ all_checkpoints = glob.glob(os.path.join(output_dir, 'checkpoint-*.pth'))
+ latest_ckpt = -1
+ for ckpt in all_checkpoints:
+ t = ckpt.split('-')[-1].split('.')[0]
+ if t.isdigit():
+ latest_ckpt = max(int(t), latest_ckpt)
+ if latest_ckpt >= 0:
+ args.resume = os.path.join(output_dir, 'checkpoint-%d.pth' % latest_ckpt)
+ print("Auto resume checkpoint: %s" % args.resume)
+
+ if args.resume:
+ if args.resume.startswith('https'):
+ checkpoint = torch.hub.load_state_dict_from_url(
+ args.resume, map_location='cpu', check_hash=True)
+ else:
+ checkpoint = torch.load(args.resume, map_location='cpu')
+ model_without_ddp.load_state_dict(checkpoint['model'])
+ print("Resume checkpoint %s" % args.resume)
+ if 'optimizer' in checkpoint and 'epoch' in checkpoint:
+ optimizer.load_state_dict(checkpoint['optimizer'])
+ args.start_epoch = checkpoint['epoch'] + 1
+ if hasattr(args, 'model_ema') and args.model_ema:
+ _load_checkpoint_for_ema(model_ema, checkpoint['model_ema'])
+ if 'scaler' in checkpoint:
+ loss_scaler.load_state_dict(checkpoint['scaler'])
+ print("With optim & sched!")
+ else:
+ # deepspeed, only support '--auto_resume'.
+ if args.auto_resume:
+ import glob
+ all_checkpoints = glob.glob(os.path.join(output_dir, 'checkpoint-*'))
+ latest_ckpt = -1
+ for ckpt in all_checkpoints:
+ t = ckpt.split('-')[-1].split('.')[0]
+ if t.isdigit():
+ latest_ckpt = max(int(t), latest_ckpt)
+ if latest_ckpt >= 0:
+ args.resume = os.path.join(output_dir, 'checkpoint-%d' % latest_ckpt)
+ print("Auto resume checkpoint: %d" % latest_ckpt)
+ _, client_states = model.load_checkpoint(args.output_dir, tag='checkpoint-%d' % latest_ckpt)
+ args.start_epoch = client_states['epoch'] + 1
+ if model_ema is not None:
+ if args.model_ema:
+ _load_checkpoint_for_ema(model_ema, client_states['model_ema'])
+
+
+# The implementation code is modified from DeiT (https://github.com/facebookresearch/deit.git)
+def load_model_and_may_interpolate(ckpt_path, model, model_key, model_prefix):
+ if ckpt_path.startswith('https'):
+ checkpoint = torch.hub.load_state_dict_from_url(
+ ckpt_path, map_location='cpu', check_hash=True)
+ else:
+ checkpoint = torch.load(ckpt_path, map_location='cpu')
+
+ print("Load ckpt from %s" % ckpt_path)
+ checkpoint_model = None
+ for model_key in model_key.split('|'):
+ if model_key in checkpoint:
+ checkpoint_model = checkpoint[model_key]
+ print("Load state_dict by model_key = %s" % model_key)
+ break
+
+ if checkpoint_model is None:
+ checkpoint_model = checkpoint
+
+ state_dict = model.state_dict()
+ for k in ['head.weight', 'head.bias']:
+ if k in checkpoint_model and checkpoint_model[k].shape != state_dict[k].shape:
+ print(f"Removing key {k} from pretrained checkpoint")
+ del checkpoint_model[k]
+
+ # interpolate position embedding
+ for pos_embed_key in ("vision_pos_embed", "pos_embed", "beit3.encoder.embed_positions.A.weight"):
+ if pos_embed_key in checkpoint_model:
+ pos_embed_checkpoint = checkpoint_model[pos_embed_key]
+ embedding_size = pos_embed_checkpoint.shape[-1]
+ if pos_embed_key == "beit3.encoder.embed_positions.A.weight":
+ # being consistent with Fairseq, which starts from 2 for position embedding
+ torchscale_model = True
+ num_patches = model.beit3.vision_embed.num_patches
+ num_extra_tokens = model.beit3.vision_embed.num_position_embeddings() + 2 - num_patches
+ else:
+ torchscale_model = False
+ num_patches = model.patch_embed.num_patches
+ num_extra_tokens = getattr(model, pos_embed_key).shape[-2] - num_patches
+ # height (== width) for the checkpoint position embedding
+ orig_size = int((pos_embed_checkpoint.shape[-2] - num_extra_tokens) ** 0.5)
+ # height (== width) for the new position embedding
+ new_size = int(num_patches ** 0.5)
+ # class_token and dist_token are kept unchanged
+ if orig_size != new_size:
+ print("Position interpolate from %dx%d to %dx%d" % (orig_size, orig_size, new_size, new_size))
+ if torchscale_model:
+ extra_tokens = pos_embed_checkpoint[:num_extra_tokens].unsqueeze(0)
+ # only the position tokens are interpolated
+ pos_tokens = pos_embed_checkpoint[num_extra_tokens:]
+ else:
+ extra_tokens = pos_embed_checkpoint[:, :num_extra_tokens]
+ # only the position tokens are interpolated
+ pos_tokens = pos_embed_checkpoint[:, num_extra_tokens:]
+ pos_tokens = pos_tokens.reshape(-1, orig_size, orig_size, embedding_size).permute(0, 3, 1, 2)
+ pos_tokens = torch.nn.functional.interpolate(
+ pos_tokens, size=(new_size, new_size), mode='bicubic', align_corners=False)
+ pos_tokens = pos_tokens.permute(0, 2, 3, 1).flatten(1, 2)
+ new_pos_embed = torch.cat((extra_tokens, pos_tokens), dim=1)
+ if torchscale_model:
+ new_pos_embed = new_pos_embed.squeeze(0)
+ checkpoint_model[pos_embed_key] = new_pos_embed
+
+ load_state_dict(model, checkpoint_model, prefix=model_prefix)
+
+
+def create_ds_config(args):
+ args.deepspeed_config = os.path.join(args.output_dir, "deepspeed_config.json")
+ with open(args.deepspeed_config, mode="w") as writer:
+ ds_config = {
+ "train_batch_size": args.batch_size * args.update_freq * get_world_size(),
+ "train_micro_batch_size_per_gpu": args.batch_size,
+ "steps_per_print": 1000,
+ "optimizer": {
+ "type": "Adam",
+ "adam_w_mode": True,
+ "params": {
+ "lr": args.lr,
+ "weight_decay": args.weight_decay,
+ "bias_correction": True,
+ "betas": [
+ args.opt_betas[0],
+ args.opt_betas[1]
+ ],
+ "eps": args.opt_eps
+ }
+ },
+ "fp16": {
+ "enabled": True,
+ "loss_scale": 0,
+ "initial_scale_power": getattr(args, "initial_scale_power", 12),
+ "loss_scale_window": 1000,
+ "hysteresis": 2,
+ "min_loss_scale": 1
+ },
+ "amp": {
+ "enabled": False,
+ "opt_level": "O2"
+ }
+ }
+
+ if args.clip_grad is not None:
+ ds_config.update({'gradient_clipping': args.clip_grad})
+
+ if args.zero_stage == 1:
+ ds_config.update({"zero_optimization": {"stage": args.zero_stage, "reduce_bucket_size": 5e8}})
+ elif args.zero_stage > 1:
+ raise NotImplementedError()
+
+ writer.write(json.dumps(ds_config, indent=2))
+
+
+def merge_batch_tensors_by_dict_key(batch):
+ batch_tensors = {}
+ for tensor_key in batch[0]:
+ if isinstance(batch[0][tensor_key], torch.Tensor):
+ batch_tensors[tensor_key] = torch.stack([d[tensor_key] for d in batch])
+ else:
+ batch_tensors[tensor_key] = torch.tensor([d[tensor_key] for d in batch], dtype=torch.long)
+ return batch_tensors
+
+
+def get_loss_scale_for_deepspeed(model):
+ optimizer = model.optimizer
+ loss_scale = None
+ if hasattr(optimizer, 'loss_scale'):
+ loss_scale = optimizer.loss_scale
+ elif hasattr(optimizer, 'cur_scale'):
+ loss_scale = optimizer.cur_scale
+ return loss_scale
+
+
+class GatherLayer(torch.autograd.Function):
+ """
+ Gather tensors from all workers with support for backward propagation:
+ This implementation does not cut the gradients as torch.distributed.all_gather does.
+ """
+ @staticmethod
+ def forward(ctx, x):
+ output = [torch.zeros_like(x) for _ in range(dist.get_world_size())]
+ dist.all_gather(output, x)
+ return tuple(output)
+ @staticmethod
+ def backward(ctx, *grads):
+ all_gradients = torch.stack(grads)
+ dist.all_reduce(all_gradients)
+ return all_gradients[dist.get_rank()]
+
+
+def gather_features(
+ image_features,
+ text_features,
+):
+ gathered_image_features = GatherLayer.apply(image_features)
+ gathered_text_features = GatherLayer.apply(text_features)
+ all_image_features = torch.cat(gathered_image_features)
+ all_text_features = torch.cat(gathered_text_features)
+
+ return all_image_features, all_text_features
+
+
+# The implementation code is modified from open_clip (https://github.com/mlfoundations/open_clip.git)
+class ClipLoss(nn.Module):
+
+ def __init__(
+ self,
+ cache_labels=False,
+ rank=0,
+ world_size=1,
+ ):
+ super().__init__()
+ self.cache_labels = cache_labels
+ self.rank = rank
+ self.world_size = world_size
+
+ # cache state
+ self.prev_num_logits = 0
+ self.labels = {}
+
+ def forward(self, image_features, text_features, logit_scale):
+ device = image_features.device
+ if self.world_size > 1:
+ all_image_features, all_text_features = gather_features(
+ image_features, text_features
+ )
+
+ logits_per_image = logit_scale * image_features @ all_text_features.T
+ logits_per_text = logit_scale * text_features @ all_image_features.T
+ else:
+ logits_per_image = logit_scale * image_features @ text_features.T
+ logits_per_text = logit_scale * text_features @ image_features.T
+
+ # calculated ground-truth and cache if enabled
+ num_logits = logits_per_image.shape[0]
+ if self.prev_num_logits != num_logits or device not in self.labels:
+ labels = torch.arange(num_logits, device=device, dtype=torch.long)
+ if self.world_size > 1:
+ labels = labels + num_logits * self.rank
+ if self.cache_labels:
+ self.labels[device] = labels
+ self.prev_num_logits = num_logits
+ else:
+ labels = self.labels[device]
+
+ total_loss = (
+ F.cross_entropy(logits_per_image, labels) +
+ F.cross_entropy(logits_per_text, labels)
+ ) / 2
+ return total_loss, logits_per_image, logits_per_text
+
+
+def write_result_to_jsonl(test_stats, result_file):
+ with open(result_file, mode="w", encoding="utf-8") as writer:
+ writer.write(json.dumps(test_stats, indent=None))
+
+
+def read_result_from_jsonl(result_file):
+ with open(result_file, mode="r", encoding="utf-8") as reader:
+ return json.load(reader)
+
+
+# The implementation code is from ViLT (https://github.com/dandelin/ViLT.git)
+class VQAScore(Metric):
+ def __init__(self, dist_sync_on_step=False):
+ super().__init__(dist_sync_on_step=dist_sync_on_step)
+ self.add_state("score", default=torch.tensor(0.0), dist_reduce_fx="sum")
+ self.add_state("total", default=torch.tensor(0.0), dist_reduce_fx="sum")
+
+ def update(self, logits, target):
+ logits, target = (
+ logits.detach().float().to(self.score.device),
+ target.detach().float().to(self.score.device),
+ )
+ logits = torch.max(logits, 1)[1]
+ one_hots = torch.zeros(*target.size()).to(target)
+ one_hots.scatter_(1, logits.view(-1, 1), 1)
+ scores = one_hots * target
+
+ self.score += scores.sum()
+ self.total += len(logits)
+
+ def compute(self):
+ return self.score / self.total
+
+
+class BertCaptioningLoss(nn.Module):
+ def __init__(self, label_smoothing, drop_worst_ratio, drop_worst_after):
+ super().__init__()
+ self.label_smoothing = label_smoothing
+ self.drop_worst_ratio = drop_worst_ratio
+ self.drop_worst_after = drop_worst_after
+ self.log_soft = nn.LogSoftmax(dim=1)
+ self.kl = nn.KLDivLoss(reduction='none')
+ self.iter = 0
+
+ def forward(self, logits, target, iter):
+ eps = self.label_smoothing
+ n_class = logits.size(1)
+ one_hot = torch.zeros_like(logits).scatter(1, target.view(-1, 1), 1)
+ one_hot = one_hot * (1 - eps) + (1 - one_hot) * eps / (n_class - 1)
+ log_prb = self.log_soft(logits)
+ loss = self.kl(log_prb, one_hot).sum(1)
+
+ if self.drop_worst_ratio > 0 and iter > self.drop_worst_after:
+ loss, _ = torch.topk(loss,
+ k=int(loss.shape[0] * (1-self.drop_worst_ratio)),
+ largest=False)
+ loss = loss.mean()
+
+ return loss
+
+
+class BeamHypotheses(object):
+ def __init__(self, n_hyp, max_length, length_penalty, early_stopping):
+ """
+ Initialize n-best list of hypotheses.
+ """
+ self.max_length = max_length - 1 # ignoring bos_token
+ self.length_penalty = length_penalty
+ self.early_stopping = early_stopping
+ self.n_hyp = n_hyp
+ self.hyp = []
+ self.worst_score = 1e9
+
+ def __len__(self):
+ """
+ Number of hypotheses in the list.
+ """
+ return len(self.hyp)
+
+ def add(self, hyp, sum_logprobs):
+ """
+ Add a new hypothesis to the list.
+ """
+ score = sum_logprobs / len(hyp) ** self.length_penalty
+ if len(self) < self.n_hyp or score > self.worst_score:
+ self.hyp.append((score, hyp))
+ if len(self) > self.n_hyp:
+ sorted_scores = sorted([(s, idx) for idx, (s, _) in enumerate(self.hyp)])
+ del self.hyp[sorted_scores[0][1]]
+ self.worst_score = sorted_scores[1][0]
+ else:
+ self.worst_score = min(score, self.worst_score)
+
+ def is_done(self, best_sum_logprobs):
+ """
+ If there are enough hypotheses and that none of the hypotheses being generated
+ can become better than the worst one in the heap, then we are done with this sentence.
+ """
+ if len(self) < self.n_hyp:
+ return False
+ elif self.early_stopping:
+ return True
+ else:
+ return self.worst_score >= best_sum_logprobs / self.max_length ** self.length_penalty
+
+
+def dump_predictions(args, result, file_suffix):
+ global_rank = get_rank()
+ jsons = None
+ if global_rank >= 0:
+ output_file = os.path.join(args.task_cache_path, f"submit_{global_rank}_{file_suffix}.json")
+ with open(output_file, "w") as fp:
+ json.dump(result, fp, indent=2)
+ torch.distributed.barrier()
+
+ if global_rank == 0:
+ world_size = get_world_size()
+ jsons = []
+ for i in range(world_size):
+ each_file = os.path.join(args.task_cache_path, f"submit_{i}_{file_suffix}.json")
+ with open(each_file, "r") as fp:
+ jsons += json.load(fp)
+
+ new_jsons = []
+ res_dict = dict()
+ if args.task in ["coco_captioning", "nocaps"]:
+ qid_key = "image_id"
+ else:
+ # for VQAv2
+ qid_key = "question_id"
+ for item in jsons:
+ if item[qid_key] in res_dict:
+ continue
+ new_jsons.append(item)
+ res_dict[item[qid_key]] = item
+ jsons = new_jsons
+
+ torch.distributed.barrier()
+ os.remove(output_file)
+ else:
+ jsons = result
+
+ result_file = os.path.join(args.output_dir, f"submit_{file_suffix}.json")
+ if jsons is not None:
+ with open(result_file, "w") as fp:
+ json.dump(jsons, fp, indent=2)
+ print("Infer %d examples into %s" % (len(jsons), result_file))
+ return result_file
+
+
+# The evaluation code is from BLIP (https://github.com/salesforce/BLIP)
+# For nocaps, please submit the prediction file to the evaluate server (https://eval.ai/web/challenges/challenge-page/355/overview) to obtain the final results
+def coco_caption_eval(gt_dir, results_file, split):
+ from pycocotools.coco import COCO
+ from pycocoevalcap.eval import COCOEvalCap
+ from torchvision.datasets.utils import download_url
+
+ urls = {'coco_captioning_val': 'https://storage.googleapis.com/sfr-vision-language-research/datasets/coco_karpathy_val_gt.json',
+ 'coco_captioning_test': 'https://storage.googleapis.com/sfr-vision-language-research/datasets/coco_karpathy_test_gt.json',
+ 'nocaps_val': 'https://github.com/addf400/files/releases/download/beit3/nocaps_val_gt.json'}
+ filenames = {'coco_captioning_val':'coco_karpathy_val_gt.json',
+ 'coco_captioning_test':'coco_karpathy_test_gt.json',
+ 'nocaps_val':'nocaps_val_gt.json'}
+
+ download_url(urls[split], gt_dir)
+ annotation_file = os.path.join(gt_dir, filenames[split])
+
+ # create coco object and coco_result object
+ coco = COCO(annotation_file)
+ coco_result = coco.loadRes(results_file)
+
+ # create coco_eval object by taking coco and coco_result
+ coco_eval = COCOEvalCap(coco, coco_result)
+
+ # evaluate results
+ # SPICE will take a few minutes the first time, but speeds up due to caching
+ coco_eval.evaluate()
+
+ res_dict = dict()
+ for metric, score in coco_eval.eval.items():
+ res_dict[metric] = score
+
+ return res_dict
diff --git a/py/evf_sam/utils/ade20k_classes.json b/py/evf_sam/utils/ade20k_classes.json
new file mode 100644
index 0000000..1f96e61
--- /dev/null
+++ b/py/evf_sam/utils/ade20k_classes.json
@@ -0,0 +1,30 @@
+[
+ "wall", "building", "sky", "floor", "tree", "ceiling", "road",
+ "bed", "windowpane", "grass", "cabinet", "sidewalk",
+ "person", "earth", "door", "table", "mountain", "plant",
+ "curtain", "chair", "car", "water", "painting", "sofa",
+ "shelf", "house", "sea", "mirror", "rug", "field", "armchair",
+ "seat", "fence", "desk", "rock", "wardrobe", "lamp",
+ "bathtub", "railing", "cushion", "base", "box", "column",
+ "signboard", "chest of drawers", "counter", "sand", "sink",
+ "skyscraper", "fireplace", "refrigerator", "grandstand",
+ "path", "stairs", "runway", "case", "pool table", "pillow",
+ "screen door", "stairway", "river", "bridge", "bookcase",
+ "blind", "coffee table", "toilet", "flower", "book", "hill",
+ "bench", "countertop", "stove", "palm", "kitchen island",
+ "computer", "swivel chair", "boat", "bar", "arcade machine",
+ "hovel", "bus", "towel", "light", "truck", "tower",
+ "chandelier", "awning", "streetlight", "booth",
+ "television receiver", "airplane", "dirt track", "apparel",
+ "pole", "land", "bannister", "escalator", "ottoman", "bottle",
+ "buffet", "poster", "stage", "van", "ship", "fountain",
+ "conveyer belt", "canopy", "washer", "plaything",
+ "swimming pool", "stool", "barrel", "basket", "waterfall",
+ "tent", "bag", "minibike", "cradle", "oven", "ball", "food",
+ "step", "tank", "trade name", "microwave", "pot", "animal",
+ "bicycle", "lake", "dishwasher", "screen", "blanket",
+ "sculpture", "hood", "sconce", "vase", "traffic light",
+ "tray", "ashcan", "fan", "pier", "crt screen", "plate",
+ "monitor", "bulletin board", "shower", "radiator", "glass",
+ "clock", "flag"
+]
\ No newline at end of file
diff --git a/py/evf_sam/utils/aug.py b/py/evf_sam/utils/aug.py
new file mode 100644
index 0000000..a01a5eb
--- /dev/null
+++ b/py/evf_sam/utils/aug.py
@@ -0,0 +1,117 @@
+# Copyright (c) Meta Platforms, Inc. and affiliates.
+# All rights reserved.
+
+# This source code is licensed under the license found in the
+# LICENSE file in the root directory of this source tree.
+
+from copy import deepcopy
+from typing import Tuple
+
+import numpy as np
+import torch
+from torch.nn import functional as F
+from torchvision.transforms.functional import resize # type: ignore
+from torchvision.transforms.functional import to_pil_image
+import random
+
+
+class RandomScale:
+ """
+ Resizes images to the longest side 'target_length', as well as provides
+ methods for resizing coordinates and boxes. Provides methods for
+ transforming both numpy array and batched torch tensors.
+ """
+
+ def __init__(self, max_length: int, min_length: int) -> None:
+ self.max_length = max_length
+ self.min_length = min_length
+
+ def apply_image(self, image: np.ndarray) -> np.ndarray:
+ """
+ Expects a numpy array with shape HxWxC in uint8 format.
+ """
+ target_size = self.get_preprocess_shape(
+ image.shape[0], image.shape[1], self.max_length, self.min_length
+ )
+ return np.array(resize(to_pil_image(image), target_size))
+
+ def apply_coords(
+ self, coords: np.ndarray, original_size: Tuple[int, ...]
+ ) -> np.ndarray:
+ """
+ Expects a numpy array of length 2 in the final dimension. Requires the
+ original image size in (H, W) format.
+ """
+ old_h, old_w = original_size
+ new_h, new_w = self.get_preprocess_shape(
+ original_size[0], original_size[1], self.max_length, self.min_length
+ )
+ coords = deepcopy(coords).astype(float)
+ coords[..., 0] = coords[..., 0] * (new_w / old_w)
+ coords[..., 1] = coords[..., 1] * (new_h / old_h)
+ return coords
+
+ def apply_boxes(
+ self, boxes: np.ndarray, original_size: Tuple[int, ...]
+ ) -> np.ndarray:
+ """
+ Expects a numpy array shape Bx4. Requires the original image size
+ in (H, W) format.
+ """
+ boxes = self.apply_coords(boxes.reshape(-1, 2, 2), original_size)
+ return boxes.reshape(-1, 4)
+
+ def apply_image_torch(self, image: torch.Tensor) -> torch.Tensor:
+ """
+ Expects batched images with shape BxCxHxW and float format. This
+ transformation may not exactly match apply_image. apply_image is
+ the transformation expected by the model.
+ """
+ # Expects an image in BCHW format. May not exactly match apply_image.
+ target_size = self.get_preprocess_shape(
+ image.shape[0], image.shape[1], self.max_length, self.min_length
+ )
+ return F.interpolate(
+ image, target_size, mode="bilinear", align_corners=False, antialias=True
+ )
+
+ def apply_coords_torch(
+ self, coords: torch.Tensor, original_size: Tuple[int, ...]
+ ) -> torch.Tensor:
+ """
+ Expects a torch tensor with length 2 in the last dimension. Requires the
+ original image size in (H, W) format.
+ """
+ old_h, old_w = original_size
+ new_h, new_w = self.get_preprocess_shape(
+ original_size[0], original_size[1], self.max_length, self.min_length
+ )
+ coords = deepcopy(coords).to(torch.float)
+ coords[..., 0] = coords[..., 0] * (new_w / old_w)
+ coords[..., 1] = coords[..., 1] * (new_h / old_h)
+ return coords
+
+ def apply_boxes_torch(
+ self, boxes: torch.Tensor, original_size: Tuple[int, ...]
+ ) -> torch.Tensor:
+ """
+ Expects a torch tensor with shape Bx4. Requires the original image
+ size in (H, W) format.
+ """
+ boxes = self.apply_coords_torch(boxes.reshape(-1, 2, 2), original_size)
+ return boxes.reshape(-1, 4)
+
+ @staticmethod
+ def get_preprocess_shape(
+ oldh: int, oldw: int, max_length: int, min_length: int
+ ) -> Tuple[int, int]:
+ """
+ Compute the output size given input size and target long side length.
+ """
+ max_scale = max_length * 1.0 / max(oldh, oldw)
+ min_scale = min_length * 1.0 / max(oldh, oldw)
+ scale = min_scale + random.random() * (max_scale-min_scale)
+ newh, neww = oldh * scale, oldw * scale
+ neww = int(neww + 0.5)
+ newh = int(newh + 0.5)
+ return (newh, neww)
diff --git a/py/evf_sam/utils/data_processing.py b/py/evf_sam/utils/data_processing.py
new file mode 100644
index 0000000..d47a80f
--- /dev/null
+++ b/py/evf_sam/utils/data_processing.py
@@ -0,0 +1,90 @@
+import glob
+import json
+import os
+
+import cv2
+import numpy as np
+
+
+def get_mask_from_json(json_path, img):
+ try:
+ with open(json_path, "r") as r:
+ anno = json.loads(r.read())
+ except:
+ with open(json_path, "r", encoding="cp1252") as r:
+ anno = json.loads(r.read())
+
+ inform = anno["shapes"]
+ comments = anno["text"]
+ is_sentence = anno["is_sentence"]
+
+ height, width = img.shape[:2]
+
+ ### sort polies by area
+ area_list = []
+ valid_poly_list = []
+ for i in inform:
+ label_id = i["label"]
+ points = i["points"]
+ if "flag" == label_id.lower(): ## meaningless deprecated annotations
+ continue
+
+ tmp_mask = np.zeros((height, width), dtype=np.uint8)
+ cv2.polylines(tmp_mask, np.array([points], dtype=np.int32), True, 1, 1)
+ cv2.fillPoly(tmp_mask, np.array([points], dtype=np.int32), 1)
+ tmp_area = tmp_mask.sum()
+
+ area_list.append(tmp_area)
+ valid_poly_list.append(i)
+
+ ### ground-truth mask
+ sort_index = np.argsort(area_list)[::-1].astype(np.int32)
+ sort_index = list(sort_index)
+ sort_inform = []
+ for s_idx in sort_index:
+ sort_inform.append(valid_poly_list[s_idx])
+
+ mask = np.zeros((height, width), dtype=np.uint8)
+ for i in sort_inform:
+ label_id = i["label"]
+ points = i["points"]
+
+ if "ignore" in label_id.lower():
+ label_value = 255 # ignored during evaluation
+ else:
+ label_value = 1 # target
+
+ cv2.polylines(mask, np.array([points], dtype=np.int32), True, label_value, 1)
+ cv2.fillPoly(mask, np.array([points], dtype=np.int32), label_value)
+
+ return mask, comments, is_sentence
+
+
+if __name__ == "__main__":
+ data_dir = "./train"
+ vis_dir = "./vis"
+
+ if not os.path.exists(vis_dir):
+ os.makedirs(vis_dir)
+
+ json_path_list = sorted(glob.glob(data_dir + "/*.json"))
+ for json_path in json_path_list:
+ img_path = json_path.replace(".json", ".jpg")
+ img = cv2.imread(img_path)[:, :, ::-1]
+
+ # In generated mask, value 1 denotes valid target region, and value 255 stands for region ignored during evaluaiton.
+ mask, comments, is_sentence = get_mask_from_json(json_path, img)
+
+ ## visualization. Green for target, and red for ignore.
+ valid_mask = (mask == 1).astype(np.float32)[:, :, None]
+ ignore_mask = (mask == 255).astype(np.float32)[:, :, None]
+ vis_img = img * (1 - valid_mask) * (1 - ignore_mask) + (
+ (np.array([0, 255, 0]) * 0.6 + img * 0.4) * valid_mask
+ + (np.array([255, 0, 0]) * 0.6 + img * 0.4) * ignore_mask
+ )
+ vis_img = np.concatenate([img, vis_img], 1)
+ vis_path = os.path.join(
+ vis_dir, json_path.split("/")[-1].replace(".json", ".jpg")
+ )
+ cv2.imwrite(vis_path, vis_img[:, :, ::-1])
+ print("Visualization has been saved to: ", vis_path)
diff --git a/py/evf_sam/utils/dataset.py b/py/evf_sam/utils/dataset.py
new file mode 100644
index 0000000..9969078
--- /dev/null
+++ b/py/evf_sam/utils/dataset.py
@@ -0,0 +1,386 @@
+import glob
+import os
+import random
+
+import cv2
+import numpy as np
+import torch
+import torch.nn.functional as F
+from pycocotools import mask
+
+from model.segment_anything.utils.transforms import ResizeLongestSide
+
+from .data_processing import get_mask_from_json
+from .refer import REFER
+from .refer_seg_dataset import ReferSegDataset
+from .sem_seg_dataset import SemSegDataset
+from torchvision import transforms
+import json
+from PIL import Image
+
+def collate_fn(
+ batch, tokenizer=None, local_rank=-1
+):
+ image_path_list = []
+ images_list = []
+ images_evf_list = []
+ masks_list = []
+ label_list = []
+ resize_list = []
+ sampled_classes_list = []
+ offset_list = [0]
+ cnt = 0
+ inferences = []
+ for (
+ image_path,
+ images,
+ images_evf,
+ masks,
+ label,
+ resize,
+ sampled_classes,
+ inference,
+ ) in batch:
+ image_path_list.append(image_path)
+ images_list.append(images)
+ images_evf_list.append(images_evf)
+ label_list.append(label)
+ masks_list.append(masks.float())
+ resize_list.append(resize)
+ sampled_classes_list.extend(sampled_classes)
+ cnt += len(sampled_classes)
+ offset_list.append(cnt)
+ inferences.append(inference)
+
+ input_ids = [
+ tokenizer(prompt, return_tensors="pt").input_ids[0]
+ for prompt in sampled_classes_list
+ ]
+
+ input_ids = torch.nn.utils.rnn.pad_sequence(
+ input_ids, batch_first=True, padding_value=tokenizer.pad_token_id
+ )
+ attention_masks = input_ids.ne(tokenizer.pad_token_id)
+
+ if inferences[0] == False:
+ truncate_len = tokenizer.model_max_length
+
+ if input_ids.shape[1] > truncate_len:
+ input_ids = input_ids[:, :truncate_len]
+ targets = targets[:, :truncate_len]
+ attention_masks = attention_masks[:, :truncate_len]
+
+ return {
+ "image_paths": image_path_list,
+ "images": torch.stack(images_list, dim=0),
+ "images_evf": torch.stack(images_evf_list, dim=0),
+ "input_ids": input_ids,
+ "attention_masks": attention_masks,
+ "masks_list": masks_list,
+ "label_list": label_list,
+ "resize_list": resize_list,
+ "offset": torch.LongTensor(offset_list),
+ "sampled_classes_list": sampled_classes_list,
+ "inference": inferences[0],
+ }
+
+
+class HybridDataset(torch.utils.data.Dataset):
+ pixel_mean = torch.Tensor([123.675, 116.28, 103.53]).view(-1, 1, 1)
+ pixel_std = torch.Tensor([58.395, 57.12, 57.375]).view(-1, 1, 1)
+ img_size = 1024
+ ignore_label = 255
+
+ def __init__(
+ self,
+ base_image_dir,
+ tokenizer,
+ samples_per_epoch=500 * 8 * 2 * 10,
+ precision: str = "fp32",
+ image_size: int = 224,
+ num_classes_per_sample: int = 3,
+ exclude_val=False,
+ dataset="sem_seg||refer_seg",
+ sample_rate=[9, 3, 3, 1],
+ sem_seg_data="ade20k||cocostuff||pascal_part||mapillary",
+ refer_seg_data="refclef||refcoco||refcoco+||refcocog",
+ explanatory=-1,
+ model_type="ori",
+ transform=ResizeLongestSide(1024),
+ ):
+ self.transform=transform
+ self.model_type = model_type
+ self.exclude_val = exclude_val
+ self.dataset = dataset
+ self.samples_per_epoch = samples_per_epoch
+ self.explanatory = explanatory
+ self.num_classes_per_sample = num_classes_per_sample
+ sample_rate = np.array(sample_rate)
+ self.sample_rate = sample_rate / sample_rate.sum()
+
+ self.base_image_dir = base_image_dir
+ self.image_size = image_size
+ self.tokenizer = tokenizer
+ self.precision = precision
+
+ self.datasets = dataset.split("||")
+
+ self.all_datasets = []
+ for dataset in self.datasets:
+ if dataset == "sem_seg":
+ self.all_datasets.append(
+ SemSegDataset(
+ base_image_dir,
+ tokenizer,
+ samples_per_epoch,
+ precision,
+ image_size,
+ num_classes_per_sample,
+ exclude_val,
+ sem_seg_data,
+ self.model_type,
+ self.transform
+ )
+ )
+ elif dataset == "refer_seg":
+ self.all_datasets.append(
+ ReferSegDataset(
+ base_image_dir,
+ tokenizer,
+ samples_per_epoch,
+ precision,
+ image_size,
+ num_classes_per_sample,
+ exclude_val,
+ refer_seg_data,
+ self.model_type,
+ self.transform
+ )
+ )
+
+ def __len__(self):
+ return self.samples_per_epoch
+
+ def __getitem__(self, idx):
+ ind = np.random.choice(list(range(len(self.datasets))), p=self.sample_rate)
+ data = self.all_datasets[ind]
+ inference = False
+ return *data[0], inference
+
+
+def init_ade20k(base_image_dir):
+ with open("utils/ade20k_classes.json", "r") as f:
+ ade20k_classes = json.load(f)
+ ade20k_classes = np.array(ade20k_classes)
+ image_ids = sorted(
+ os.listdir(os.path.join(base_image_dir, "ade20k/images", "validation"))
+ )
+ ade20k_image_ids = []
+ for x in image_ids:
+ if x.endswith(".jpg"):
+ ade20k_image_ids.append(x[:-4])
+ ade20k_images = []
+ for image_id in ade20k_image_ids: # self.descriptions:
+ ade20k_images.append(
+ os.path.join(
+ base_image_dir,
+ "ade20k",
+ "images",
+ "validation",
+ "{}.jpg".format(image_id),
+ )
+ )
+ ade20k_labels = [
+ x.replace(".jpg", ".png").replace("images", "annotations")
+ for x in ade20k_images
+ ]
+ print("ade20k: ", len(ade20k_images))
+ return ade20k_classes, ade20k_images, ade20k_labels
+
+
+class ValDataset(torch.utils.data.Dataset):
+ pixel_mean = torch.Tensor([123.675, 116.28, 103.53]).view(-1, 1, 1)
+ pixel_std = torch.Tensor([58.395, 57.12, 57.375]).view(-1, 1, 1)
+ img_size = 1024
+ ignore_label = 255
+
+ def __init__(
+ self,
+ base_image_dir,
+ tokenizer,
+ val_dataset,
+ image_size=224,
+ model_type="ori"
+ ):
+ self.model_type = model_type
+ self.base_image_dir = base_image_dir
+ splits = val_dataset.split("|")
+ if len(splits) == 3:
+ ds, splitBy, split = splits
+ base_image_dir = os.path.join(base_image_dir, "refer_seg")
+ refer_api = REFER(base_image_dir, ds, splitBy)
+ ref_ids_val = refer_api.getRefIds(split=split)
+ images_ids_val = refer_api.getImgIds(ref_ids=ref_ids_val)
+ refs_val = refer_api.loadRefs(ref_ids=ref_ids_val)
+ refer_seg_ds = {}
+ refer_seg_ds["images"] = []
+ loaded_images = refer_api.loadImgs(image_ids=images_ids_val)
+ for item in loaded_images:
+ item = item.copy()
+ if ds == "refclef":
+ item["file_name"] = os.path.join(
+ base_image_dir, "images/saiapr_tc-12", item["file_name"]
+ )
+ elif ds in ["refcoco", "refcoco+", "refcocog", "grefcoco"]:
+ item["file_name"] = os.path.join(
+ base_image_dir,
+ "images/mscoco/images/train2014",
+ item["file_name"],
+ )
+ refer_seg_ds["images"].append(item)
+ refer_seg_ds["annotations"] = refer_api.Anns # anns_val
+
+ img2refs = {}
+ for ref in refs_val:
+ image_id = ref["image_id"]
+ img2refs[image_id] = img2refs.get(image_id, []) + [
+ ref,
+ ]
+ refer_seg_ds["img2refs"] = img2refs
+ self.refer_seg_ds = refer_seg_ds
+ self.data_type = "refer_seg"
+ elif val_dataset=="ade":
+ ds = "ade"
+ self.classes, self.images, self.labels = init_ade20k(base_image_dir)
+ self.data_type = "sem_seg"
+
+
+ self.ds = ds
+ self.tokenizer = tokenizer
+ self.transform = ResizeLongestSide(1024)
+ self.image_preprocessor = transforms.Compose([
+ transforms.ToTensor(),
+ transforms.Resize((image_size, image_size), interpolation=3),
+ transforms.Normalize(mean=(0.5, 0.5, 0.5), std=(0.5, 0.5, 0.5))
+ ])
+ def __len__(self):
+ if self.data_type == "refer_seg":
+ return len(self.refer_seg_ds["images"])
+ else:
+ return len(self.images)
+
+ def preprocess(self, x: torch.Tensor) -> torch.Tensor:
+ """Normalize pixel values and pad to a square input."""
+ # Normalize colors
+ x = (x - self.pixel_mean) / self.pixel_std
+
+ if self.model_type=="effi" or self.model_type=="sam2":
+ x = F.interpolate(x.unsqueeze(0), (self.img_size, self.img_size), mode="bilinear").squeeze(0)
+ else:
+ # Pad
+ h, w = x.shape[-2:]
+ padh = self.img_size - h
+ padw = self.img_size - w
+ x = F.pad(x, (0, padw, 0, padh))
+ return x
+
+ def __getitem__(self, idx):
+ if self.data_type == "refer_seg":
+ refer_seg_ds = self.refer_seg_ds
+ images = refer_seg_ds["images"]
+ annotations = refer_seg_ds["annotations"]
+ img2refs = refer_seg_ds["img2refs"]
+
+ image_info = images[idx]
+ image_path = image_info["file_name"]
+ image_id = image_info["id"]
+
+ refs = img2refs[image_id]
+ if len(refs) == 0:
+ raise ValueError("image {} has no refs".format(image_id))
+
+ sents = []
+ ann_ids = []
+ for ref in refs:
+ for sent in ref["sentences"]:
+ sents.append(sent["sent"].strip().lower())
+ ann_ids.append(ref["ann_id"])
+
+ sampled_sents = sents
+ sampled_ann_ids = ann_ids
+ image = cv2.imread(image_path)
+ image = cv2.cvtColor(image, cv2.COLOR_BGR2RGB)
+ is_sentence = False
+
+ elif self.data_type == "sem_seg":
+ image_path = self.images[idx]
+ label_path = self.labels[idx]
+ label = Image.open(label_path)
+ label = np.array(label)
+ label[label == 0] = 255
+ label -= 1
+ label[label == 254] = 255
+ unique_label = np.unique(label).tolist()
+ if 255 in unique_label:
+ unique_label.remove(255)
+
+ sampled_sents = [self.classes[class_id] for class_id in unique_label]
+
+ img = cv2.imread(image_path)
+ image = cv2.cvtColor(img, cv2.COLOR_BGR2RGB)
+ class_ids = unique_label
+ label = torch.from_numpy(label).long()
+ masks = []
+ for class_id in class_ids:
+ masks.append(label == class_id)
+ masks = torch.stack(masks, dim=0)
+
+ # preprocess image for evf
+ image_evf = self.image_preprocessor(image)
+
+ # preprocess image for sam
+ image = self.transform.apply_image(image)
+ resize = image.shape[:2]
+ image = self.preprocess(torch.from_numpy(image).permute(2, 0, 1).contiguous())
+
+ if self.data_type == "refer_seg":
+ masks = []
+ for i, ann_id in enumerate(sampled_ann_ids):
+ ann = annotations[ann_id]
+ if len(ann["segmentation"]) == 0 and sampled_sents[i] != "":
+ m = np.zeros((image_info["height"], image_info["width"], 1))
+ else:
+ if type(ann["segmentation"][0]) == list: # polygon
+ rle = mask.frPyObjects(
+ ann["segmentation"],
+ image_info["height"],
+ image_info["width"],
+ )
+ else:
+ rle = ann["segmentation"]
+ for i in range(len(rle)):
+ if not isinstance(rle[i]["counts"], bytes):
+ rle[i]["counts"] = rle[i]["counts"].encode()
+ m = mask.decode(rle)
+ m = np.sum(
+ m, axis=2
+ ) # sometimes there are multiple binary map (corresponding to multiple segs)
+ m = m.astype(np.uint8) # convert to np.uint8
+ masks.append(m)
+
+ if not isinstance(masks, torch.Tensor):
+ masks = np.stack(masks, axis=0)
+ masks = torch.from_numpy(masks)
+ labels = torch.ones(masks.shape[1], masks.shape[2]) * self.ignore_label
+ inference = True
+
+ return (
+ image_path,
+ image,
+ image_evf,
+ masks,
+ labels,
+ resize,
+ sampled_sents,
+ inference,
+ )
diff --git a/py/evf_sam/utils/grefcoco.py b/py/evf_sam/utils/grefcoco.py
new file mode 100644
index 0000000..7e1a49f
--- /dev/null
+++ b/py/evf_sam/utils/grefcoco.py
@@ -0,0 +1,193 @@
+import contextlib
+import copy
+import io
+import logging
+import os
+import random
+
+import numpy as np
+import pycocotools.mask as mask_util
+from detectron2.structures import Boxes, BoxMode, PolygonMasks, RotatedBoxes
+from detectron2.utils.file_io import PathManager
+from fvcore.common.timer import Timer
+from PIL import Image
+
+"""
+This file contains functions to parse RefCOCO-format annotations into dicts in "Detectron2 format".
+"""
+
+
+logger = logging.getLogger(__name__)
+
+__all__ = ["load_refcoco_json"]
+
+
+def load_grefcoco_json(
+ refer_root,
+ dataset_name,
+ splitby,
+ split,
+ image_root,
+ extra_annotation_keys=None,
+ extra_refer_keys=None,
+):
+ if dataset_name == "refcocop":
+ dataset_name = "refcoco+"
+ if dataset_name == "refcoco" or dataset_name == "refcoco+":
+ splitby == "unc"
+ if dataset_name == "refcocog":
+ assert splitby == "umd" or splitby == "google"
+
+ dataset_id = "_".join([dataset_name, splitby, split])
+
+ from .grefer import G_REFER
+
+ logger.info("Loading dataset {} ({}-{}) ...".format(dataset_name, splitby, split))
+ logger.info("Refcoco root: {}".format(refer_root))
+ timer = Timer()
+ refer_root = PathManager.get_local_path(refer_root)
+ with contextlib.redirect_stdout(io.StringIO()):
+ refer_api = G_REFER(data_root=refer_root, dataset=dataset_name, splitBy=splitby)
+ if timer.seconds() > 1:
+ logger.info(
+ "Loading {} takes {:.2f} seconds.".format(dataset_id, timer.seconds())
+ )
+
+ ref_ids = refer_api.getRefIds(split=split)
+ img_ids = refer_api.getImgIds(ref_ids)
+ refs = refer_api.loadRefs(ref_ids)
+ imgs = [refer_api.loadImgs(ref["image_id"])[0] for ref in refs]
+ anns = [refer_api.loadAnns(ref["ann_id"]) for ref in refs]
+ imgs_refs_anns = list(zip(imgs, refs, anns))
+
+ logger.info(
+ "Loaded {} images, {} referring object sets in G_RefCOCO format from {}".format(
+ len(img_ids), len(ref_ids), dataset_id
+ )
+ )
+
+ dataset_dicts = []
+
+ ann_keys = ["iscrowd", "bbox", "category_id"] + (extra_annotation_keys or [])
+ ref_keys = ["raw", "sent_id"] + (extra_refer_keys or [])
+
+ ann_lib = {}
+
+ NT_count = 0
+ MT_count = 0
+
+ for img_dict, ref_dict, anno_dicts in imgs_refs_anns:
+ record = {}
+ record["source"] = "grefcoco"
+ record["file_name"] = os.path.join(image_root, img_dict["file_name"])
+ record["height"] = img_dict["height"]
+ record["width"] = img_dict["width"]
+ image_id = record["image_id"] = img_dict["id"]
+
+ # Check that information of image, ann and ref match each other
+ # This fails only when the data parsing logic or the annotation file is buggy.
+ assert ref_dict["image_id"] == image_id
+ assert ref_dict["split"] == split
+ if not isinstance(ref_dict["ann_id"], list):
+ ref_dict["ann_id"] = [ref_dict["ann_id"]]
+
+ # No target samples
+ if None in anno_dicts:
+ assert anno_dicts == [None]
+ assert ref_dict["ann_id"] == [-1]
+ record["empty"] = True
+ obj = {key: None for key in ann_keys if key in ann_keys}
+ obj["bbox_mode"] = BoxMode.XYWH_ABS
+ obj["empty"] = True
+ obj = [obj]
+
+ # Multi target samples
+ else:
+ record["empty"] = False
+ obj = []
+ for anno_dict in anno_dicts:
+ ann_id = anno_dict["id"]
+ if anno_dict["iscrowd"]:
+ continue
+ assert anno_dict["image_id"] == image_id
+ assert ann_id in ref_dict["ann_id"]
+
+ if ann_id in ann_lib:
+ ann = ann_lib[ann_id]
+ else:
+ ann = {key: anno_dict[key] for key in ann_keys if key in anno_dict}
+ ann["bbox_mode"] = BoxMode.XYWH_ABS
+ ann["empty"] = False
+
+ segm = anno_dict.get("segmentation", None)
+ assert segm # either list[list[float]] or dict(RLE)
+ if isinstance(segm, dict):
+ if isinstance(segm["counts"], list):
+ # convert to compressed RLE
+ segm = mask_util.frPyObjects(segm, *segm["size"])
+ else:
+ # filter out invalid polygons (< 3 points)
+ segm = [
+ poly
+ for poly in segm
+ if len(poly) % 2 == 0 and len(poly) >= 6
+ ]
+ if len(segm) == 0:
+ num_instances_without_valid_segmentation += 1
+ continue # ignore this instance
+ ann["segmentation"] = segm
+ ann_lib[ann_id] = ann
+
+ obj.append(ann)
+
+ record["annotations"] = obj
+
+ # Process referring expressions
+ sents = ref_dict["sentences"]
+ for sent in sents:
+ ref_record = record.copy()
+ ref = {key: sent[key] for key in ref_keys if key in sent}
+ ref["ref_id"] = ref_dict["ref_id"]
+ ref_record["sentence"] = ref
+ dataset_dicts.append(ref_record)
+ # if ref_record['empty']:
+ # NT_count += 1
+ # else:
+ # MT_count += 1
+
+ # logger.info("NT samples: %d, MT samples: %d", NT_count, MT_count)
+
+ # Debug mode
+ # return dataset_dicts[:100]
+
+ return dataset_dicts
+
+
+if __name__ == "__main__":
+ """
+ Test the COCO json dataset loader.
+
+ Usage:
+ python -m detectron2.data.datasets.coco \
+ path/to/json path/to/image_root dataset_name
+
+ "dataset_name" can be "coco_2014_minival_100", or other
+ pre-registered ones
+ """
+ import sys
+
+ REFCOCO_PATH = "/mnt/lustre/hhding/code/ReLA/datasets"
+ COCO_TRAIN_2014_IMAGE_ROOT = "/mnt/lustre/hhding/code/ReLA/datasets/images"
+ REFCOCO_DATASET = "grefcoco"
+ REFCOCO_SPLITBY = "unc"
+ REFCOCO_SPLIT = "train"
+
+
+ dicts = load_grefcoco_json(
+ REFCOCO_PATH,
+ REFCOCO_DATASET,
+ REFCOCO_SPLITBY,
+ REFCOCO_SPLIT,
+ COCO_TRAIN_2014_IMAGE_ROOT,
+ )
+ print(1)
diff --git a/py/evf_sam/utils/grefer.py b/py/evf_sam/utils/grefer.py
new file mode 100644
index 0000000..3c881c5
--- /dev/null
+++ b/py/evf_sam/utils/grefer.py
@@ -0,0 +1,352 @@
+"""
+grefer v0.1
+This interface provides access to gRefCOCO.
+
+The following API functions are defined:
+G_REFER - REFER api class
+getRefIds - get ref ids that satisfy given filter conditions.
+getAnnIds - get ann ids that satisfy given filter conditions.
+getImgIds - get image ids that satisfy given filter conditions.
+getCatIds - get category ids that satisfy given filter conditions.
+loadRefs - load refs with the specified ref ids.
+loadAnns - load anns with the specified ann ids.
+loadImgs - load images with the specified image ids.
+loadCats - load category names with the specified category ids.
+getRefBox - get ref's bounding box [x, y, w, h] given the ref_id
+showRef - show image, segmentation or box of the referred object with the ref
+getMaskByRef - get mask and area of the referred object given ref or ref ids
+getMask - get mask and area of the referred object given ref
+showMask - show mask of the referred object given ref
+"""
+
+import itertools
+import json
+import os.path as osp
+import pickle
+import time
+
+import matplotlib.pyplot as plt
+import numpy as np
+import skimage.io as io
+from matplotlib.collections import PatchCollection
+from matplotlib.patches import Polygon, Rectangle
+from pycocotools import mask
+
+
+class G_REFER:
+ def __init__(self, data_root, dataset="grefcoco", splitBy="unc"):
+ # provide data_root folder which contains grefcoco
+ print("loading dataset %s into memory..." % dataset)
+ self.ROOT_DIR = osp.abspath(osp.dirname(__file__))
+ self.DATA_DIR = osp.join(data_root, dataset)
+ if dataset in ["grefcoco"]:
+ self.IMAGE_DIR = osp.join(data_root, "images/train2014")
+ else:
+ raise KeyError("No refer dataset is called [%s]" % dataset)
+
+ tic = time.time()
+
+ # load refs from data/dataset/refs(dataset).json
+ self.data = {}
+ self.data["dataset"] = dataset
+
+ ref_file = osp.join(self.DATA_DIR, f"grefs({splitBy}).p")
+ if osp.exists(ref_file):
+ self.data["refs"] = pickle.load(open(ref_file, "rb"), fix_imports=True)
+ else:
+ ref_file = osp.join(self.DATA_DIR, f"grefs({splitBy}).json")
+ if osp.exists(ref_file):
+ self.data["refs"] = json.load(open(ref_file, "rb"))
+ else:
+ raise FileNotFoundError("JSON file not found")
+
+ # load annotations from data/dataset/instances.json
+ instances_file = osp.join(self.DATA_DIR, "instances.json")
+ instances = json.load(open(instances_file, "r"))
+ self.data["images"] = instances["images"]
+ self.data["annotations"] = instances["annotations"]
+ self.data["categories"] = instances["categories"]
+
+ # create index
+ self.createIndex()
+ print("DONE (t=%.2fs)" % (time.time() - tic))
+
+ @staticmethod
+ def _toList(x):
+ return x if isinstance(x, list) else [x]
+
+ @staticmethod
+ def match_any(a, b):
+ a = a if isinstance(a, list) else [a]
+ b = b if isinstance(b, list) else [b]
+ return set(a) & set(b)
+
+ def createIndex(self):
+ # create sets of mapping
+ # 1) Refs: {ref_id: ref}
+ # 2) Anns: {ann_id: ann}
+ # 3) Imgs: {image_id: image}
+ # 4) Cats: {category_id: category_name}
+ # 5) Sents: {sent_id: sent}
+ # 6) imgToRefs: {image_id: refs}
+ # 7) imgToAnns: {image_id: anns}
+ # 8) refToAnn: {ref_id: ann}
+ # 9) annToRef: {ann_id: ref}
+ # 10) catToRefs: {category_id: refs}
+ # 11) sentToRef: {sent_id: ref}
+ # 12) sentToTokens: {sent_id: tokens}
+ print("creating index...")
+ # fetch info from instances
+ Anns, Imgs, Cats, imgToAnns = {}, {}, {}, {}
+ Anns[-1] = None
+ for ann in self.data["annotations"]:
+ Anns[ann["id"]] = ann
+ imgToAnns[ann["image_id"]] = imgToAnns.get(ann["image_id"], []) + [ann]
+ for img in self.data["images"]:
+ Imgs[img["id"]] = img
+ for cat in self.data["categories"]:
+ Cats[cat["id"]] = cat["name"]
+
+ # fetch info from refs
+ Refs, imgToRefs, refToAnn, annToRef, catToRefs = {}, {}, {}, {}, {}
+ Sents, sentToRef, sentToTokens = {}, {}, {}
+ availableSplits = []
+ for ref in self.data["refs"]:
+ # ids
+ ref_id = ref["ref_id"]
+ ann_id = ref["ann_id"]
+ category_id = ref["category_id"]
+ image_id = ref["image_id"]
+
+ if ref["split"] not in availableSplits:
+ availableSplits.append(ref["split"])
+
+ # add mapping related to ref
+ if ref_id in Refs:
+ print("Duplicate ref id")
+ Refs[ref_id] = ref
+ imgToRefs[image_id] = imgToRefs.get(image_id, []) + [ref]
+
+ category_id = self._toList(category_id)
+ added_cats = []
+ for cat in category_id:
+ if cat not in added_cats:
+ added_cats.append(cat)
+ catToRefs[cat] = catToRefs.get(cat, []) + [ref]
+
+ ann_id = self._toList(ann_id)
+ refToAnn[ref_id] = [Anns[ann] for ann in ann_id]
+ for ann_id_n in ann_id:
+ annToRef[ann_id_n] = annToRef.get(ann_id_n, []) + [ref]
+
+ # add mapping of sent
+ for sent in ref["sentences"]:
+ Sents[sent["sent_id"]] = sent
+ sentToRef[sent["sent_id"]] = ref
+ sentToTokens[sent["sent_id"]] = sent["tokens"]
+
+ # create class members
+ self.Refs = Refs
+ self.Anns = Anns
+ self.Imgs = Imgs
+ self.Cats = Cats
+ self.Sents = Sents
+ self.imgToRefs = imgToRefs
+ self.imgToAnns = imgToAnns
+ self.refToAnn = refToAnn
+ self.annToRef = annToRef
+ self.catToRefs = catToRefs
+ self.sentToRef = sentToRef
+ self.sentToTokens = sentToTokens
+ self.availableSplits = availableSplits
+ print("index created.")
+
+ def getRefIds(self, image_ids=[], cat_ids=[], split=[]):
+ image_ids = self._toList(image_ids)
+ cat_ids = self._toList(cat_ids)
+ split = self._toList(split)
+
+ for s in split:
+ if s not in self.availableSplits:
+ raise ValueError(f"Invalid split name: {s}")
+
+ refs = self.data["refs"]
+
+ if len(image_ids) > 0:
+ lists = [self.imgToRefs[image_id] for image_id in image_ids]
+ refs = list(itertools.chain.from_iterable(lists))
+ if len(cat_ids) > 0:
+ refs = [ref for ref in refs if self.match_any(ref["category_id"], cat_ids)]
+ if len(split) > 0:
+ refs = [ref for ref in refs if ref["split"] in split]
+
+ ref_ids = [ref["ref_id"] for ref in refs]
+ return ref_ids
+
+ def getAnnIds(self, image_ids=[], ref_ids=[]):
+ image_ids = self._toList(image_ids)
+ ref_ids = self._toList(ref_ids)
+
+ if any([len(image_ids), len(ref_ids)]):
+ if len(image_ids) > 0:
+ lists = [
+ self.imgToAnns[image_id]
+ for image_id in image_ids
+ if image_id in self.imgToAnns
+ ]
+ anns = list(itertools.chain.from_iterable(lists))
+ else:
+ anns = self.data["annotations"]
+ ann_ids = [ann["id"] for ann in anns]
+ if len(ref_ids) > 0:
+ lists = [self.Refs[ref_id]["ann_id"] for ref_id in ref_ids]
+ anns_by_ref_id = list(itertools.chain.from_iterable(lists))
+ ann_ids = list(set(ann_ids).intersection(set(anns_by_ref_id)))
+ else:
+ ann_ids = [ann["id"] for ann in self.data["annotations"]]
+
+ return ann_ids
+
+ def getImgIds(self, ref_ids=[]):
+ ref_ids = self._toList(ref_ids)
+
+ if len(ref_ids) > 0:
+ image_ids = list(set([self.Refs[ref_id]["image_id"] for ref_id in ref_ids]))
+ else:
+ image_ids = self.Imgs.keys()
+ return image_ids
+
+ def getCatIds(self):
+ return self.Cats.keys()
+
+ def loadRefs(self, ref_ids=[]):
+ return [self.Refs[ref_id] for ref_id in self._toList(ref_ids)]
+
+ def loadAnns(self, ann_ids=[]):
+ if isinstance(ann_ids, str):
+ ann_ids = int(ann_ids)
+ return [self.Anns[ann_id] for ann_id in self._toList(ann_ids)]
+
+ def loadImgs(self, image_ids=[]):
+ return [self.Imgs[image_id] for image_id in self._toList(image_ids)]
+
+ def loadCats(self, cat_ids=[]):
+ return [self.Cats[cat_id] for cat_id in self._toList(cat_ids)]
+
+ def getRefBox(self, ref_id):
+ anns = self.refToAnn[ref_id]
+ return [ann["bbox"] for ann in anns] # [x, y, w, h]
+
+ def showRef(self, ref, seg_box="seg"):
+ ax = plt.gca()
+ # show image
+ image = self.Imgs[ref["image_id"]]
+ I = io.imread(osp.join(self.IMAGE_DIR, image["file_name"]))
+ ax.imshow(I)
+ # show refer expression
+ for sid, sent in enumerate(ref["sentences"]):
+ print("%s. %s" % (sid + 1, sent["sent"]))
+ # show segmentations
+ if seg_box == "seg":
+ ann_id = ref["ann_id"]
+ ann = self.Anns[ann_id]
+ polygons = []
+ color = []
+ c = "none"
+ if type(ann["segmentation"][0]) == list:
+ # polygon used for refcoco*
+ for seg in ann["segmentation"]:
+ poly = np.array(seg).reshape((len(seg) / 2, 2))
+ polygons.append(Polygon(poly, True, alpha=0.4))
+ color.append(c)
+ p = PatchCollection(
+ polygons,
+ facecolors=color,
+ edgecolors=(1, 1, 0, 0),
+ linewidths=3,
+ alpha=1,
+ )
+ ax.add_collection(p) # thick yellow polygon
+ p = PatchCollection(
+ polygons,
+ facecolors=color,
+ edgecolors=(1, 0, 0, 0),
+ linewidths=1,
+ alpha=1,
+ )
+ ax.add_collection(p) # thin red polygon
+ else:
+ # mask used for refclef
+ rle = ann["segmentation"]
+ m = mask.decode(rle)
+ img = np.ones((m.shape[0], m.shape[1], 3))
+ color_mask = np.array([2.0, 166.0, 101.0]) / 255
+ for i in range(3):
+ img[:, :, i] = color_mask[i]
+ ax.imshow(np.dstack((img, m * 0.5)))
+ # show bounding-box
+ elif seg_box == "box":
+ ann_id = ref["ann_id"]
+ ann = self.Anns[ann_id]
+ bbox = self.getRefBox(ref["ref_id"])
+ box_plot = Rectangle(
+ (bbox[0], bbox[1]),
+ bbox[2],
+ bbox[3],
+ fill=False,
+ edgecolor="green",
+ linewidth=3,
+ )
+ ax.add_patch(box_plot)
+
+ def getMask(self, ann):
+ if not ann:
+ return None
+ if ann["iscrowd"]:
+ raise ValueError("Crowd object")
+ image = self.Imgs[ann["image_id"]]
+ if type(ann["segmentation"][0]) == list: # polygon
+ rle = mask.frPyObjects(ann["segmentation"], image["height"], image["width"])
+ else:
+ rle = ann["segmentation"]
+
+ m = mask.decode(rle)
+ m = np.sum(
+ m, axis=2
+ ) # sometimes there are multiple binary map (corresponding to multiple segs)
+ m = m.astype(np.uint8) # convert to np.uint8
+ # compute area
+ area = sum(mask.area(rle)) # should be close to ann['area']
+ return {"mask": m, "area": area}
+
+ def getMaskByRef(self, ref=None, ref_id=None, merge=False):
+ if not ref and not ref_id:
+ raise ValueError
+ if ref:
+ ann_ids = ref["ann_id"]
+ ref_id = ref["ref_id"]
+ else:
+ ann_ids = self.getAnnIds(ref_ids=ref_id)
+
+ if ann_ids == [-1]:
+ img = self.Imgs[self.Refs[ref_id]["image_id"]]
+ return {
+ "mask": np.zeros([img["height"], img["width"]], dtype=np.uint8),
+ "empty": True,
+ }
+
+ anns = self.loadAnns(ann_ids)
+ mask_list = [self.getMask(ann) for ann in anns if not ann["iscrowd"]]
+
+ if merge:
+ merged_masks = sum([mask["mask"] for mask in mask_list])
+ merged_masks[np.where(merged_masks > 1)] = 1
+ return {"mask": merged_masks, "empty": False}
+ else:
+ return mask_list
+
+ def showMask(self, ref):
+ M = self.getMask(ref)
+ msk = M["mask"]
+ ax = plt.gca()
+ ax.imshow(msk)
diff --git a/py/evf_sam/utils/refer.py b/py/evf_sam/utils/refer.py
new file mode 100644
index 0000000..3b4cea7
--- /dev/null
+++ b/py/evf_sam/utils/refer.py
@@ -0,0 +1,391 @@
+__author__ = "licheng"
+
+"""
+This interface provides access to four datasets:
+1) refclef
+2) refcoco
+3) refcoco+
+4) refcocog
+split by unc and google
+
+The following API functions are defined:
+REFER - REFER api class
+getRefIds - get ref ids that satisfy given filter conditions.
+getAnnIds - get ann ids that satisfy given filter conditions.
+getImgIds - get image ids that satisfy given filter conditions.
+getCatIds - get category ids that satisfy given filter conditions.
+loadRefs - load refs with the specified ref ids.
+loadAnns - load anns with the specified ann ids.
+loadImgs - load images with the specified image ids.
+loadCats - load category names with the specified category ids.
+getRefBox - get ref's bounding box [x, y, w, h] given the ref_id
+showRef - show image, segmentation or box of the referred object with the ref
+getMask - get mask and area of the referred object given ref
+showMask - show mask of the referred object given ref
+"""
+
+import itertools
+import json
+import os.path as osp
+import pickle
+import sys
+import time
+from pprint import pprint
+
+import matplotlib.pyplot as plt
+import numpy as np
+import skimage.io as io
+from matplotlib.collections import PatchCollection
+from matplotlib.patches import Polygon, Rectangle
+from pycocotools import mask
+
+
+class REFER:
+ def __init__(self, data_root, dataset="refcoco", splitBy="unc"):
+ # provide data_root folder which contains refclef, refcoco, refcoco+ and refcocog
+ # also provide dataset name and splitBy information
+ # e.g., dataset = 'refcoco', splitBy = 'unc'
+ print("loading dataset %s into memory..." % dataset)
+ self.ROOT_DIR = osp.abspath(osp.dirname(__file__))
+ self.DATA_DIR = osp.join(data_root, dataset)
+ if dataset in ["refcoco", "refcoco+", "refcocog"]:
+ self.IMAGE_DIR = osp.join(data_root, "images/mscoco/images/train2014")
+ elif dataset == "refclef":
+ self.IMAGE_DIR = osp.join(data_root, "images/saiapr_tc-12")
+ else:
+ print("No refer dataset is called [%s]" % dataset)
+ sys.exit()
+
+ self.dataset = dataset
+
+ # load refs from data/dataset/refs(dataset).json
+ tic = time.time()
+
+ ref_file = osp.join(self.DATA_DIR, "refs(" + splitBy + ").p")
+ print("ref_file: ", ref_file)
+ self.data = {}
+ self.data["dataset"] = dataset
+ self.data["refs"] = pickle.load(open(ref_file, "rb"))
+
+ # load annotations from data/dataset/instances.json
+ instances_file = osp.join(self.DATA_DIR, "instances.json")
+ instances = json.load(open(instances_file, "rb"))
+ self.data["images"] = instances["images"]
+ self.data["annotations"] = instances["annotations"]
+ self.data["categories"] = instances["categories"]
+
+ # create index
+ self.createIndex()
+ print("DONE (t=%.2fs)" % (time.time() - tic))
+
+ def createIndex(self):
+ # create sets of mapping
+ # 1) Refs: {ref_id: ref}
+ # 2) Anns: {ann_id: ann}
+ # 3) Imgs: {image_id: image}
+ # 4) Cats: {category_id: category_name}
+ # 5) Sents: {sent_id: sent}
+ # 6) imgToRefs: {image_id: refs}
+ # 7) imgToAnns: {image_id: anns}
+ # 8) refToAnn: {ref_id: ann}
+ # 9) annToRef: {ann_id: ref}
+ # 10) catToRefs: {category_id: refs}
+ # 11) sentToRef: {sent_id: ref}
+ # 12) sentToTokens: {sent_id: tokens}
+ print("creating index...")
+ # fetch info from instances
+ Anns, Imgs, Cats, imgToAnns = {}, {}, {}, {}
+ for ann in self.data["annotations"]:
+ Anns[ann["id"]] = ann
+ imgToAnns[ann["image_id"]] = imgToAnns.get(ann["image_id"], []) + [ann]
+ for img in self.data["images"]:
+ Imgs[img["id"]] = img
+ for cat in self.data["categories"]:
+ Cats[cat["id"]] = cat["name"]
+
+ # fetch info from refs
+ Refs, imgToRefs, refToAnn, annToRef, catToRefs = {}, {}, {}, {}, {}
+ Sents, sentToRef, sentToTokens = {}, {}, {}
+ for ref in self.data["refs"]:
+ # ids
+ ref_id = ref["ref_id"]
+ ann_id = ref["ann_id"]
+ category_id = ref["category_id"]
+ image_id = ref["image_id"]
+
+ # add mapping related to ref
+ Refs[ref_id] = ref
+ imgToRefs[image_id] = imgToRefs.get(image_id, []) + [ref]
+ catToRefs[category_id] = catToRefs.get(category_id, []) + [ref]
+ refToAnn[ref_id] = Anns[ann_id]
+ annToRef[ann_id] = ref
+
+ # add mapping of sent
+ for sent in ref["sentences"]:
+ Sents[sent["sent_id"]] = sent
+ sentToRef[sent["sent_id"]] = ref
+ sentToTokens[sent["sent_id"]] = sent["tokens"]
+
+ # create class members
+ self.Refs = Refs
+ self.Anns = Anns
+ self.Imgs = Imgs
+ self.Cats = Cats
+ self.Sents = Sents
+ self.imgToRefs = imgToRefs
+ self.imgToAnns = imgToAnns
+ self.refToAnn = refToAnn
+ self.annToRef = annToRef
+ self.catToRefs = catToRefs
+ self.sentToRef = sentToRef
+ self.sentToTokens = sentToTokens
+ print("index created.")
+
+ def getRefIds(self, image_ids=[], cat_ids=[], ref_ids=[], split=""):
+ image_ids = image_ids if type(image_ids) == list else [image_ids]
+ cat_ids = cat_ids if type(cat_ids) == list else [cat_ids]
+ ref_ids = ref_ids if type(ref_ids) == list else [ref_ids]
+
+ if len(image_ids) == len(cat_ids) == len(ref_ids) == len(split) == 0:
+ refs = self.data["refs"]
+ else:
+ if not len(image_ids) == 0:
+ refs = [self.imgToRefs[image_id] for image_id in image_ids]
+ else:
+ refs = self.data["refs"]
+ if not len(cat_ids) == 0:
+ refs = [ref for ref in refs if ref["category_id"] in cat_ids]
+ if not len(ref_ids) == 0:
+ refs = [ref for ref in refs if ref["ref_id"] in ref_ids]
+ if not len(split) == 0:
+ if split in ["testA", "testB", "testC"]:
+ refs = [
+ ref for ref in refs if split[-1] in ref["split"]
+ ] # we also consider testAB, testBC, ...
+ elif split in ["testAB", "testBC", "testAC"]:
+ refs = [
+ ref for ref in refs if ref["split"] == split
+ ] # rarely used I guess...
+ elif split == "test":
+ refs = [ref for ref in refs if "test" in ref["split"]]
+ elif split == "train" or split == "val":
+ refs = [ref for ref in refs if ref["split"] == split]
+ else:
+ print("No such split [%s]" % split)
+ sys.exit()
+ ref_ids = [ref["ref_id"] for ref in refs]
+ return ref_ids
+
+ def getAnnIds(self, image_ids=[], cat_ids=[], ref_ids=[]):
+ image_ids = image_ids if type(image_ids) == list else [image_ids]
+ cat_ids = cat_ids if type(cat_ids) == list else [cat_ids]
+ ref_ids = ref_ids if type(ref_ids) == list else [ref_ids]
+
+ if len(image_ids) == len(cat_ids) == len(ref_ids) == 0:
+ ann_ids = [ann["id"] for ann in self.data["annotations"]]
+ else:
+ if not len(image_ids) == 0:
+ lists = [
+ self.imgToAnns[image_id]
+ for image_id in image_ids
+ if image_id in self.imgToAnns
+ ] # list of [anns]
+ anns = list(itertools.chain.from_iterable(lists))
+ else:
+ anns = self.data["annotations"]
+ if not len(cat_ids) == 0:
+ anns = [ann for ann in anns if ann["category_id"] in cat_ids]
+ ann_ids = [ann["id"] for ann in anns]
+ if not len(ref_ids) == 0:
+ ids = set(ann_ids).intersection(
+ set([self.Refs[ref_id]["ann_id"] for ref_id in ref_ids])
+ )
+ return ann_ids
+
+ def getImgIds(self, ref_ids=[]):
+ ref_ids = ref_ids if type(ref_ids) == list else [ref_ids]
+
+ if not len(ref_ids) == 0:
+ image_ids = list(set([self.Refs[ref_id]["image_id"] for ref_id in ref_ids]))
+ else:
+ image_ids = self.Imgs.keys()
+ return image_ids
+
+ def getCatIds(self):
+ return self.Cats.keys()
+
+ def loadRefs(self, ref_ids=[]):
+ if type(ref_ids) == list:
+ return [self.Refs[ref_id] for ref_id in ref_ids]
+ elif type(ref_ids) == int:
+ return [self.Refs[ref_ids]]
+
+ def loadAnns(self, ann_ids=[]):
+ if type(ann_ids) == list:
+ return [self.Anns[ann_id] for ann_id in ann_ids]
+ elif type(ann_ids) == int or type(ann_ids) == unicode:
+ return [self.Anns[ann_ids]]
+
+ def loadImgs(self, image_ids=[]):
+ if type(image_ids) == list:
+ return [self.Imgs[image_id] for image_id in image_ids]
+ elif type(image_ids) == int:
+ return [self.Imgs[image_ids]]
+
+ def loadCats(self, cat_ids=[]):
+ if type(cat_ids) == list:
+ return [self.Cats[cat_id] for cat_id in cat_ids]
+ elif type(cat_ids) == int:
+ return [self.Cats[cat_ids]]
+
+ def getRefBox(self, ref_id):
+ ref = self.Refs[ref_id]
+ ann = self.refToAnn[ref_id]
+ return ann["bbox"] # [x, y, w, h]
+
+ def showRef(self, ref, seg_box="seg"):
+ ax = plt.gca()
+ # show image
+ image = self.Imgs[ref["image_id"]]
+ I = io.imread(osp.join(self.IMAGE_DIR, image["file_name"]))
+ ax.imshow(I)
+ # show refer expression
+ for sid, sent in enumerate(ref["sentences"]):
+ print("%s. %s" % (sid + 1, sent["sent"]))
+ # show segmentations
+ if seg_box == "seg":
+ ann_id = ref["ann_id"]
+ ann = self.Anns[ann_id]
+ polygons = []
+ color = []
+ c = "none"
+ if type(ann["segmentation"][0]) == list:
+ # polygon used for refcoco*
+ for seg in ann["segmentation"]:
+ poly = np.array(seg).reshape((len(seg) / 2, 2))
+ polygons.append(Polygon(poly, True, alpha=0.4))
+ color.append(c)
+ p = PatchCollection(
+ polygons,
+ facecolors=color,
+ edgecolors=(1, 1, 0, 0),
+ linewidths=3,
+ alpha=1,
+ )
+ ax.add_collection(p) # thick yellow polygon
+ p = PatchCollection(
+ polygons,
+ facecolors=color,
+ edgecolors=(1, 0, 0, 0),
+ linewidths=1,
+ alpha=1,
+ )
+ ax.add_collection(p) # thin red polygon
+ else:
+ # mask used for refclef
+ rle = ann["segmentation"]
+ m = mask.decode(rle)
+ img = np.ones((m.shape[0], m.shape[1], 3))
+ color_mask = np.array([2.0, 166.0, 101.0]) / 255
+ for i in range(3):
+ img[:, :, i] = color_mask[i]
+ ax.imshow(np.dstack((img, m * 0.5)))
+ # show bounding-box
+ elif seg_box == "box":
+ ann_id = ref["ann_id"]
+ ann = self.Anns[ann_id]
+ bbox = self.getRefBox(ref["ref_id"])
+ box_plot = Rectangle(
+ (bbox[0], bbox[1]),
+ bbox[2],
+ bbox[3],
+ fill=False,
+ edgecolor="green",
+ linewidth=3,
+ )
+ ax.add_patch(box_plot)
+
+ def getMask(self, ref):
+ # return mask, area and mask-center
+ ann = self.refToAnn[ref["ref_id"]]
+ image = self.Imgs[ref["image_id"]]
+ if type(ann["segmentation"][0]) == list: # polygon
+ rle = mask.frPyObjects(ann["segmentation"], image["height"], image["width"])
+ else:
+ rle = ann["segmentation"]
+ m = mask.decode(rle)
+ m = np.sum(
+ m, axis=2
+ ) # sometimes there are multiple binary map (corresponding to multiple segs)
+ m = m.astype(np.uint8) # convert to np.uint8
+ # compute area
+ area = sum(mask.area(rle)) # should be close to ann['area']
+ return {"mask": m, "area": area}
+ # # position
+ # position_x = np.mean(np.where(m==1)[1]) # [1] means columns (matlab style) -> x (c style)
+ # position_y = np.mean(np.where(m==1)[0]) # [0] means rows (matlab style) -> y (c style)
+ # # mass position (if there were multiple regions, we use the largest one.)
+ # label_m = label(m, connectivity=m.ndim)
+ # regions = regionprops(label_m)
+ # if len(regions) > 0:
+ # largest_id = np.argmax(np.array([props.filled_area for props in regions]))
+ # largest_props = regions[largest_id]
+ # mass_y, mass_x = largest_props.centroid
+ # else:
+ # mass_x, mass_y = position_x, position_y
+ # # if centroid is not in mask, we find the closest point to it from mask
+ # if m[mass_y, mass_x] != 1:
+ # print('Finding closes mask point ...')
+ # kernel = np.ones((10, 10),np.uint8)
+ # me = cv2.erode(m, kernel, iterations = 1)
+ # points = zip(np.where(me == 1)[0].tolist(), np.where(me == 1)[1].tolist()) # row, col style
+ # points = np.array(points)
+ # dist = np.sum((points - (mass_y, mass_x))**2, axis=1)
+ # id = np.argsort(dist)[0]
+ # mass_y, mass_x = points[id]
+ # # return
+ # return {'mask': m, 'area': area, 'position_x': position_x, 'position_y': position_y, 'mass_x': mass_x, 'mass_y': mass_y}
+ # # show image and mask
+ # I = io.imread(osp.join(self.IMAGE_DIR, image['file_name']))
+ # plt.figure()
+ # plt.imshow(I)
+ # ax = plt.gca()
+ # img = np.ones( (m.shape[0], m.shape[1], 3) )
+ # color_mask = np.array([2.0,166.0,101.0])/255
+ # for i in range(3):
+ # img[:,:,i] = color_mask[i]
+ # ax.imshow(np.dstack( (img, m*0.5) ))
+ # plt.show()
+
+ def showMask(self, ref):
+ M = self.getMask(ref)
+ msk = M["mask"]
+ ax = plt.gca()
+ ax.imshow(msk)
+
+
+if __name__ == "__main__":
+ refer = REFER(dataset="refcocog", splitBy="google")
+ ref_ids = refer.getRefIds()
+ print(len(ref_ids))
+
+ print(len(refer.Imgs))
+ print(len(refer.imgToRefs))
+
+ ref_ids = refer.getRefIds(split="train")
+ print("There are %s training referred objects." % len(ref_ids))
+
+ for ref_id in ref_ids:
+ ref = refer.loadRefs(ref_id)[0]
+ if len(ref["sentences"]) < 2:
+ continue
+
+ pprint(ref)
+ print("The label is %s." % refer.Cats[ref["category_id"]])
+ plt.figure()
+ refer.showRef(ref, seg_box="box")
+ plt.show()
+
+ # plt.figure()
+ # refer.showMask(ref)
+ # plt.show()
diff --git a/py/evf_sam/utils/refer_seg_dataset.py b/py/evf_sam/utils/refer_seg_dataset.py
new file mode 100644
index 0000000..fc8bdcb
--- /dev/null
+++ b/py/evf_sam/utils/refer_seg_dataset.py
@@ -0,0 +1,254 @@
+import os
+import random
+
+import cv2
+import numpy as np
+import torch
+import torch.nn.functional as F
+from pycocotools import mask
+
+from model.segment_anything.utils.transforms import ResizeLongestSide
+
+from .grefer import G_REFER
+from .refer import REFER
+from torchvision import transforms
+
+
+class ReferSegDataset(torch.utils.data.Dataset):
+ pixel_mean = torch.Tensor([123.675, 116.28, 103.53]).view(-1, 1, 1)
+ pixel_std = torch.Tensor([58.395, 57.12, 57.375]).view(-1, 1, 1)
+ img_size = 1024
+ ignore_label = 255
+
+ def __init__(
+ self,
+ base_image_dir,
+ tokenizer,
+ samples_per_epoch=500 * 8 * 2 * 10,
+ precision: str = "fp32",
+ image_size: int = 224,
+ num_classes_per_sample: int = 3,
+ exclude_val=False,
+ refer_seg_data="refclef||refcoco||refcoco+||refcocog",
+ model_type="ori",
+ transform=ResizeLongestSide(1024),
+ ):
+ self.model_type = model_type
+ self.exclude_val = exclude_val
+ self.samples_per_epoch = samples_per_epoch
+ self.num_classes_per_sample = num_classes_per_sample
+
+ self.base_image_dir = base_image_dir
+ self.tokenizer = tokenizer
+ self.precision = precision
+ self.transform = transform
+ self.image_preprocessor = transforms.Compose([
+ transforms.ToTensor(),
+ transforms.Resize((image_size, image_size), interpolation=3),
+ transforms.Normalize(mean=(0.5, 0.5, 0.5), std=(0.5, 0.5, 0.5))
+ ])
+
+ DATA_DIR = os.path.join(base_image_dir, "refer_seg")
+ self.refer_seg_ds_list = refer_seg_data.split(
+ "||"
+ ) # ['refclef', 'refcoco', 'refcoco+', 'refcocog']
+ self.refer_seg_data = {}
+ for ds in self.refer_seg_ds_list:
+ if ds == "refcocog":
+ splitBy = "umd"
+ else:
+ splitBy = "unc"
+
+ if ds == "grefcoco":
+ refer_api = G_REFER(DATA_DIR, ds, splitBy)
+ else:
+ refer_api = REFER(DATA_DIR, ds, splitBy)
+
+ ref_ids_train = refer_api.getRefIds(split="train")
+ images_ids_train = refer_api.getImgIds(ref_ids=ref_ids_train)
+ refs_train = refer_api.loadRefs(ref_ids=ref_ids_train)
+
+ refer_seg_ds = {}
+ refer_seg_ds["images"] = []
+ loaded_images = refer_api.loadImgs(image_ids=images_ids_train)
+
+ for item in loaded_images:
+ item = item.copy()
+ if ds == "refclef":
+ item["file_name"] = os.path.join(
+ DATA_DIR, "images/saiapr_tc-12", item["file_name"]
+ )
+ else:
+ item["file_name"] = os.path.join(
+ DATA_DIR, "images/mscoco/images/train2014", item["file_name"]
+ )
+ refer_seg_ds["images"].append(item)
+ refer_seg_ds["annotations"] = refer_api.Anns # anns_train
+
+ print(
+ "dataset {} (refs {}) (train split) has {} images and {} annotations.".format(
+ ds,
+ splitBy,
+ len(refer_seg_ds["images"]),
+ len(refer_seg_ds["annotations"]),
+ )
+ )
+
+ img2refs = {}
+ for ref in refs_train:
+ image_id = ref["image_id"]
+ img2refs[image_id] = img2refs.get(image_id, []) + [
+ ref,
+ ]
+ refer_seg_ds["img2refs"] = img2refs
+ self.refer_seg_data[ds] = refer_seg_ds
+
+ def __len__(self):
+ return self.samples_per_epoch
+
+ def preprocess(self, x: torch.Tensor) -> torch.Tensor:
+ """Normalize pixel values and pad to a square input."""
+ if self.model_type=="hq":
+ h, w = x.shape[-2:]
+ padh = self.img_size - h
+ padw = self.img_size - w
+ x = F.pad(x, (0, padw, 0, padh), value=128)
+
+ # Normalize colors
+ x = (x - self.pixel_mean) / self.pixel_std
+
+ if self.model_type=="effi" or self.model_type=="sam2":
+ x = F.interpolate(x.unsqueeze(0), (self.img_size, self.img_size), mode="bilinear").squeeze(0)
+ else:
+ # Pad
+ h, w = x.shape[-2:]
+ padh = self.img_size - h
+ padw = self.img_size - w
+ x = F.pad(x, (0, padw, 0, padh))
+ return x
+
+ def __getitem__(self, idx):
+ ds = random.randint(0, len(self.refer_seg_ds_list) - 1)
+ ds = self.refer_seg_ds_list[ds]
+ refer_seg_ds = self.refer_seg_data[ds]
+ images = refer_seg_ds["images"]
+ annotations = refer_seg_ds["annotations"]
+ img2refs = refer_seg_ds["img2refs"]
+ idx = random.randint(0, len(images) - 1)
+ image_info = images[idx]
+ image_path = image_info["file_name"]
+ image_id = image_info["id"]
+ refs = img2refs[image_id]
+ if len(refs) == 0:
+ return self.__getitem__(0)
+
+ sents = []
+ ann_ids = []
+ for ref in refs:
+ for sent in ref["sentences"]:
+ text = sent["sent"]
+ sents.append(text)
+ ann_ids.append(ref["ann_id"])
+ if len(sents) >= self.num_classes_per_sample:
+ sampled_inds = np.random.choice(
+ list(range(len(sents))), size=self.num_classes_per_sample, replace=False
+ )
+ else:
+ sampled_inds = list(range(len(sents)))
+ sampled_sents = np.vectorize(sents.__getitem__)(sampled_inds).tolist()
+ # sampled_ann_ids = np.vectorize(ann_ids.__getitem__)(sampled_inds).tolist()
+ sampled_ann_ids = [ann_ids[ind] for ind in sampled_inds]
+ sampled_classes = sampled_sents
+ image = cv2.imread(image_path)
+ image = cv2.cvtColor(image, cv2.COLOR_BGR2RGB)
+
+ # preprocess image for evf
+ image_evf = self.image_preprocessor(image)
+
+ image = self.transform.apply_image(image) # preprocess image for sam
+ resize = image.shape[:2]
+
+ image = self.preprocess(torch.from_numpy(image).permute(2, 0, 1).contiguous())
+
+ flag = False
+ masks = []
+ for ann_id in sampled_ann_ids:
+ if isinstance(ann_id, list):
+ flag = True
+ if -1 in ann_id:
+ assert len(ann_id) == 1
+ m = np.zeros((image_info["height"], image_info["width"])).astype(
+ np.uint8
+ )
+ else:
+ m_final = np.zeros(
+ (image_info["height"], image_info["width"])
+ ).astype(np.uint8)
+ for ann_id_i in ann_id:
+ ann = annotations[ann_id_i]
+
+ if len(ann["segmentation"]) == 0:
+ m = np.zeros(
+ (image_info["height"], image_info["width"])
+ ).astype(np.uint8)
+ else:
+ if type(ann["segmentation"][0]) == list: # polygon
+ rle = mask.frPyObjects(
+ ann["segmentation"],
+ image_info["height"],
+ image_info["width"],
+ )
+ else:
+ rle = ann["segmentation"]
+ for i in range(len(rle)):
+ if not isinstance(rle[i]["counts"], bytes):
+ rle[i]["counts"] = rle[i]["counts"].encode()
+ m = mask.decode(rle)
+ m = np.sum(
+ m, axis=2
+ ) # sometimes there are multiple binary map (corresponding to multiple segs)
+ m = m.astype(np.uint8) # convert to np.uint8
+ m_final = m_final | m
+ m = m_final
+ masks.append(m)
+ continue
+
+ ann = annotations[ann_id]
+
+ if len(ann["segmentation"]) == 0:
+ m = np.zeros((image_info["height"], image_info["width"])).astype(
+ np.uint8
+ )
+ masks.append(m)
+ continue
+
+ if type(ann["segmentation"][0]) == list: # polygon
+ rle = mask.frPyObjects(
+ ann["segmentation"], image_info["height"], image_info["width"]
+ )
+ else:
+ rle = ann["segmentation"]
+ for i in range(len(rle)):
+ if not isinstance(rle[i]["counts"], bytes):
+ rle[i]["counts"] = rle[i]["counts"].encode()
+ m = mask.decode(rle)
+ m = np.sum(
+ m, axis=2
+ ) # sometimes there are multiple binary map (corresponding to multiple segs)
+ m = m.astype(np.uint8) # convert to np.uint8
+ masks.append(m)
+
+ masks = np.stack(masks, axis=0)
+
+ masks = torch.from_numpy(masks)
+ label = torch.ones(masks.shape[1], masks.shape[2]) * self.ignore_label
+
+ return (
+ image_path,
+ image,
+ image_evf,
+ masks,
+ label,
+ resize,
+ sampled_classes,
+ )
diff --git a/py/evf_sam/utils/sem_seg_dataset.py b/py/evf_sam/utils/sem_seg_dataset.py
new file mode 100644
index 0000000..b19fd1a
--- /dev/null
+++ b/py/evf_sam/utils/sem_seg_dataset.py
@@ -0,0 +1,290 @@
+import glob
+import json
+import os
+import random
+
+import cv2
+import numpy as np
+import torch
+import torch.nn.functional as F
+from PIL import Image
+from pycocotools.coco import COCO
+
+from model.segment_anything.utils.transforms import ResizeLongestSide
+from torchvision import transforms
+
+def init_mapillary(base_image_dir):
+ mapillary_data_root = os.path.join(base_image_dir, "mapillary")
+ with open(os.path.join(mapillary_data_root, "config_v2.0.json")) as f:
+ mapillary_classes = json.load(f)["labels"]
+ mapillary_classes = [x["readable"].lower() for x in mapillary_classes]
+ mapillary_classes = np.array(mapillary_classes)
+ mapillary_labels = sorted(
+ glob.glob(
+ os.path.join(mapillary_data_root, "training", "v2.0", "labels", "*.png")
+ )
+ )
+ mapillary_images = [
+ x.replace(".png", ".jpg").replace("v2.0/labels", "images")
+ for x in mapillary_labels
+ ]
+ print("mapillary: ", len(mapillary_images))
+ return mapillary_classes, mapillary_images, mapillary_labels
+
+
+def init_ade20k(base_image_dir):
+ with open("utils/ade20k_classes.json", "r") as f:
+ ade20k_classes = json.load(f)
+ ade20k_classes = np.array(ade20k_classes)
+ image_ids = sorted(
+ os.listdir(os.path.join(base_image_dir, "ade20k/images", "training"))
+ )
+ ade20k_image_ids = []
+ for x in image_ids:
+ if x.endswith(".jpg"):
+ ade20k_image_ids.append(x[:-4])
+ ade20k_images = []
+ for image_id in ade20k_image_ids: # self.descriptions:
+ ade20k_images.append(
+ os.path.join(
+ base_image_dir,
+ "ade20k",
+ "images",
+ "training",
+ "{}.jpg".format(image_id),
+ )
+ )
+ ade20k_labels = [
+ x.replace(".jpg", ".png").replace("images", "annotations")
+ for x in ade20k_images
+ ]
+ print("ade20k: ", len(ade20k_images))
+ return ade20k_classes, ade20k_images, ade20k_labels
+
+def init_paco_lvis(base_image_dir):
+ coco_api_paco_lvis = COCO(
+ os.path.join(
+ base_image_dir, "vlpart", "paco", "annotations", "paco_lvis_v1_train.json"
+ )
+ )
+ all_classes = coco_api_paco_lvis.loadCats(coco_api_paco_lvis.getCatIds())
+ class_map_paco_lvis = {}
+ for cat in all_classes:
+ cat_split = cat["name"].strip().split(":")
+ if len(cat_split) == 1:
+ name = cat_split[0].split("_(")[0]
+ else:
+ assert len(cat_split) == 2
+ obj, part = cat_split
+ obj = obj.split("_(")[0]
+ part = part.split("_(")[0]
+ name = (obj, part)
+ class_map_paco_lvis[cat["id"]] = name
+ img_ids = coco_api_paco_lvis.getImgIds()
+ print("paco_lvis: ", len(img_ids))
+ return class_map_paco_lvis, img_ids, coco_api_paco_lvis
+
+
+def init_pascal_part(base_image_dir):
+ coco_api_pascal_part = COCO(
+ os.path.join(base_image_dir, "vlpart", "pascal_part", "train.json")
+ )
+ all_classes = coco_api_pascal_part.loadCats(coco_api_pascal_part.getCatIds())
+ class_map_pascal_part = {}
+ for cat in all_classes:
+ cat_main, cat_part = cat["name"].strip().split(":")
+ name = (cat_main, cat_part)
+ class_map_pascal_part[cat["id"]] = name
+ img_ids = coco_api_pascal_part.getImgIds()
+ print("pascal_part: ", len(img_ids))
+ return class_map_pascal_part, img_ids, coco_api_pascal_part
+
+
+class SemSegDataset(torch.utils.data.Dataset):
+ pixel_mean = torch.Tensor([123.675, 116.28, 103.53]).view(-1, 1, 1)
+ pixel_std = torch.Tensor([58.395, 57.12, 57.375]).view(-1, 1, 1)
+ img_size = 1024
+ ignore_label = 255
+
+ def __init__(
+ self,
+ base_image_dir,
+ tokenizer,
+ samples_per_epoch=500 * 8 * 2 * 10,
+ precision: str = "fp32",
+ image_size: int = 224,
+ num_classes_per_sample: int = 3,
+ exclude_val=False,
+ sem_seg_data="ade20k||pascal_part||mapillary",
+ model_type="ori",
+ transform=ResizeLongestSide(1024),
+ ):
+ self.model_type = model_type
+ self.exclude_val = exclude_val
+ self.samples_per_epoch = samples_per_epoch
+ self.num_classes_per_sample = num_classes_per_sample
+
+ self.base_image_dir = base_image_dir
+ self.tokenizer = tokenizer
+ self.precision = precision
+ self.transform = transform
+ self.image_preprocessor = transforms.Compose([
+ transforms.ToTensor(),
+ transforms.Resize((image_size, image_size), interpolation=3),
+ transforms.Normalize(mean=(0.5, 0.5, 0.5), std=(0.5, 0.5, 0.5))
+ ])
+
+ self.data2list = {}
+ self.data2classes = {}
+
+ self.sem_seg_datas = sem_seg_data.split("||")
+ for ds in self.sem_seg_datas:
+ classes, images, labels = eval("init_{}".format(ds))(base_image_dir)
+ self.data2list[ds] = (images, labels)
+ self.data2classes[ds] = classes
+
+
+ def __len__(self):
+ return self.samples_per_epoch
+
+ def preprocess(self, x: torch.Tensor) -> torch.Tensor:
+ """Normalize pixel values and pad to a square input."""
+ if self.model_type=="hq":
+ h, w = x.shape[-2:]
+ padh = self.img_size - h
+ padw = self.img_size - w
+ x = F.pad(x, (0, padw, 0, padh), value=128)
+
+ # Normalize colors
+ x = (x - self.pixel_mean) / self.pixel_std
+
+ if self.model_type=="effi" or self.model_type=="sam2":
+ x = F.interpolate(x.unsqueeze(0), (self.img_size, self.img_size), mode="bilinear").squeeze(0)
+ else:
+ # Pad
+ h, w = x.shape[-2:]
+ padh = self.img_size - h
+ padw = self.img_size - w
+ x = F.pad(x, (0, padw, 0, padh))
+ return x
+
+ def __getitem__(self, idx):
+ ds = random.randint(0, len(self.sem_seg_datas) - 1)
+ ds = self.sem_seg_datas[ds]
+
+ if ds in ["pascal_part"]:
+ class_map = self.data2classes[ds]
+ img_ids, coco_api = self.data2list[ds]
+ idx = random.randint(0, len(img_ids) - 1)
+ img_id = img_ids[idx]
+ image_info = coco_api.loadImgs([img_id])[0]
+ file_name = image_info["file_name"]
+ file_name = os.path.join(
+ "VOCdevkit", "VOC2010", "JPEGImages", file_name
+ )
+ image_path = os.path.join(self.base_image_dir, "vlpart", ds, file_name)
+
+ image = cv2.imread(image_path)
+ image = cv2.cvtColor(image, cv2.COLOR_BGR2RGB)
+
+ # preprocess image for evf
+ image_evf = self.image_preprocessor(image)
+
+ image = self.transform.apply_image(image) # preprocess image for sam
+ resize = image.shape[:2]
+ annIds = coco_api.getAnnIds(imgIds=image_info["id"])
+ anns = coco_api.loadAnns(annIds)
+ if len(anns) == 0:
+ return self.__getitem__(0)
+ if len(anns) >= self.num_classes_per_sample:
+ sampled_anns = np.random.choice(
+ anns, size=self.num_classes_per_sample, replace=False
+ ).tolist()
+ else:
+ sampled_anns = anns
+ sampled_classes = []
+ for ann in sampled_anns:
+ sampled_cls = class_map[ann["category_id"]]
+ if isinstance(sampled_cls, tuple):
+ obj, part = sampled_cls
+ if random.random() < 0.5:
+ name = obj + " " + part
+ else:
+ name = "the {} of the {}".format(part, obj)
+ else:
+ name = sampled_cls
+ sampled_classes.append(name)
+
+ elif ds in ["ade20k", "mapillary"]:
+ image, labels = self.data2list[ds]
+ idx = random.randint(0, len(image) - 1)
+ image_path = image[idx]
+ label_path = labels[idx]
+ label = Image.open(label_path)
+ label = np.array(label)
+ if ds == "ade20k":
+ label[label == 0] = 255
+ label -= 1
+ label[label == 254] = 255
+
+ img = cv2.imread(image_path)
+ image = cv2.cvtColor(img, cv2.COLOR_BGR2RGB)
+ # preprocess image for evf
+ image_evf = self.image_preprocessor(image)
+ image = self.transform.apply_image(image) # preprocess image for sam
+ resize = image.shape[:2]
+ unique_label = np.unique(label).tolist()
+ if 255 in unique_label:
+ unique_label.remove(255)
+ if len(unique_label) == 0:
+ return self.__getitem__(0)
+
+ classes = [self.data2classes[ds][class_id] for class_id in unique_label]
+ if len(classes) >= self.num_classes_per_sample:
+ sampled_classes = np.random.choice(
+ classes, size=self.num_classes_per_sample, replace=False
+ ).tolist()
+ else:
+ sampled_classes = classes
+
+ class_ids = []
+ for sampled_cls in sampled_classes:
+ assert len(sampled_cls.split("||")) == 1
+
+ if ds in ["paco_lvis", "pascal_part"]:
+ continue
+
+ class_id = self.data2classes[ds].tolist().index(sampled_cls)
+ class_ids.append(class_id)
+
+ image = self.preprocess(torch.from_numpy(image).permute(2, 0, 1).contiguous())
+
+ if ds in ["pascal_part"]:
+ masks = []
+ for ann in sampled_anns:
+ try:
+ masks.append(coco_api.annToMask(ann))
+ except Exception as e:
+ print(e)
+ return self.__getitem__(0)
+
+ masks = np.stack(masks, axis=0)
+ masks = torch.from_numpy(masks)
+ label = torch.ones(masks.shape[1], masks.shape[2]) * self.ignore_label
+
+ else:
+ label = torch.from_numpy(label).long()
+ masks = []
+ for class_id in class_ids:
+ masks.append(label == class_id)
+ masks = torch.stack(masks, dim=0)
+ # sampled_classes = ["all "+_ for _ in sampled_classes]
+ return (
+ image_path,
+ image,
+ image_evf,
+ masks,
+ label,
+ resize,
+ sampled_classes,
+ )
diff --git a/py/evf_sam/utils/utils.py b/py/evf_sam/utils/utils.py
new file mode 100644
index 0000000..4900182
--- /dev/null
+++ b/py/evf_sam/utils/utils.py
@@ -0,0 +1,127 @@
+from enum import Enum
+
+import numpy as np
+import torch
+import torch.distributed as dist
+
+IGNORE_INDEX = -100
+
+class Summary(Enum):
+ NONE = 0
+ AVERAGE = 1
+ SUM = 2
+ COUNT = 3
+
+
+class AverageMeter(object):
+ """Computes and stores the average and current value"""
+
+ def __init__(self, name, fmt=":f", summary_type=Summary.AVERAGE):
+ self.name = name
+ self.fmt = fmt
+ self.summary_type = summary_type
+ self.reset()
+
+ def reset(self):
+ self.val = 0
+ self.avg = 0
+ self.sum = 0
+ self.count = 0
+
+ def update(self, val, n=1):
+ self.val = val
+ self.sum += val * n
+ self.count += n
+ self.avg = self.sum / self.count
+
+ def all_reduce(self):
+ device = "cuda" if torch.cuda.is_available() else "cpu"
+ if isinstance(self.sum, np.ndarray):
+ total = torch.tensor(
+ self.sum.tolist()
+ + [
+ self.count,
+ ],
+ dtype=torch.float32,
+ device=device,
+ )
+ else:
+ total = torch.tensor(
+ [self.sum, self.count], dtype=torch.float32, device=device
+ )
+
+ dist.all_reduce(total, dist.ReduceOp.SUM, async_op=False)
+ if total.shape[0] > 2:
+ self.sum, self.count = total[:-1].cpu().numpy(), total[-1].cpu().item()
+ else:
+ self.sum, self.count = total.tolist()
+ self.avg = self.sum / (self.count + 1e-5)
+
+ def __str__(self):
+ fmtstr = "{name} {val" + self.fmt + "} ({avg" + self.fmt + "})"
+ return fmtstr.format(**self.__dict__)
+
+ def summary(self):
+ fmtstr = ""
+ if self.summary_type is Summary.NONE:
+ fmtstr = ""
+ elif self.summary_type is Summary.AVERAGE:
+ fmtstr = "{name} {avg:.3f}"
+ elif self.summary_type is Summary.SUM:
+ fmtstr = "{name} {sum:.3f}"
+ elif self.summary_type is Summary.COUNT:
+ fmtstr = "{name} {count:.3f}"
+ else:
+ raise ValueError("invalid summary type %r" % self.summary_type)
+
+ return fmtstr.format(**self.__dict__)
+
+
+def intersectionAndUnionGPU(output, target, K, ignore_index=255):
+ # 'K' classes, output and target sizes are N or N * L or N * H * W, each value in range 0 to K - 1.
+ assert output.dim() in [1, 2, 3]
+ assert output.shape == target.shape
+ output = output.view(-1)
+ target = target.view(-1)
+ output[target == ignore_index] = ignore_index
+ intersection = output[output == target]
+ area_intersection = torch.histc(intersection, bins=K, min=0, max=K - 1)
+ area_output = torch.histc(output, bins=K, min=0, max=K - 1)
+ area_target = torch.histc(target, bins=K, min=0, max=K - 1)
+ area_union = area_output + area_target - area_intersection
+ return area_intersection, area_union, area_target
+
+
+class ProgressMeter(object):
+ def __init__(self, num_batches, meters, prefix=""):
+ self.batch_fmtstr = self._get_batch_fmtstr(num_batches)
+ self.meters = meters
+ self.prefix = prefix
+
+ def display(self, batch):
+ entries = [self.prefix + self.batch_fmtstr.format(batch)]
+ entries += [str(meter) for meter in self.meters]
+ print("\t".join(entries))
+
+ def display_summary(self):
+ entries = [" *"]
+ entries += [meter.summary() for meter in self.meters]
+ print(" ".join(entries))
+
+ def _get_batch_fmtstr(self, num_batches):
+ num_digits = len(str(num_batches // 1))
+ fmt = "{:" + str(num_digits) + "d}"
+ return "[" + fmt + "/" + fmt.format(num_batches) + "]"
+
+
+def dict_to_cuda(input_dict):
+ for k, v in input_dict.items():
+ if isinstance(input_dict[k], torch.Tensor):
+ input_dict[k] = v.cuda(non_blocking=True)
+ elif (
+ isinstance(input_dict[k], list)
+ and len(input_dict[k]) > 0
+ and isinstance(input_dict[k][0], torch.Tensor)
+ ):
+ input_dict[k] = [ele.cuda(non_blocking=True) for ele in v]
+ return input_dict
diff --git a/py/evf_sam_ultra.py b/py/evf_sam_ultra.py
new file mode 100644
index 0000000..72e1da6
--- /dev/null
+++ b/py/evf_sam_ultra.py
@@ -0,0 +1,120 @@
+# layerstyle advance
+
+'''
+推理部分代码来自https://github.com/hustvl/EVF-SAM
+'''
+import sys
+
+from .imagefunc import *
+sys.path.append(os.path.join(os.path.dirname(__file__), 'evf_sam'))
+from evf_sam.evf_sam_inference import evf_sam_main
+class EVF_SAM_Ultra:
+
+ def __init__(self):
+ self.NODE_NAME = 'EVF_SAM Ultra'
+ pass
+
+ @classmethod
+ def INPUT_TYPES(cls):
+ # model_list = ["evf-sam2","evf-sam", "evf-sam2-multitask", "evf-sam-multitask"]
+ model_list = ["evf-sam2", "evf-sam"]
+ precision_list = ["fp16", "bf16", "fp32"]
+ load_in_bit_list = ["full", "8", "4"]
+ method_list = ['VITMatte', 'VITMatte(local)', 'PyMatting', 'GuidedFilter', ]
+ device_list = ['cuda', 'cpu']
+ return {"required":
+ {
+ "image": ("IMAGE",),
+ "model": (model_list,),
+ "precision": (precision_list,),
+ "load_in_bit": (load_in_bit_list,),
+ "prompt": ("STRING", {"default": "subject"}),
+ "detail_method": (method_list,),
+ "detail_erode": ("INT", {"default": 6, "min": 1, "max": 255, "step": 1}),
+ "detail_dilate": ("INT", {"default": 4, "min": 1, "max": 255, "step": 1}),
+ "black_point": ("FLOAT", {"default": 0.01, "min": 0.01, "max": 0.98, "step": 0.01, "display": "slider"}),
+ "white_point": ("FLOAT", {"default": 0.99, "min": 0.02, "max": 0.99, "step": 0.01, "display": "slider"}),
+ "process_detail": ("BOOLEAN", {"default": True}),
+ "device": (device_list,),
+ "max_megapixels": ("FLOAT", {"default": 2.0, "min": 1, "max": 999, "step": 0.1}),
+ }
+ }
+
+ RETURN_TYPES = ("IMAGE", "MASK",)
+ RETURN_NAMES = ("image", "mask",)
+ FUNCTION = "evf_sam_ultra"
+ CATEGORY = '😺dzNodes/LayerMask'
+
+ def evf_sam_ultra(self, image, model, precision, load_in_bit, prompt,
+ detail_method, detail_erode, detail_dilate, black_point, white_point,
+ process_detail, device, max_megapixels,
+ ):
+
+ ret_images = []
+ ret_masks = []
+
+ if detail_method == 'VITMatte(local)':
+ local_files_only = True
+ else:
+ local_files_only = False
+
+ if model == 'evf-sam2' or model == 'evf-sam2-multitask':
+ model_type = 'sam2'
+ elif model == 'evf-sam' or model == 'evf-sam-multitask':
+ model_type = 'ori'
+ else:
+ model_type = 'effi'
+
+ if load_in_bit == 'full':
+ load_in_bit = 16
+ else:
+ load_in_bit = int(load_in_bit)
+
+ model_path = ""
+ model_folder_name = 'EVF-SAM'
+ try:
+ model_path = os.path.join(
+ os.path.normpath(folder_paths.folder_names_and_paths[model_folder_name][0][0]), model)
+ except:
+ pass
+ if not os.path.exists(model_path):
+ model_path = os.path.join(folder_paths.models_dir, model_folder_name, model)
+
+ for i in image:
+ i = torch.unsqueeze(i, 0)
+ orig_image = tensor2pil(i).convert('RGB')
+ sys.path.append(os.path.dirname(os.path.abspath(__file__)))
+
+ mask_image = evf_sam_main(model_path, model_type, precision, load_in_bit, orig_image, prompt)
+ _mask = pil2tensor(mask_image)
+
+ detail_range = detail_erode + detail_dilate
+ if process_detail:
+ if detail_method == 'GuidedFilter':
+ _mask = guided_filter_alpha(i, _mask, detail_range // 6 + 1)
+ _mask = tensor2pil(histogram_remap(_mask, black_point, white_point))
+ elif detail_method == 'PyMatting':
+ _mask = tensor2pil(mask_edge_detail(i, _mask, detail_range // 8 + 1, black_point, white_point))
+ else:
+ _trimap = generate_VITMatte_trimap(_mask, detail_erode, detail_dilate)
+ _mask = generate_VITMatte(orig_image, _trimap, local_files_only=local_files_only, device=device,
+ max_megapixels=max_megapixels)
+ _mask = tensor2pil(histogram_remap(pil2tensor(_mask), black_point, white_point))
+ else:
+ _mask = mask2image(_mask)
+
+ ret_image = RGB2RGBA(orig_image, _mask.convert('L'))
+ ret_images.append(pil2tensor(ret_image))
+ ret_masks.append(image2mask(_mask))
+
+ log(f"{self.NODE_NAME} Processed {len(ret_images)} image(s).", message_type='finish')
+ return (torch.cat(ret_images, dim=0), torch.cat(ret_masks, dim=0),)
+
+NODE_CLASS_MAPPINGS = {
+ "LayerMask: EVFSAMUltra": EVF_SAM_Ultra
+}
+
+NODE_DISPLAY_NAME_MAPPINGS = {
+ "LayerMask: EVFSAMUltra": "LayerMask: EVF-SAM Ultra(Advance)"
+}
+
diff --git a/py/florence2_ultra.py b/py/florence2_ultra.py
new file mode 100644
index 0000000..4749f99
--- /dev/null
+++ b/py/florence2_ultra.py
@@ -0,0 +1,615 @@
+# layerstyle advance
+
+import io
+from unittest.mock import patch
+import matplotlib.pyplot as plt
+import matplotlib.patches as patches
+import colorsys
+from transformers.dynamic_module_utils import get_imports
+import comfy.model_management
+from .imagefunc import *
+
+colormap = ['blue', 'orange', 'green', 'purple', 'brown', 'pink', 'gray', 'olive', 'cyan', 'red',
+ 'lime', 'indigo', 'violet', 'aqua', 'magenta', 'coral', 'gold', 'tan', 'skyblue']
+
+device = comfy.model_management.get_torch_device()
+
+fl2_model_repos = {
+ "base": "microsoft/Florence-2-base",
+ "base-ft": "microsoft/Florence-2-base-ft",
+ "large": "microsoft/Florence-2-large",
+ "large-ft": "microsoft/Florence-2-large-ft",
+ "DocVQA": "HuggingFaceM4/Florence-2-DocVQA",
+ "SD3-Captioner": "gokaygokay/Florence-2-SD3-Captioner",
+ "base-PromptGen": "MiaoshouAI/Florence-2-base-PromptGen",
+ "CogFlorence-2-Large-Freeze": "thwri/CogFlorence-2-Large-Freeze",
+ "CogFlorence-2.1-Large": "thwri/CogFlorence-2.1-Large",
+ "base-PromptGen-v1.5":"MiaoshouAI/Florence-2-base-PromptGen-v1.5",
+ "large-PromptGen-v1.5":"MiaoshouAI/Florence-2-large-PromptGen-v1.5",
+ "base-PromptGen-v2.0":"MiaoshouAI/Florence-2-base-PromptGen-v2.0",
+ "large-PromptGen-v2.0":"MiaoshouAI/Florence-2-large-PromptGen-v2.0"
+}
+
+def fixed_get_imports(filename) -> list[str]:
+ """Workaround for FlashAttention"""
+ if os.path.basename(filename) != "modeling_florence2.py":
+ return get_imports(filename)
+ imports = get_imports(filename)
+ try:
+ imports.remove("flash_attn")
+ except:
+ pass
+ return imports
+
+def load_model(version):
+ florence_path = os.path.join(folder_paths.models_dir, "florence2")
+ os.makedirs(florence_path, exist_ok=True)
+
+ model_path = os.path.join(florence_path, version)
+ attention = 'sdpa'
+
+ if not os.path.exists(model_path):
+ log(f"Downloading Florence2 {version} model...")
+ repo_id = fl2_model_repos[version]
+ from huggingface_hub import snapshot_download
+ snapshot_download(repo_id=repo_id, local_dir=model_path, ignore_patterns=["*.md", "*.txt"])
+
+ try:
+ with patch("transformers.dynamic_module_utils.get_imports", fixed_get_imports):
+ # model = AutoModelForCausalLM.from_pretrained(model_path, trust_remote_code=True)
+ model = AutoModelForCausalLM.from_pretrained(model_path, attn_implementation=attention, device_map=device,
+ torch_dtype=torch.float32, trust_remote_code=True)
+ processor = AutoProcessor.from_pretrained(model_path, trust_remote_code=True)
+ except Exception as e:
+ try:
+ model = AutoModelForCausalLM.from_pretrained(model_path, attn_implementation=attention, device_map=device,
+ torch_dtype=torch.float32, trust_remote_code=True)
+ processor = AutoTokenizer.from_pretrained(model_path, trust_remote_code=True)
+ except Exception as e:
+ sys.path.append(model_path)
+ # Import the Florence modules
+ if version == 'large-PromptGen-v1.5':
+ from florence2_large.modeling_florence2 import Florence2ForConditionalGeneration
+ from florence2_large.configuration_florence2 import Florence2Config
+ elif version == 'base-PromptGen-v1.5':
+ from florence2_base_ft.modeling_florence2 import Florence2ForConditionalGeneration
+ from florence2_base_ft.configuration_florence2 import Florence2Config
+ else:
+ log(f"Error loading model or tokenizer: {str(e)}", message_type='error')
+ return (None, None)
+
+ # Load the model configuration
+ model_config = Florence2Config.from_pretrained(model_path)
+ # Load the model
+ with patch("transformers.dynamic_module_utils.get_imports", fixed_get_imports):
+ model = Florence2ForConditionalGeneration.from_pretrained(
+ model_path,
+ config=model_config,
+ attn_implementation=attention,
+ device_map=device
+ ).to(device)
+
+ processor = AutoProcessor.from_pretrained(model_path, trust_remote_code=True)
+
+ return (model.to(device), processor)
+
+def fig_to_pil(fig):
+ buf = io.BytesIO()
+ fig.savefig(buf, format='png', dpi=100, bbox_inches='tight', pad_inches=0)
+ buf.seek(0)
+ pil = Image.open(buf)
+ plt.close()
+ return pil
+
+def plot_bbox(image, data):
+ fig, ax = plt.subplots()
+ fig.set_size_inches(image.width / 100, image.height / 100)
+ ax.imshow(image)
+ for i, (bbox, label) in enumerate(zip(data['bboxes'], data['labels'])):
+ x1, y1, x2, y2 = bbox
+ rect = patches.Rectangle((x1, y1), x2 - x1, y2 - y1, linewidth=1, edgecolor='r', facecolor='none')
+ ax.add_patch(rect)
+ enum_label = f"{i}: {label}"
+ plt.text(x1 + 7, y1 + 17, enum_label, color='white', fontsize=8, bbox=dict(facecolor='red', alpha=0.5))
+ ax.axis('off')
+ return fig
+
+def generate_color(index, total_colors=25):
+ # Generate color by varying the hue to maximize difference between colors
+ hue = (index / total_colors) % 1.0 # Normalize hue to be between 0 and 1
+ saturation = 0.65 # Keep saturation constant
+ lightness = 0.5 # Keep lightness constant
+
+ # Convert HSL to RGB, then to hexadecimal
+ r, g, b = colorsys.hls_to_rgb(hue, lightness, saturation)
+ return f'#{int(r * 255):02X}{int(g * 255):02X}{int(b * 255):02X}'
+
+def plot_mask_bbox(image, data):
+ fig, ax = plt.subplots()
+ fig.set_size_inches(image.width / 100, image.height / 100)
+ ax.imshow(image)
+ num_bboxes = len(data['bboxes'])
+ for i, (bbox, label) in enumerate(list(zip(data['bboxes'], data['labels']))[1:], start=1):
+ x1, y1, x2, y2 = bbox
+ if x2 < x1:
+ x1, y1, x2, y2 = x2, y2, x1, y1
+ color = generate_color(i, total_colors=num_bboxes)
+ rect = patches.Rectangle((x1, y1), x2 - x1, y2 - y1, linewidth=1, edgecolor=color, facecolor='none')
+ ax.add_patch(rect)
+ enum_label = f"{i}: {label}"
+ plt.text(x1 + 7, y1 + 17, enum_label, color='white', fontsize=8, bbox=dict(facecolor=color, alpha=0.5))
+ ax.axis('off')
+ return fig
+
+def plot_mask(image, data, indexes):
+ # Create a black background image (mode "1" for binary, "L" for grayscale)
+ mask = Image.new("L", (image.width, image.height), 0) # Black background
+ fig, ax = plt.subplots()
+ fig.set_size_inches(mask.width / 100, mask.height / 100)
+ ax.imshow(mask, cmap='gray') # Display the mask in grayscale
+ ax.set_facecolor('black') # Set the axes background to black
+ fig.patch.set_facecolor('black') # Set the figure background to black
+ for i, (bbox, label) in enumerate(list(zip(data['bboxes'], data['labels']))[1:], start=1):
+ x1, y1, x2, y2 = bbox
+ if x2 < x1:
+ x1, y1, x2, y2 = x2, y2, x1, y1
+ rect = patches.Rectangle((x1, y1), x2 - x1, y2 - y1, linewidth=1, edgecolor='w', facecolor='w')
+ if i in indexes:
+ ax.add_patch(rect)
+ ax.axis('off')
+ return fig
+
+def draw_polygons(image, prediction, fill_mask=False):
+ output_image = copy.deepcopy(image)
+ draw = ImageDraw.Draw(output_image)
+ scale = 1
+ for polygons, label in zip(prediction['polygons'], prediction['labels']):
+ color = random.choice(colormap)
+ fill_color = color if fill_mask else None
+ for _polygon in polygons:
+ _polygon = np.array(_polygon).reshape(-1, 2)
+ if len(_polygon) < 3:
+ print('Invalid polygon:', _polygon)
+ continue
+ _polygon = (_polygon * scale).reshape(-1).tolist()
+ if fill_mask:
+ draw.polygon(_polygon, outline=color, fill=fill_color)
+ else:
+ draw.polygon(_polygon, outline=color)
+ draw.text((_polygon[0] + 8, _polygon[1] + 2), label, fill=color)
+ return output_image
+
+
+def convert_to_od_format(data):
+ od_results = {
+ 'bboxes': data.get('bboxes', []),
+ 'labels': data.get('bboxes_labels', [])
+ }
+ return od_results
+
+
+def draw_ocr_bboxes(image, prediction):
+ scale = 1
+ output_image = copy.deepcopy(image)
+ draw = ImageDraw.Draw(output_image)
+ bboxes, labels = prediction['quad_boxes'], prediction['labels']
+ for box, label in zip(bboxes, labels):
+ color = random.choice(colormap)
+ new_box = (np.array(box) * scale).tolist()
+ draw.polygon(new_box, width=3, outline=color)
+ draw.text((new_box[0] + 8, new_box[1] + 2),
+ "{}".format(label),
+ align="right",
+ fill=color)
+ return output_image
+
+
+def run_example(model, processor, task_prompt, image, max_new_tokens, num_beams, do_sample, text_input=None):
+ if text_input is None:
+ prompt = task_prompt
+ else:
+ prompt = task_prompt + text_input
+ inputs = processor(text=prompt, images=image, return_tensors="pt").to(device)
+ generated_ids = model.generate(
+ input_ids=inputs["input_ids"],
+ pixel_values=inputs["pixel_values"],
+ max_new_tokens=max_new_tokens,
+ early_stopping=False,
+ do_sample=do_sample,
+ num_beams=num_beams,
+ )
+ generated_text = processor.batch_decode(generated_ids, skip_special_tokens=False)[0]
+ parsed_answer = processor.post_process_generation(
+ generated_text,
+ task=task_prompt,
+ image_size=(image.width, image.height)
+ )
+ return parsed_answer
+
+
+def process_image(model, processor, image, task_prompt, max_new_tokens, num_beams, do_sample, fill_mask, text_input=None):
+ if task_prompt == 'caption':
+ task_prompt = ''
+ result = run_example(model, processor, task_prompt, image, max_new_tokens, num_beams, do_sample)
+ return result[task_prompt], None
+ elif task_prompt == 'detailed caption':
+ task_prompt = ''
+ result = run_example(model, processor, task_prompt, image, max_new_tokens, num_beams, do_sample)
+ return result[task_prompt], None
+ elif task_prompt == 'more detailed caption':
+ task_prompt = ''
+ result = run_example(model, processor, task_prompt, image, max_new_tokens, num_beams, do_sample)
+ return result[task_prompt], None
+ elif task_prompt == 'object detection':
+ task_prompt = ''
+ results = run_example(model, processor, task_prompt, image, max_new_tokens, num_beams, do_sample)
+ fig = plot_bbox(image, results[''])
+ return results[task_prompt], fig_to_pil(fig)
+ elif task_prompt == 'dense region caption':
+ task_prompt = ''
+ results = run_example(model, processor, task_prompt, image, max_new_tokens, num_beams, do_sample)
+ fig = plot_bbox(image, results[''])
+ return results[task_prompt], fig_to_pil(fig)
+ elif task_prompt == 'region proposal':
+ task_prompt = ''
+ results = run_example(model, processor, task_prompt, image, max_new_tokens, num_beams, do_sample)
+ fig = plot_bbox(image, results[''])
+ return results[task_prompt], fig_to_pil(fig)
+ elif task_prompt == 'region proposal (mask)':
+ task_prompt = ''
+ results = run_example(model, processor, task_prompt, image, max_new_tokens, num_beams, do_sample)
+ indexes = []
+ if isinstance(text_input, str):
+ for i in text_input.split(','):
+ try:
+ indexes.append(int(i))
+ except ValueError:
+ print(f"{i} is nit an instance of int")
+ if len(indexes) > 0:
+ fig = plot_mask(image, results[''], indexes)
+ pil = fig_to_pil(fig).resize((image.width, image.height), Image.Resampling.LANCZOS)
+ else:
+ fig = plot_mask_bbox(image, results[''])
+ pil = fig_to_pil(fig)
+ return results[task_prompt], pil
+ elif task_prompt == 'caption to phrase grounding':
+ task_prompt = ''
+ results = run_example(model, processor, task_prompt, image, max_new_tokens, num_beams, do_sample, text_input)
+ fig = plot_bbox(image, results[''])
+ return results[task_prompt], fig_to_pil(fig)
+ elif task_prompt == 'referring expression segmentation':
+ task_prompt = ''
+ results = run_example(model, processor, task_prompt, image, max_new_tokens, num_beams, do_sample, text_input)
+ output_image = draw_polygons(image, results[''], fill_mask)
+ return results[task_prompt], output_image
+ elif task_prompt == 'region to segmentation':
+ task_prompt = ''
+ results = run_example(model, processor, task_prompt, image, max_new_tokens, num_beams, do_sample, text_input)
+ output_image = draw_polygons(image, results[''], fill_mask)
+ return results[task_prompt], output_image
+ elif task_prompt == 'open vocabulary detection':
+ task_prompt = ''
+ results = run_example(model, processor, task_prompt, image, max_new_tokens, num_beams, do_sample, text_input)
+ bbox_results = convert_to_od_format(results[''])
+ fig = plot_bbox(image, bbox_results)
+ return bbox_results, fig_to_pil(fig)
+ elif task_prompt == 'region to category':
+ task_prompt = ''
+ results = run_example(model, processor, task_prompt, image, max_new_tokens, num_beams, do_sample, text_input)
+ return results[task_prompt], None
+ elif task_prompt == 'region to description':
+ task_prompt = ''
+ results = run_example(model, processor, task_prompt, image, max_new_tokens, num_beams, do_sample, text_input)
+ return results[task_prompt], None
+ elif task_prompt == 'OCR':
+ task_prompt = ''
+ result = run_example(model, processor, task_prompt, image, max_new_tokens, num_beams, do_sample)
+ return result[task_prompt], None
+ elif task_prompt == 'OCR with region':
+ task_prompt = ''
+ results = run_example(model, processor, task_prompt, image, max_new_tokens, num_beams, do_sample)
+ output_image = draw_ocr_bboxes(image, results[''])
+ output_results = {'bboxes': results[task_prompt].get('quad_boxes', []),
+ 'labels': results[task_prompt].get('labels', [])}
+ return output_results, output_image
+ # gokaygokay/Florence-2-SD3-Captioner task
+ elif task_prompt == 'description':
+ task_prompt = ''
+ result = run_example(model, processor, task_prompt, image, max_new_tokens, num_beams, do_sample)
+ return result[task_prompt], None
+ # MiaoshouAI/Florence-2-large-PromptGen-v1.5 task
+ elif task_prompt == 'generate tags(PromptGen 1.5)':
+ task_prompt = ''
+ result = run_example(model, processor, task_prompt, image, max_new_tokens, num_beams, do_sample)
+ return result[task_prompt], None
+ elif task_prompt == 'mixed caption(PromptGen 1.5)':
+ task_prompt = ''
+ result = run_example(model, processor, task_prompt, image, max_new_tokens, num_beams, do_sample)
+ return result[task_prompt], None
+ elif task_prompt == 'mixed caption plus(PromptGen 2.0)':
+ task_prompt = ''
+ result = run_example(model, processor, task_prompt, image, max_new_tokens, num_beams, do_sample)
+ return result[task_prompt], None
+ elif task_prompt == 'analyze(PromptGen 2.0)':
+ task_prompt = '<>'
+ result = run_example(model, processor, task_prompt, image, max_new_tokens, num_beams, do_sample)
+ return result[task_prompt], None
+
+ else:
+ return "", None # Return empty string and None for unknown task prompts
+
+
+def remove_angle_bracket_content(text):
+ import re
+ # 正则表达式匹配 "<>" 包围的内容,包括尖括号本身
+ pattern = r'<[^>]*>'
+ # 使用 re.sub 替换匹配的内容为空字符串
+ cleaned_text = re.sub(pattern, '', text)
+ return cleaned_text
+
+
+def decode_f_bboxes(F_BBOXES):
+ if isinstance(F_BBOXES, str):
+ return (torch.zeros(1, 512, 512, dtype=torch.float32), F_BBOXES)
+
+ width = F_BBOXES["width"]
+ height = F_BBOXES["height"]
+ mask = np.zeros((height, width), dtype=np.uint8)
+
+ x1_c = width
+ y1_c = height
+ x2_c = y2_c = 0
+ label = ""
+ if "bboxes" in F_BBOXES:
+ for idx in range(len(F_BBOXES["bboxes"])):
+ bbox = F_BBOXES["bboxes"][idx]
+
+ new_label = F_BBOXES["labels"][idx].removeprefix("")
+ if new_label not in label:
+ if idx > 0:
+ label = label + ", "
+ label = label + new_label
+
+ if len(bbox) == 4:
+ x1, y1, x2, y2 = int(bbox[0]), int(bbox[1]), int(bbox[2]), int(bbox[3])
+ elif len(bbox) == 8:
+ x1 = int(min(bbox[0::2]))
+ x2 = int(max(bbox[0::2]))
+ y1 = int(min(bbox[1::2]))
+ y2 = int(max(bbox[1::2]))
+ else:
+ continue
+
+ x1_c = min(x1_c, x1)
+ y1_c = min(y1_c, y1)
+ x2_c = max(x2_c, x2)
+ y2_c = max(y2_c, y2)
+
+ mask[y1:y2, x1:x2] = 1
+
+ else:
+ image = Image.new('RGB', (width, height), color='black')
+ draw = ImageDraw.Draw(image)
+
+ x1_c = width
+ y1_c = height
+ x2_c = y2_c = 0
+
+ for polygon in F_BBOXES["polygons"][0]:
+ _polygon = np.array(polygon).reshape(-1, 2)
+ if len(_polygon) < 3:
+ print('Invalid polygon:', _polygon)
+ continue
+
+ draw.polygon(_polygon.flatten().tolist(), outline='white', fill='white')
+
+ x1_c = min(x1_c, int(min(polygon[0::2])))
+ x2_c = max(x2_c, int(max(polygon[0::2])))
+ y1_c = min(y1_c, int(min(polygon[1::2])))
+ y2_c = max(y2_c, int(max(polygon[1::2])))
+
+ mask = np.asarray(image)[..., 0].astype(np.float32) / 255
+
+ mask = torch.from_numpy(mask.astype(np.float32)).unsqueeze(0)
+ # label = remove_angle_bracket_content(label)
+ return (mask, label)
+
+
+class LS_LoadFlorence2Model:
+ def __init__(self):
+ self.model = None
+ self.processor = None
+ self.version = None
+
+ @classmethod
+ def INPUT_TYPES(s):
+ model_list = list(fl2_model_repos.keys())
+ return {
+ "required": {
+ "version": (model_list,{"default": model_list[0]}),
+ },
+ }
+
+ RETURN_TYPES = ("FLORENCE2",)
+ RETURN_NAMES = ("florence2_model",)
+ FUNCTION = "load"
+ CATEGORY = '😺dzNodes/LayerMask'
+
+ def load(self, version):
+ if self.version != version:
+ self.model, self.processor = load_model(version)
+ self.version = version
+
+ return ({'model': self.model, 'processor': self.processor, 'version': self.version, 'device': device},)
+
+
+class Florence2Ultra:
+ def __init__(self):
+ self.NODE_NAME = 'Florence2Ultra'
+
+ @classmethod
+ def INPUT_TYPES(s):
+ segment_task_list = [
+ "region to segmentation",
+ "referring expression segmentation",
+ "open vocabulary detection",
+ ]
+ method_list = ['VITMatte', 'VITMatte(local)', 'PyMatting', 'GuidedFilter', ]
+ device_list = ['cuda','cpu']
+ return {
+ "required": {
+ "florence2_model": ("FLORENCE2",),
+ "image": ("IMAGE",),
+ "task": (segment_task_list,{"default": segment_task_list[0]}),
+ "text_input": ("STRING", {"default": "subject"}),
+ "detail_method": (method_list,),
+ "detail_erode": ("INT", {"default": 6, "min": 1, "max": 255, "step": 1}),
+ "detail_dilate": ("INT", {"default": 6, "min": 1, "max": 255, "step": 1}),
+ "black_point": ("FLOAT", {"default": 0.01, "min": 0.01, "max": 0.98, "step": 0.01, "display": "slider"}),
+ "white_point": ("FLOAT", {"default": 0.99, "min": 0.02, "max": 0.99, "step": 0.01, "display": "slider"}),
+ "process_detail": ("BOOLEAN", {"default": True}),
+ "device": (device_list,),
+ "max_megapixels": ("FLOAT", {"default": 2.0, "min": 1, "max": 999, "step": 0.1}),
+ },
+ }
+
+ RETURN_TYPES = ("IMAGE", "MASK",)
+ RETURN_NAMES = ("image", "mask",)
+ FUNCTION = "florence2_ultra"
+ CATEGORY = '😺dzNodes/LayerMask'
+
+ def florence2_ultra(self, florence2_model, image, task, text_input,
+ detail_method, detail_erode, detail_dilate,
+ black_point, white_point, process_detail, device, max_megapixels):
+ max_new_tokens = 512
+ num_beams = 3
+ do_sample = False
+ fill_mask = False
+
+ ret_images = []
+ ret_masks = []
+
+ if detail_method == 'VITMatte(local)':
+ local_files_only = True
+ else:
+ local_files_only = False
+
+ model = florence2_model['model']
+ processor = florence2_model['processor']
+
+ for i in image:
+ img = tensor2pil(i).convert("RGB")
+
+ results, _ = process_image(model, processor, img, task,
+ max_new_tokens, num_beams, do_sample,
+ fill_mask, text_input)
+
+ if isinstance(results, dict):
+ results["width"] = img.width
+ results["height"] = img.height
+
+ _mask, _ = decode_f_bboxes(results)
+
+ if process_detail:
+ detail_range = detail_erode + detail_dilate
+ if detail_method == 'GuidedFilter':
+ _mask = guided_filter_alpha(i, _mask, detail_range // 6 + 1)
+ _mask = tensor2pil(histogram_remap(_mask, black_point, white_point))
+ elif detail_method == 'PyMatting':
+ _mask = tensor2pil(mask_edge_detail(i, _mask, detail_range // 8 + 1, black_point, white_point))
+ else:
+ _trimap = generate_VITMatte_trimap(_mask, detail_erode, detail_dilate)
+ _mask = generate_VITMatte(img, _trimap, local_files_only=local_files_only, device=device, max_megapixels=max_megapixels)
+ _mask = tensor2pil(histogram_remap(pil2tensor(_mask), black_point, white_point))
+ else:
+ _mask = tensor2pil(_mask)
+
+ ret_image = RGB2RGBA(img, _mask.convert('L'))
+ ret_images.append(pil2tensor(ret_image))
+ ret_masks.append(image2mask(_mask))
+
+ return (torch.cat(ret_images, dim=0), torch.cat(ret_masks, dim=0),)
+
+
+class Florence2Image2Prompt:
+
+ def __init__(self):
+ self.NODE_NAME = 'Florence2Image2Prompt'
+
+ @classmethod
+ def INPUT_TYPES(s):
+ caption_task_list = [
+ "caption",
+ "detailed caption",
+ "more detailed caption",
+ 'description',
+ 'generate tags(PromptGen 1.5)',
+ 'mixed caption(PromptGen 1.5)',
+ 'mixed caption plus(PromptGen 2.0)',
+ 'analyze(PromptGen 2.0)',
+ "object detection",
+ "dense region caption",
+ "region proposal",
+ "region proposal (mask)",
+ "caption to phrase grounding",
+ "open vocabulary detection",
+ "region to category",
+ "region to description",
+ "OCR",
+ "OCR with region",
+ ]
+ return {
+ "required": {
+ "florence2_model": ("FLORENCE2",),
+ "image": ("IMAGE",),
+ "task": (caption_task_list,{"default": caption_task_list[2]}),
+ "text_input": ("STRING", {"default": ""}),
+ "max_new_tokens": ("INT", {"default": 1024, "step": 1}),
+ "num_beams": ("INT", {"default": 3, "min": 1, "step": 1}),
+ "do_sample": ('BOOLEAN', {"default": False}),
+ "fill_mask": ('BOOLEAN', {"default": False}),
+ },
+ }
+
+ RETURN_TYPES = ("STRING", "IMAGE",)
+ RETURN_NAMES = ("text", "preview_image",)
+ FUNCTION = "florence2_image2prompt"
+ CATEGORY = '😺dzNodes/LayerUtility/Prompt'
+
+ def florence2_image2prompt(self, florence2_model, image, task, text_input,
+ max_new_tokens, num_beams, do_sample, fill_mask):
+
+ model = florence2_model['model']
+ processor = florence2_model['processor']
+
+ img = tensor2pil(image[0])
+ caption = ""
+ results, output_image = process_image(model, processor, img, task, max_new_tokens, num_beams,
+ do_sample, fill_mask,
+ text_input)
+
+ if isinstance(results, dict):
+ results["width"] = img.width
+ results["height"] = img.height
+
+ if output_image == None:
+ output_image = image[0].detach().clone().unsqueeze(0)
+ else:
+ output_image = np.asarray(output_image).astype(np.float32) / 255
+ output_image = torch.from_numpy(output_image).unsqueeze(0)
+
+ _, caption = decode_f_bboxes(results)
+
+ return (remove_angle_bracket_content(caption), output_image,)
+
+NODE_CLASS_MAPPINGS = {
+ "LayerMask: Florence2Ultra": Florence2Ultra,
+ "LayerMask: LoadFlorence2Model": LS_LoadFlorence2Model,
+ "LayerUtility: Florence2Image2Prompt": Florence2Image2Prompt
+}
+
+NODE_DISPLAY_NAME_MAPPINGS = {
+ "LayerMask: Florence2Ultra": "LayerMask: Florence2 Ultra(Advance)",
+ "LayerMask: LoadFlorence2Model": "LayerMask: Load Florence2 Model(Advance)",
+ "LayerUtility: Florence2Image2Prompt": "LayerUtility: Florence2 Image2Prompt(Advance)"
+}
diff --git a/py/get_color_tone.py b/py/get_color_tone.py
new file mode 100644
index 0000000..52404d5
--- /dev/null
+++ b/py/get_color_tone.py
@@ -0,0 +1,46 @@
+# layer style advance
+import torch
+from .imagefunc import log, tensor2pil, gaussian_blur, get_image_color_tone, get_image_color_average, RGB_to_HSV, Hex_to_RGB
+
+class GetColorTone:
+
+ def __init__(self):
+ self.NODE_NAME = 'GetColorTone'
+
+ @classmethod
+ def INPUT_TYPES(self):
+ mode_list = ['main_color', 'average']
+ return {
+ "required": {
+ "image": ("IMAGE", ), #
+ "mode": (mode_list,), # 主色/平均色
+ },
+ "optional": {
+ }
+ }
+
+ RETURN_TYPES = ("STRING", "LIST")
+ RETURN_NAMES = ("RGB color in HEX", "HSV color in list")
+ FUNCTION = 'get_color_tone'
+ CATEGORY = '😺dzNodes/LayerUtility'
+
+ def get_color_tone(self, image, mode,):
+ if image.shape[0] > 0:
+ image = torch.unsqueeze(image[0], 0)
+ _canvas = tensor2pil(image).convert('RGB')
+ _canvas = gaussian_blur(_canvas, int((_canvas.width + _canvas.height) / 200))
+ if mode == 'main_color':
+ ret_color = get_image_color_tone(_canvas)
+ else:
+ ret_color = get_image_color_average(_canvas)
+ hsv_color = RGB_to_HSV(Hex_to_RGB(ret_color))
+ log(f"{self.NODE_NAME}: color is {ret_color}/{hsv_color}", message_type='finish')
+ return (ret_color, hsv_color)
+
+NODE_CLASS_MAPPINGS = {
+ "LayerUtility: GetColorTone": GetColorTone
+}
+
+NODE_DISPLAY_NAME_MAPPINGS = {
+ "LayerUtility: GetColorTone": "LayerUtility: GetColorTone(Advance)"
+}
\ No newline at end of file
diff --git a/py/get_color_tone_v2.py b/py/get_color_tone_v2.py
new file mode 100644
index 0000000..9587687
--- /dev/null
+++ b/py/get_color_tone_v2.py
@@ -0,0 +1,123 @@
+# layerstyle advance
+
+import torch
+from PIL import Image
+from .imagefunc import log, tensor2pil, pil2tensor, image2mask, gaussian_blur, get_image_color_tone, get_image_color_average
+from .imagefunc import RGB_to_HSV, Hex_to_RGB, pixel_spread, RMBG, expand_mask
+
+
+
+class GetColorToneV2:
+
+ def __init__(self):
+ self.NODE_NAME = 'GetColorToneV2'
+
+ @classmethod
+ def INPUT_TYPES(self):
+ remove_background_list = ['none','BiRefNet', 'RMBG 1.4',]
+ subject_list = ['mask','entire', 'background', 'subject']
+ mode_list = ['main_color', 'average']
+ return {
+ "required": {
+ "image": ("IMAGE", ), #
+ "mode": (mode_list,), # 主色/平均色
+ "color_of": (subject_list,),
+ "remove_bkgd_method": (remove_background_list,),
+ "invert_mask": ("BOOLEAN", {"default": False}), # 反转mask#
+ "mask_grow": ("INT", {"default": 0, "min": -999, "max": 999, "step": 1}),
+ },
+ "optional": {
+ "mask": ("MASK",), #
+ }
+ }
+
+ RETURN_TYPES = ("IMAGE", "STRING", "LIST", "MASK")
+ RETURN_NAMES = ("image", "color_in_hex", "HSV color in list", "mask",)
+ FUNCTION = 'get_color_tone_v2'
+ CATEGORY = '😺dzNodes/LayerUtility'
+
+ def get_color_tone_v2(self, image, mode, remove_bkgd_method, color_of, invert_mask, mask_grow,
+ mask=None
+ ):
+
+ _images = []
+ _masks = []
+ ret_images = []
+ ret_masks = []
+ need_rmbg = False
+ for i in image:
+ _images.append(torch.unsqueeze(i, 0))
+ m = tensor2pil(i)
+ if m.mode == 'RGBA':
+ _masks.append(1 - image2mask(m.split()[-1]))
+ else:
+ _masks.append(pil2tensor(Image.new("L", (m.width, m.height), color="white")))
+ if remove_bkgd_method != 'none':
+ need_rmbg = True
+
+ if mask is not None:
+ if mask.dim() == 2:
+ mask = torch.unsqueeze(mask, 0)
+ _masks = []
+ for m in mask:
+ _masks.append(torch.unsqueeze(m, 0))
+ need_rmbg = False
+
+ max_batch = max(len(_images), len(_masks))
+
+ if remove_bkgd_method == 'BiRefNet':
+ from .birefnet_legacy import BiRefNetRemoveBackground
+ birefnetrmbg = BiRefNetRemoveBackground()
+
+ for i in range(max_batch):
+ _image = _images[i] if i < len(_images) else _images[-1]
+ _image = tensor2pil(_image).convert("RGB")
+ if need_rmbg:
+ if remove_bkgd_method == 'BiRefNet':
+ _mask = birefnetrmbg.generate_mask(_image)
+ else:
+ _mask = RMBG(_image)
+ _mask = image2mask(_mask)
+ else:
+ _mask = _masks[i] if i < len(_masks) else _masks[-1]
+
+ if invert_mask:
+ _mask = 1 - _mask
+
+ if mask_grow != 0:
+ _mask = expand_mask(_mask, mask_grow, 0) # 扩张,模糊
+
+ if color_of == 'entire':
+ blured_image = gaussian_blur(_image, int((_image.width + _image.height) / 400))
+ else:
+ if color_of == 'background':
+ _mask = 1 - _mask
+ _mask = tensor2pil(_mask)
+ pixel_spread_image = pixel_spread(_image, _mask.convert('RGB'))
+ blured_image = gaussian_blur(pixel_spread_image, int((_image.width + _image.height) / 400))
+
+ ret_color = '#000000'
+ if mode == 'main_color' and color_of != 'mask':
+ ret_color = get_image_color_tone(blured_image)
+ elif mode == 'average' and color_of != 'mask':
+ ret_color = get_image_color_average(blured_image)
+ elif mode == 'main_color' and color_of == 'mask':
+ ret_color = get_image_color_tone(blured_image, mask=_mask)
+ elif mode == 'average' and color_of == 'mask':
+ ret_color = get_image_color_average(blured_image, mask=_mask)
+
+ ret_image = Image.new('RGB', size=_image.size, color=ret_color)
+ ret_images.append(pil2tensor(ret_image))
+ ret_masks.append(pil2tensor(_mask))
+ hsv_color = RGB_to_HSV(Hex_to_RGB(ret_color))
+ log(f"{self.NODE_NAME} Processed {len(ret_images)} image(s). Color is {ret_color}/{hsv_color}", message_type='finish')
+
+ return (torch.cat(ret_images, dim=0), ret_color, hsv_color, torch.cat(ret_masks, dim=0),)
+
+NODE_CLASS_MAPPINGS = {
+ "LayerUtility: GetColorToneV2": GetColorToneV2
+}
+
+NODE_DISPLAY_NAME_MAPPINGS = {
+ "LayerUtility: GetColorToneV2": "LayerUtility: GetColorTone V2(Advance)"
+}
\ No newline at end of file
diff --git a/py/human_parts_ultra.py b/py/human_parts_ultra.py
new file mode 100644
index 0000000..6bc7f0f
--- /dev/null
+++ b/py/human_parts_ultra.py
@@ -0,0 +1,199 @@
+# layerstyle advance
+
+import os
+from typing import Tuple
+import torch
+import numpy as np
+from PIL import Image, ImageEnhance
+import folder_paths
+from .imagefunc import pil2tensor, tensor2pil, image2mask, mask2image, log, RGB2RGBA, histogram_remap
+from .imagefunc import generate_VITMatte_trimap, generate_VITMatte, mask_edge_detail, guided_filter_alpha
+
+models_dir_path = os.path.join(folder_paths.models_dir, "onnx", "human-parts")
+model_url = "https://huggingface.co/Metal3d/deeplabv3p-resnet50-human/resolve/main/deeplabv3p-resnet50-human.onnx"
+model_name = os.path.basename(model_url)
+model_path = os.path.join(models_dir_path, "deeplabv3p-resnet50-human.onnx")
+
+
+class LS_HumanPartsUltra:
+ """
+ This node is used to get a mask of the human parts in the image.
+
+ The model used is DeepLabV3+ with a ResNet50 backbone trained
+ by Keras-io, converted to ONNX format.
+
+ """
+
+ def __init__(self):
+ self.NODE_NAME = 'HumanPartsUltra'
+
+ RETURN_TYPES = ("IMAGE", "MASK",)
+ RETURN_NAMES = ("image", "mask",)
+ FUNCTION = "human_parts_ultra"
+ CATEGORY = '😺dzNodes/LayerMask'
+
+ @classmethod
+ def INPUT_TYPES(cls):
+ method_list = ['VITMatte', 'VITMatte(local)', 'PyMatting', 'GuidedFilter', ]
+ device_list = ['cuda', 'cpu']
+ return {
+ "required": {
+ "image": ("IMAGE",),
+ "face": ("BOOLEAN", {"default": False, "label_on": "enabled(脸)", "label_off": "disabled(脸)"}),
+ "hair": ("BOOLEAN", {"default": False, "label_on": "enabled(头发)", "label_off": "disabled(头发)"}),
+ "glasses": ("BOOLEAN", {"default": False, "label_on": "enabled(眼镜)", "label_off": "disabled(眼镜)"}),
+ "top_clothes": ("BOOLEAN", {"default": False, "label_on": "enabled(上装)", "label_off": "disabled(上装)"}),
+ "bottom_clothes": ("BOOLEAN", {"default": False, "label_on": "enabled(下装)", "label_off": "disabled(下装)"}),
+ "torso_skin": ("BOOLEAN", {"default": False, "label_on": "enabled(躯干)", "label_off": "disabled(躯干)"}),
+ "left_arm": ("BOOLEAN", {"default": False, "label_on": "enabled(左臂)", "label_off": "disabled(左臂)"}),
+ "right_arm": ("BOOLEAN", {"default": False, "label_on": "enabled(右臂)", "label_off": "disabled(右臂)"}),
+ "left_leg": ("BOOLEAN", {"default": False, "label_on": "enabled(左腿)", "label_off": "disabled(左腿)"}),
+ "right_leg": ("BOOLEAN", {"default": False, "label_on": "enabled(右腿)", "label_off": "disabled(右腿)"}),
+ "left_foot": ("BOOLEAN", {"default": False, "label_on": "enabled(左脚)", "label_off": "disabled(左脚)"}),
+ "right_foot": ("BOOLEAN", {"default": False, "label_on": "enabled(右脚)", "label_off": "disabled(右脚)"}),
+ "detail_method": (method_list,),
+ "detail_erode": ("INT", {"default": 8, "min": 1, "max": 255, "step": 1}),
+ "detail_dilate": ("INT", {"default": 6, "min": 1, "max": 255, "step": 1}),
+ "black_point": (
+ "FLOAT", {"default": 0.01, "min": 0.01, "max": 0.98, "step": 0.01, "display": "slider"}),
+ "white_point": (
+ "FLOAT", {"default": 0.99, "min": 0.02, "max": 0.99, "step": 0.01, "display": "slider"}),
+ "process_detail": ("BOOLEAN", {"default": True}),
+ "device": (device_list,),
+ "max_megapixels": ("FLOAT", {"default": 2.0, "min": 1, "max": 999, "step": 0.1}),
+ }
+ }
+
+ def human_parts_ultra(self, image, face, hair, glasses, top_clothes, bottom_clothes,
+ torso_skin, left_arm, right_arm, left_leg, right_leg, left_foot, right_foot,
+ detail_method, detail_erode, detail_dilate, black_point, white_point,
+ process_detail, device, max_megapixels):
+ """
+ Return a Tensor with the mask of the human parts in the image.
+ """
+ import onnxruntime as ort
+
+ model = ort.InferenceSession(model_path, providers=['TensorrtExecutionProvider', 'CUDAExecutionProvider', 'CPUExecutionProvider'])
+ ret_images = []
+ ret_masks = []
+ for img in image:
+ orig_image = tensor2pil(img).convert('RGB')
+
+ human_parts_mask, _ = self.get_mask(orig_image, model=model, rotation=0, background=False,
+ face=face, hair=hair, glasses=glasses,
+ top_clothes=top_clothes, bottom_clothes=bottom_clothes,
+ torso_skin=torso_skin, left_arm=left_arm, right_arm=right_arm,
+ left_leg=left_leg, right_leg=right_leg,
+ left_foot=right_foot, right_foot=right_foot)
+ _mask = tensor2pil(human_parts_mask).convert('L')
+ brightness_image = ImageEnhance.Brightness(_mask)
+ _mask = brightness_image.enhance(factor=1.08)
+ _mask = image2mask(_mask)
+
+ if detail_method == 'VITMatte(local)':
+ local_files_only = True
+ else:
+ local_files_only = False
+ detail_range = detail_erode + detail_dilate
+ if process_detail:
+ if detail_method == 'GuidedFilter':
+ _mask = guided_filter_alpha(img.unsqueeze(0), _mask, detail_range // 6 + 1)
+ _mask = tensor2pil(histogram_remap(_mask, black_point, white_point))
+ elif detail_method == 'PyMatting':
+ _mask = tensor2pil(mask_edge_detail(img.unsqueeze(0), _mask, detail_range // 8 + 1, black_point, white_point))
+ else:
+ _trimap = generate_VITMatte_trimap(_mask, detail_erode, detail_dilate)
+ _mask = generate_VITMatte(orig_image, _trimap, local_files_only=local_files_only, device=device,
+ max_megapixels=max_megapixels)
+ _mask = tensor2pil(histogram_remap(pil2tensor(_mask), black_point, white_point))
+ else:
+ _mask = mask2image(_mask)
+
+ ret_image = RGB2RGBA(orig_image, _mask.convert('L'))
+ ret_images.append(pil2tensor(ret_image))
+ ret_masks.append(image2mask(_mask))
+
+ log(f"{self.NODE_NAME} Processed {len(ret_images)} image(s).", message_type='finish')
+ return (torch.cat(ret_images, dim=0), torch.cat(ret_masks, dim=0),)
+
+ def get_mask(self, pil_image:Image, model, rotation:float, **kwargs) -> tuple:
+ """
+ Return a Tensor with the mask of the human parts in the image.
+
+ The rotation parameter is not used for now. The idea is to propose rotation to help
+ the model to detect the human parts in the image if the character is not in a casual position.
+ Several tests have been done, but the model seems to fail to detect the human parts in these cases,
+ and the rotation does not help.
+ """
+
+ # classes used in the model
+ classes = {
+ "background": 0,
+ "hair": 2,
+ "glasses": 4,
+ "top_clothes": 5,
+ "bottom_clothes": 9,
+ "torso_skin": 10,
+ "face": 13,
+ "left_arm": 14,
+ "right_arm": 15,
+ "left_leg": 16,
+ "right_leg": 17,
+ "left_foot": 18,
+ "right_foot": 19,
+ }
+
+ original_size = pil_image.size # to resize the mask later
+ # resize to 512x512 as the model expects
+ pil_image = pil_image.resize((512, 512))
+ center = (256, 256)
+
+ if rotation != 0:
+ pil_image = pil_image.rotate(rotation, center=center)
+
+ # normalize the image
+ image_np = np.array(pil_image).astype(np.float32) / 127.5 - 1
+ image_np = np.expand_dims(image_np, axis=0)
+
+ # use the onnx model to get the mask
+ input_name = model.get_inputs()[0].name
+ output_name = model.get_outputs()[0].name
+ result = model.run([output_name], {input_name: image_np})
+ result = np.array(result[0]).argmax(axis=3).squeeze(0)
+
+ score: int = 0
+
+ mask = np.zeros_like(result)
+ for class_name, enabled in kwargs.items():
+ if enabled and class_name in classes:
+ class_index = classes[class_name]
+ detected = result == class_index
+ mask[detected] = 255
+ score += mask.sum()
+
+ # back to the original size
+ mask_image = Image.fromarray(mask.astype(np.uint8), mode="L")
+ if rotation != 0:
+ mask_image = mask_image.rotate(-rotation, center=center)
+
+ mask_image = mask_image.resize(original_size)
+
+ # and back to numpy...
+ mask = np.array(mask_image).astype(np.float32) / 255
+
+ # add 2 dimensions to match the expected output
+ mask = np.expand_dims(mask, axis=0)
+ mask = np.expand_dims(mask, axis=0)
+ # ensure to return a "binary mask_image"
+
+ del image_np, result # free up memory, maybe not necessary
+ return (torch.from_numpy(mask.astype(np.uint8)), score)
+
+
+NODE_CLASS_MAPPINGS = {
+ "LayerMask: HumanPartsUltra": LS_HumanPartsUltra
+}
+
+NODE_DISPLAY_NAME_MAPPINGS = {
+ "LayerMask: HumanPartsUltra": "LayerMask: Human Parts Ultra(Advance)"
+}
diff --git a/py/image_auto_crop.py b/py/image_auto_crop.py
new file mode 100644
index 0000000..8095319
--- /dev/null
+++ b/py/image_auto_crop.py
@@ -0,0 +1,165 @@
+# layerstyle advance
+
+from .imagefunc import *
+from .segment_anything_func import *
+
+
+
+class ImageAutoCrop:
+
+ def __init__(self):
+ self.NODE_NAME = 'ImageAutoCrop'
+
+ @classmethod
+ def INPUT_TYPES(self):
+ matting_method_list = ['RMBG 1.4', 'SegmentAnything']
+ detect_mode = ['min_bounding_rect', 'max_inscribed_rect', 'mask_area']
+ ratio_list = ['1:1', '3:2', '4:3', '16:9', '2:3', '3:4', '9:16', 'custom', 'detect_mask']
+ return {
+ "required": {
+ "image": ("IMAGE", ), #
+ "background_color": ("STRING", {"default": "#FFFFFF"}), # 背景颜色
+ "aspect_ratio": (ratio_list,),
+ "proportional_width": ("INT", {"default": 2, "min": 1, "max": 999, "step": 1}),
+ "proportional_height": ("INT", {"default": 1, "min": 1, "max": 999, "step": 1}),
+ "scale_to_longest_side": ("BOOLEAN", {"default": True}), # 是否按长边缩放
+ "longest_side": ("INT", {"default": 1024, "min": 4, "max": 999999, "step": 1}),
+ "detect": (detect_mode,),
+ "border_reserve": ("INT", {"default": 100, "min": -9999, "max": 9999, "step": 1}),
+ "ultra_detail_range": ("INT", {"default": 0, "min": 0, "max": 256, "step": 1}),
+ "matting_method": (matting_method_list,),
+ "sam_model": (list_sam_model(),),
+ "grounding_dino_model": (list_groundingdino_model(),),
+ "sam_threshold": ("FLOAT", {"default": 0.3, "min": 0, "max": 1.0, "step": 0.01}),
+ "sam_prompt": ("STRING", {"default": "subject"}),
+ },
+ "optional": {
+ }
+ }
+
+ RETURN_TYPES = ("IMAGE", "IMAGE", "MASK",)
+ RETURN_NAMES = ("cropped_image", "box_preview", "cropped_mask",)
+ FUNCTION = 'image_auto_crop'
+ CATEGORY = '😺dzNodes/LayerUtility'
+
+ def image_auto_crop(self, image, detect, border_reserve, aspect_ratio, proportional_width, proportional_height,
+ background_color, ultra_detail_range, scale_to_longest_side, longest_side,
+ matting_method, sam_model, grounding_dino_model, sam_threshold, sam_prompt
+ ):
+
+ ret_images = []
+ ret_box_previews = []
+ ret_masks = []
+ input_images = []
+ input_masks = []
+ crop_boxs = []
+
+ for l in image:
+ input_images.append(torch.unsqueeze(l, 0))
+ m = tensor2pil(l)
+ if m.mode == 'RGBA':
+ input_masks.append(m.split()[-1])
+
+ if len(input_masks) > 0 and len(input_masks) != len(input_images):
+ input_masks = []
+ log(f"Warning, {self.NODE_NAME} unable align alpha to image, drop it.", message_type='warning')
+
+ if aspect_ratio == 'custom':
+ ratio = proportional_width / proportional_height
+ elif aspect_ratio == 'detect_mask':
+ ratio = 0
+ else:
+ s = aspect_ratio.split(":")
+ ratio = int(s[0]) / int(s[1])
+ side_limit = longest_side if scale_to_longest_side else 0
+
+ for i in range(len(input_images)):
+ _image = tensor2pil(input_images[i]).convert('RGB')
+ if len(input_masks) > 0:
+ _mask = input_masks[i]
+ else:
+ if matting_method == 'SegmentAnything':
+ sam_model = load_sam_model(sam_model)
+ dino_model = load_groundingdino_model(grounding_dino_model)
+ item = _image.convert('RGBA')
+ boxes = groundingdino_predict(dino_model, item, sam_prompt, sam_threshold)
+ (_, _mask) = sam_segment(sam_model, item, boxes)
+ _mask = mask2image(_mask[0])
+ else:
+ _mask = RMBG(_image)
+ if ultra_detail_range:
+ _mask = tensor2pil(mask_edge_detail(input_images[i], pil2tensor(_mask), ultra_detail_range, 0.01, 0.99))
+ bluredmask = gaussian_blur(_mask, 20).convert('L')
+ x = 0
+ y = 0
+ width = 0
+ height = 0
+ x_offset = 0
+ y_offset = 0
+ if detect == "min_bounding_rect":
+ (x, y, width, height) = min_bounding_rect(bluredmask)
+ elif detect == "max_inscribed_rect":
+ (x, y, width, height) = max_inscribed_rect(bluredmask)
+ else:
+ (x, y, width, height) = mask_area(bluredmask)
+ canvas_width, canvas_height = _image.size
+ x1 = x - border_reserve
+ y1 = y - border_reserve
+ x2 = x + width + border_reserve
+ y2 = y + height + border_reserve
+ if x1 < 0:
+ canvas_width -= x1
+ x_offset = -x1
+ if y1 < 0:
+ canvas_height -= y1
+ y_offset = -y1
+ if x2 > _image.width:
+ canvas_width += x2 - _image.width
+ if y2 > _image.height:
+ canvas_height += y2 - _image.height
+ crop_box = (x1 + x_offset, y1 + y_offset, width + border_reserve*2, height + border_reserve*2)
+ crop_boxs.append(crop_box)
+ if len(crop_boxs) > 0: # 批量图强制使用同一尺寸
+ crop_box = crop_boxs[0]
+ if aspect_ratio == 'detect_mask':
+ ratio = crop_box[2] / crop_box[3]
+ target_width, target_height = calculate_side_by_ratio(crop_box[2], crop_box[3], ratio,
+ longest_side=side_limit)
+ _canvas = Image.new('RGB', size=(canvas_width, canvas_height), color=background_color)
+ _mask_canvas = Image.new('L', size=(canvas_width, canvas_height), color='black')
+ if ultra_detail_range:
+ _image = pixel_spread(_image, _mask)
+ _canvas.paste(_image, box=(x_offset, y_offset), mask=_mask.convert('L'))
+ _mask_canvas.paste(_mask, box=(x_offset, y_offset))
+ preview_image = Image.new('RGB', size=(canvas_width, canvas_height), color='gray')
+ preview_image.paste(_mask, box=(x_offset, y_offset))
+ preview_image = draw_rect(preview_image,
+ crop_box[0], crop_box[1], crop_box[2], crop_box[3],
+ line_color="#F00000", line_width=(canvas_width + canvas_height)//200)
+
+ ret_image = _canvas.crop((crop_box[0], crop_box[1], crop_box[0]+crop_box[2], crop_box[1]+crop_box[3]))
+ ret_image = fit_resize_image(ret_image, target_width, target_height,
+ fit='letterbox', resize_sampler=Image.LANCZOS,
+ background_color=background_color)
+ ret_mask = _mask_canvas.crop((crop_box[0], crop_box[1], crop_box[0]+crop_box[2], crop_box[1]+crop_box[3]))
+ ret_mask = fit_resize_image(ret_mask, target_width, target_height,
+ fit='letterbox', resize_sampler=Image.LANCZOS,
+ background_color="#000000")
+ ret_images.append(pil2tensor(ret_image))
+ ret_box_previews.append(pil2tensor(preview_image))
+ ret_masks.append(image2mask(ret_mask))
+
+ log(f"{self.NODE_NAME} Processed {len(ret_images)} image(s).", message_type='finish')
+ return (torch.cat(ret_images, dim=0),
+ torch.cat(ret_box_previews, dim=0),
+ torch.cat(ret_masks, dim=0),
+ )
+
+
+NODE_CLASS_MAPPINGS = {
+ "LayerUtility: ImageAutoCrop": ImageAutoCrop
+}
+
+NODE_DISPLAY_NAME_MAPPINGS = {
+ "LayerUtility: ImageAutoCrop": "LayerUtility: ImageAutoCrop(Advance)"
+}
\ No newline at end of file
diff --git a/py/image_auto_crop_v2.py b/py/image_auto_crop_v2.py
new file mode 100644
index 0000000..c0605a6
--- /dev/null
+++ b/py/image_auto_crop_v2.py
@@ -0,0 +1,253 @@
+# layerstyle advance
+
+from .imagefunc import *
+from .segment_anything_func import *
+
+
+SAM_MODEL = None
+DINO_MODEL = None
+previous_sam_model = ""
+previous_dino_model = ""
+
+class ImageAutoCropV2:
+
+ def __init__(self):
+ self.NODE_NAME = 'ImageAutoCropV2'
+
+ @classmethod
+ def INPUT_TYPES(self):
+ matting_method_list = ['RMBG 1.4', 'SegmentAnything']
+ detect_mode = ['min_bounding_rect', 'max_inscribed_rect', 'mask_area']
+ ratio_list = ['1:1', '3:2', '4:3', '16:9', '2:3', '3:4', '9:16', 'custom', 'detect_mask', 'original']
+ scale_to_side_list = ['None', 'longest', 'shortest', 'width', 'height']
+ return {
+ "required": {
+ "image": ("IMAGE", ), #
+ "fill_background": ("BOOLEAN", {"default": True}), # 是否填充背景
+ "background_color": ("STRING", {"default": "#FFFFFF"}), # 背景颜色
+ "aspect_ratio": (ratio_list,),
+ "proportional_width": ("INT", {"default": 1, "min": 1, "max": 999, "step": 1}),
+ "proportional_height": ("INT", {"default": 1, "min": 1, "max": 999, "step": 1}),
+ "scale_to_side": (scale_to_side_list,), # 是否按长边缩放
+ "scale_to_length": ("INT", {"default": 1024, "min": 4, "max": 999999, "step": 1}),
+ "detect": (detect_mode,),
+ "border_reserve": ("INT", {"default": 100, "min": -9999, "max": 9999, "step": 1}),
+ "ultra_detail_range": ("INT", {"default": 0, "min": 0, "max": 256, "step": 1}),
+ "matting_method": (matting_method_list,),
+ "sam_model": (list_sam_model(),),
+ "grounding_dino_model": (list_groundingdino_model(),),
+ "sam_threshold": ("FLOAT", {"default": 0.3, "min": 0, "max": 1.0, "step": 0.01}),
+ "sam_prompt": ("STRING", {"default": "subject"}),
+ },
+ "optional": {
+ "mask": ("MASK",), #
+ }
+ }
+
+ RETURN_TYPES = ("IMAGE", "IMAGE", "MASK",)
+ RETURN_NAMES = ("cropped_image", "box_preview", "cropped_mask",)
+ FUNCTION = 'image_auto_crop_v2'
+ CATEGORY = '😺dzNodes/LayerUtility'
+
+ def image_auto_crop_v2(self, image, fill_background, background_color, aspect_ratio,
+ proportional_width, proportional_height,
+ scale_to_side, scale_to_length, detect, border_reserve,
+ ultra_detail_range, matting_method,
+ sam_model, grounding_dino_model, sam_threshold, sam_prompt,
+ mask=None,
+ ):
+
+ ret_images = []
+ ret_box_previews = []
+ ret_masks = []
+ input_images = []
+ input_masks = []
+ crop_boxs = []
+
+ global SAM_MODEL
+ global DINO_MODEL
+ global previous_sam_model
+ global previous_dino_model
+
+ for l in image:
+ input_images.append(torch.unsqueeze(l, 0))
+ m = tensor2pil(l)
+ if m.mode == 'RGBA':
+ input_masks.append(m.split()[-1])
+ if mask is not None:
+ if mask.dim() == 2:
+ mask = torch.unsqueeze(mask, 0)
+ input_masks = []
+ for m in mask:
+ input_masks.append(tensor2pil(torch.unsqueeze(m, 0)).convert('L'))
+
+ if len(input_masks) > 0 and len(input_masks) != len(input_images):
+ input_masks = []
+ log(f"Warning, {self.NODE_NAME} unable align alpha to image, drop it.", message_type='warning')
+
+ fit = 'letterbox'
+ if aspect_ratio == 'custom':
+ ratio = proportional_width / proportional_height
+ elif aspect_ratio == 'original':
+ _image = tensor2pil(input_images[0])
+ ratio = _image.width / _image.height
+ elif aspect_ratio == 'detect_mask':
+ ratio = 0
+ fit = 'fill'
+ else:
+ s = aspect_ratio.split(":")
+ ratio = int(s[0]) / int(s[1])
+
+ for i in range(len(input_images)):
+ _image = tensor2pil(input_images[i]).convert('RGB')
+
+ if len(input_masks) > 0:
+ _mask = input_masks[i]
+ else:
+ if matting_method == 'SegmentAnything':
+ if previous_sam_model != sam_model:
+ SAM_MODEL = load_sam_model(sam_model)
+ previous_sam_model = sam_model
+ if previous_dino_model != grounding_dino_model:
+ DINO_MODEL = load_groundingdino_model(grounding_dino_model)
+ previous_dino_model = grounding_dino_model
+ item = _image.convert('RGBA')
+ boxes = groundingdino_predict(DINO_MODEL, item, sam_prompt, sam_threshold)
+ (_, _mask) = sam_segment(SAM_MODEL, item, boxes)
+ _mask = mask2image(_mask[0])
+ else:
+ _mask = RMBG(_image)
+ if ultra_detail_range:
+ _mask = tensor2pil(mask_edge_detail(input_images[i], pil2tensor(_mask), ultra_detail_range, 0.01, 0.99))
+ bluredmask = gaussian_blur(_mask, 20).convert('L')
+ x = 0
+ y = 0
+ width = 0
+ height = 0
+ x_offset = 0
+ y_offset = 0
+ if detect == "min_bounding_rect":
+ (x, y, width, height) = min_bounding_rect(bluredmask)
+ elif detect == "max_inscribed_rect":
+ (x, y, width, height) = max_inscribed_rect(bluredmask)
+ else:
+ (x, y, width, height) = mask_area(bluredmask)
+
+ canvas_width, canvas_height = _image.size
+
+ x1 = x - border_reserve
+ y1 = y - border_reserve
+ x2 = x + width + border_reserve
+ y2 = y + height + border_reserve
+
+ if x1 < 0:
+ if fill_background:
+ canvas_width -= x1
+ x_offset = -x1
+ else:
+ x1 = 0
+ if y1 < 0:
+ if fill_background:
+ canvas_height -= y1
+ y_offset = -y1
+ else:
+ y1 = 0
+ if x2 > _image.width:
+ if fill_background:
+ canvas_width += x2 - _image.width
+ else:
+ x2 = _image.width
+ if y2 > _image.height:
+ if fill_background:
+ canvas_height += y2 - _image.height
+ else:
+ y2 = _image.height
+
+ if fill_background:
+ crop_box = (x1 + x_offset, y1 + y_offset, width + border_reserve*2, height + border_reserve*2)
+ else:
+ crop_box = (x1, y1, x2 - x1, y2 - y1)
+ crop_boxs.append(crop_box)
+ if len(crop_boxs) > 0: # 批量图强制使用同一尺寸
+ crop_box = crop_boxs[0]
+
+ orig_width = crop_box[2]
+ orig_height = crop_box[3]
+ if aspect_ratio == 'detect_mask':
+ ratio = orig_width / orig_height
+
+ # calculate target width and height
+ if orig_width > orig_height:
+ if scale_to_side == 'longest':
+ target_width = scale_to_length
+ target_height = int(target_width / ratio)
+ elif scale_to_side == 'shortest':
+ target_height = scale_to_length
+ target_width = int(target_height * ratio)
+ elif scale_to_side == 'width':
+ target_width = scale_to_length
+ target_height = int(target_width / ratio)
+ elif scale_to_side == 'height':
+ target_height = scale_to_length
+ target_width = int(target_height * ratio)
+ else:
+ target_width = orig_width
+ target_height = int(target_width / ratio)
+ else:
+ if scale_to_side == 'longest':
+ target_height = scale_to_length
+ target_width = int(target_height * ratio)
+ elif scale_to_side == 'shortest':
+ target_width = scale_to_length
+ target_height = int(target_width / ratio)
+ elif scale_to_side == 'width':
+ target_width = scale_to_length
+ target_height = int(target_width / ratio)
+ elif scale_to_side == 'height':
+ target_height = scale_to_length
+ target_width = int(target_height * ratio)
+ else:
+ target_height = orig_height
+ target_width = int(target_height * ratio)
+
+ _canvas = Image.new('RGB', size=(canvas_width, canvas_height), color=background_color)
+ _mask_canvas = Image.new('L', size=(canvas_width, canvas_height), color='black')
+ if ultra_detail_range:
+ _image = pixel_spread(_image, _mask)
+ if fill_background:
+ _canvas.paste(_image, box=(x_offset, y_offset), mask=_mask.convert('L'))
+ else:
+ _canvas.paste(_image, box=(x_offset, y_offset))
+ _mask_canvas.paste(_mask, box=(x_offset, y_offset))
+ preview_image = Image.new('RGB', size=(canvas_width, canvas_height), color='gray')
+ preview_image.paste(_mask, box=(x_offset, y_offset))
+ preview_image = draw_rect(preview_image,
+ crop_box[0], crop_box[1], crop_box[2], crop_box[3],
+ line_color="#F00000", line_width=(canvas_width + canvas_height)//200)
+
+ ret_image = _canvas.crop((crop_box[0], crop_box[1], crop_box[0]+crop_box[2], crop_box[1]+crop_box[3]))
+ ret_image = fit_resize_image(ret_image, target_width, target_height,
+ fit=fit, resize_sampler=Image.LANCZOS,
+ background_color=background_color)
+ ret_mask = _mask_canvas.crop((crop_box[0], crop_box[1], crop_box[0]+crop_box[2], crop_box[1]+crop_box[3]))
+ ret_mask = fit_resize_image(ret_mask, target_width, target_height,
+ fit=fit, resize_sampler=Image.LANCZOS,
+ background_color="#000000")
+ ret_images.append(pil2tensor(ret_image))
+ ret_box_previews.append(pil2tensor(preview_image))
+ ret_masks.append(image2mask(ret_mask))
+
+ log(f"{self.NODE_NAME} Processed {len(ret_images)} image(s).", message_type='finish')
+ return (torch.cat(ret_images, dim=0),
+ torch.cat(ret_box_previews, dim=0),
+ torch.cat(ret_masks, dim=0),
+ )
+
+
+NODE_CLASS_MAPPINGS = {
+ "LayerUtility: ImageAutoCrop V2": ImageAutoCropV2
+}
+
+NODE_DISPLAY_NAME_MAPPINGS = {
+ "LayerUtility: ImageAutoCrop V2": "LayerUtility: ImageAutoCrop V2(Advance)"
+}
\ No newline at end of file
diff --git a/py/image_auto_crop_v3.py b/py/image_auto_crop_v3.py
new file mode 100644
index 0000000..511737b
--- /dev/null
+++ b/py/image_auto_crop_v3.py
@@ -0,0 +1,193 @@
+# layerstyle advance
+
+import torch
+import numpy as np
+import math
+from PIL import Image
+from .imagefunc import log, tensor2pil, pil2tensor, num_round_up_to_multiple, draw_rect, gaussian_blur, mask_area
+
+
+
+class ImageAutoCropV3:
+
+ def __init__(self):
+ self.NODE_NAME = 'ImageAutoCropV3'
+
+ @classmethod
+ def INPUT_TYPES(self):
+ ratio_list = ['1:1', '3:2', '4:3', '16:9', '2:3', '3:4', '9:16', 'custom', 'original']
+ scale_to_side_list = ['None', 'longest', 'shortest', 'width', 'height', 'total_pixel(kilo pixel)']
+ multiple_list = ['8', '16', '32', '64', '128', '256', '512', 'None']
+ method_mode = ['lanczos', 'bicubic', 'hamming', 'bilinear', 'box', 'nearest']
+ return {
+ "required": {
+ "image": ("IMAGE", ),
+ "aspect_ratio": (ratio_list,),
+ "proportional_width": ("INT", {"default": 1, "min": 1, "max": 99999999, "step": 1}),
+ "proportional_height": ("INT", {"default": 1, "min": 1, "max": 99999999, "step": 1}),
+ "method": (method_mode,),
+ "scale_to_side": (scale_to_side_list,),
+ "scale_to_length": ("INT", {"default": 1024, "min": 4, "max": 999999, "step": 1}),
+ "round_to_multiple": (multiple_list,),
+ },
+ "optional": {
+ "mask": ("MASK",),
+ }
+ }
+
+ RETURN_TYPES = ("IMAGE", "IMAGE",)
+ RETURN_NAMES = ("cropped_image", "box_preview",)
+ FUNCTION = 'image_auto_crop_v3'
+ CATEGORY = '😺dzNodes/LayerUtility'
+
+ def image_auto_crop_v3(self, image, aspect_ratio,
+ proportional_width, proportional_height, method,
+ scale_to_side, scale_to_length, round_to_multiple,
+ mask=None,
+ ):
+
+ ret_images = []
+ ret_box_previews = []
+ ret_masks = []
+ input_images = []
+ input_masks = []
+ crop_boxs = []
+
+ for l in image:
+ input_images.append(torch.unsqueeze(l, 0))
+ m = tensor2pil(l)
+ if m.mode == 'RGBA':
+ input_masks.append(m.split()[-1])
+ if mask is not None:
+ if mask.dim() == 2:
+ mask = torch.unsqueeze(mask, 0)
+ input_masks = []
+ for m in mask:
+ input_masks.append(tensor2pil(torch.unsqueeze(m, 0)).convert('L'))
+
+ if len(input_masks) > 0 and len(input_masks) != len(input_images):
+ input_masks = []
+ log(f"Warning, {self.NODE_NAME} unable align alpha to image, drop it.", message_type='warning')
+
+ fit = 'crop'
+ _image = tensor2pil(input_images[0])
+ (orig_width, orig_height) = _image.size
+ if aspect_ratio == 'custom':
+ ratio = proportional_width / proportional_height
+ elif aspect_ratio == 'original':
+ ratio = orig_width / orig_height
+ else:
+ s = aspect_ratio.split(":")
+ ratio = int(s[0]) / int(s[1])
+
+ resize_sampler = Image.LANCZOS
+ if method == "bicubic":
+ resize_sampler = Image.BICUBIC
+ elif method == "hamming":
+ resize_sampler = Image.HAMMING
+ elif method == "bilinear":
+ resize_sampler = Image.BILINEAR
+ elif method == "box":
+ resize_sampler = Image.BOX
+ elif method == "nearest":
+ resize_sampler = Image.NEAREST
+
+ # calculate target width and height
+ if ratio > 1:
+ if scale_to_side == 'longest':
+ target_width = scale_to_length
+ target_height = int(target_width / ratio)
+ elif scale_to_side == 'shortest':
+ target_height = scale_to_length
+ target_width = int(target_height * ratio)
+ elif scale_to_side == 'width':
+ target_width = scale_to_length
+ target_height = int(target_width / ratio)
+ elif scale_to_side == 'height':
+ target_height = scale_to_length
+ target_width = int(target_height * ratio)
+ elif scale_to_side == 'total_pixel(kilo pixel)':
+ target_width = math.sqrt(ratio * scale_to_length * 1000)
+ target_height = target_width / ratio
+ target_width = int(target_width)
+ target_height = int(target_height)
+ else:
+ target_width = orig_width
+ target_height = int(target_width / ratio)
+ else:
+ if scale_to_side == 'longest':
+ target_height = scale_to_length
+ target_width = int(target_height * ratio)
+ elif scale_to_side == 'shortest':
+ target_width = scale_to_length
+ target_height = int(target_width / ratio)
+ elif scale_to_side == 'width':
+ target_width = scale_to_length
+ target_height = int(target_width / ratio)
+ elif scale_to_side == 'height':
+ target_height = scale_to_length
+ target_width = int(target_height * ratio)
+ elif scale_to_side == 'total_pixel(kilo pixel)':
+ target_width = math.sqrt(ratio * scale_to_length * 1000)
+ target_height = target_width / ratio
+ target_width = int(target_width)
+ target_height = int(target_height)
+ else:
+ target_height = orig_height
+ target_width = int(target_height * ratio)
+
+ if round_to_multiple != 'None':
+ multiple = int(round_to_multiple)
+ target_width = num_round_up_to_multiple(target_width, multiple)
+ target_height = num_round_up_to_multiple(target_height, multiple)
+
+ for i in range(len(input_images)):
+ _image = tensor2pil(input_images[i]).convert('RGB')
+
+ if len(input_masks) > 0:
+ _mask = input_masks[i]
+ else:
+ _mask = Image.new('L', _image.size, color='black')
+
+ bluredmask = gaussian_blur(_mask, 20).convert('L')
+ (mask_x, mask_y, mask_w, mask_h) = mask_area(bluredmask)
+ orig_ratio = _image.width / _image.height
+ target_ratio = target_width / target_height
+ # crop image to target ratio
+ if orig_ratio > target_ratio: # crop LiftRight side
+ crop_w = int(_image.height * target_ratio)
+ crop_h = _image.height
+ else: # crop TopBottom side
+ crop_w = _image.width
+ crop_h = int(_image.width / target_ratio)
+ crop_x = mask_w // 2 + mask_x - crop_w // 2
+ if crop_x < 0:
+ crop_x = 0
+ if crop_x + crop_w > _image.width:
+ crop_x = _image.width - crop_w
+ crop_y = mask_h // 2 + mask_y - crop_h // 2
+ if crop_y < 0:
+ crop_y = 0
+ if crop_y + crop_h > _image.height:
+ crop_y = _image.height - crop_h
+ crop_image = _image.crop((crop_x, crop_y, crop_x + crop_w, crop_y + crop_h))
+ line_width = (_image.width + _image.height) // 200
+ preview_image = draw_rect(_image, crop_x, crop_y,
+ crop_w, crop_h,
+ line_color="#F00000", line_width=line_width)
+ ret_image = crop_image.resize((target_width, target_height), resize_sampler)
+ ret_images.append(pil2tensor(ret_image))
+ ret_box_previews.append(pil2tensor(preview_image))
+
+ log(f"{self.NODE_NAME} Processed {len(ret_images)} image(s).", message_type='finish')
+ return (torch.cat(ret_images, dim=0),
+ torch.cat(ret_box_previews, dim=0),
+ )
+
+NODE_CLASS_MAPPINGS = {
+ "LayerUtility: ImageAutoCrop V3": ImageAutoCropV3
+}
+
+NODE_DISPLAY_NAME_MAPPINGS = {
+ "LayerUtility: ImageAutoCrop V3": "LayerUtility: ImageAutoCrop V3(Advance)"
+}
\ No newline at end of file
diff --git a/py/image_reward_filter.py b/py/image_reward_filter.py
new file mode 100644
index 0000000..dbf1daa
--- /dev/null
+++ b/py/image_reward_filter.py
@@ -0,0 +1,77 @@
+# layerstyle advance
+
+from .imagefunc import *
+
+
+class ImageRewardFilter:
+
+ def __init__(self):
+ self.NODE_NAME = 'ImageRewardFilter'
+
+ @classmethod
+ def INPUT_TYPES(self):
+ return {
+ "required": {
+ "images": ("IMAGE", ),
+ "prompt": ("STRING", {"multiline": False}),
+ "output_num": ("INT", {"default": 3, "min": 1, "max": 999999, "step": 1}),
+ },
+ "optional": {
+ }
+ }
+
+ RETURN_TYPES = ("IMAGE", "IMAGE",)
+ RETURN_NAMES = ("images", 'obsolete_images',)
+ FUNCTION = 'image_reward_filter'
+ CATEGORY = '😺dzNodes/LayerUtility'
+
+ def image_reward_filter(self, images, prompt, output_num,):
+ log(f"len(images)= {len(images)}, output_num={output_num}")
+ if output_num > len(images):
+ log(f"Error: {self.NODE_NAME} skipped, because 'output_num' is greater then input images.", message_type='error')
+ return (images,)
+
+ scores = []
+ ret_images = []
+ obsolete_images = []
+
+ if not torch.cuda.is_available() :
+ device = "cpu"
+ else:
+ device = "cuda"
+
+ import ImageReward as RM
+ reward_model = RM.load("ImageReward-v1.0")
+ reward_model = reward_model.to(device=device)
+
+ with torch.no_grad():
+ for i in range(len(images)):
+ score = reward_model.score(prompt, tensor2pil(images[i]))
+ scores.append(
+ {
+ "score":score,
+ "image_index":i
+ }
+ )
+ scores = sorted(scores, key=lambda s: s['score'], reverse=True)
+
+ for i in range(len(images)):
+ if i < output_num:
+ log(f"{self.NODE_NAME} append image #{i}: {scores[i]['image_index']}, score = {scores[i]['score']}.")
+ ret_images.append(images[scores[i]['image_index']])
+ else:
+ log(f"{self.NODE_NAME} obsolete image #{i}: {scores[i]['image_index']}, score = {scores[i]['score']}.")
+ obsolete_images.append(images[scores[i]['image_index']])
+
+ log(f"{self.NODE_NAME} Processed {len(ret_images)} image(s).", message_type='finish')
+
+ return (ret_images, obsolete_images,)
+
+
+NODE_CLASS_MAPPINGS = {
+ "LayerUtility: ImageRewardFilter": ImageRewardFilter
+}
+
+NODE_DISPLAY_NAME_MAPPINGS = {
+ "LayerUtility: ImageRewardFilter": "LayerUtility: ImageRewardFilter(Obsolete)"
+}
\ No newline at end of file
diff --git a/py/imagefunc.py b/py/imagefunc.py
new file mode 100644
index 0000000..a20748f
--- /dev/null
+++ b/py/imagefunc.py
@@ -0,0 +1,2485 @@
+"""Image process functions for ComfyUI nodes
+by chflame https://github.com/chflame163
+
+@author: chflame
+@title: LayerStyle
+@nickname: LayerStyle
+@description: A set of nodes for ComfyUI that can composite layer and mask to achieve Photoshop like functionality.
+"""
+
+import os
+import sys
+sys.path.append(os.path.dirname(os.path.abspath(__file__)))
+import pickle
+import copy
+import re
+import json
+import math
+import glob
+import numpy as np
+import torch
+import scipy.ndimage
+import cv2
+import random
+import time
+from pathlib import Path
+from tqdm import tqdm
+from functools import lru_cache
+from typing import Union, List
+from PIL import Image, ImageFilter, ImageChops, ImageDraw, ImageOps, ImageEnhance, ImageFont
+from skimage import img_as_float, img_as_ubyte
+import torchvision.transforms.functional as TF
+import torch.nn.functional as F
+from transformers import AutoModel, AutoProcessor, StoppingCriteria, StoppingCriteriaList, AutoModelForCausalLM, AutoTokenizer
+from colorsys import rgb_to_hsv
+import folder_paths
+import comfy.model_management
+from .blendmodes import *
+
+def log(message:str, message_type:str='info'):
+ name = 'LayerStyle'
+
+ if message_type == 'error':
+ message = '\033[1;41m' + message + '\033[m'
+ elif message_type == 'warning':
+ message = '\033[1;31m' + message + '\033[m'
+ elif message_type == 'finish':
+ message = '\033[1;32m' + message + '\033[m'
+ else:
+ message = '\033[1;33m' + message + '\033[m'
+ print(f"# 😺dzNodes: {name} -> {message}")
+
+try:
+ from cv2.ximgproc import guidedFilter
+except ImportError as e:
+ # print(e)
+ log(f"Cannot import name 'guidedFilter' from 'cv2.ximgproc'"
+ f"\nA few nodes cannot works properly, while most nodes are not affected. Please REINSTALL package 'opencv-contrib-python'."
+ f"\nFor detail refer to \033[4mhttps://github.com/chflame163/ComfyUI_LayerStyle/issues/5\033[0m")
+
+
+
+'''warpper'''
+
+# create a wrapper function that can apply a function to multiple images in a batch while passing all other arguments to the function
+def apply_to_batch(func):
+ def wrapper(self, image, *args, **kwargs):
+ images = []
+ for img in image:
+ images.append(func(self, img, *args, **kwargs))
+ batch_tensor = torch.cat(images, dim=0)
+ return (batch_tensor,)
+ return wrapper
+
+
+'''pickle'''
+
+
+def read_image(filename:str) -> Image:
+ return Image.open(filename)
+
+def pickle_to_file(obj:object, file_path:str):
+ with open(file_path, 'wb') as f:
+ pickle.dump(obj, f)
+
+def load_pickle(file_name:str) -> object:
+ with open(file_name, 'rb') as f:
+ obj = pickle.load(f)
+ return obj
+
+def load_light_leak_images() -> list:
+ file = os.path.join(folder_paths.models_dir, "layerstyle", "light_leak.pkl")
+ return load_pickle(file)
+
+'''Converter'''
+
+def cv22ski(cv2_image:np.ndarray) -> np.array:
+ return img_as_float(cv2_image)
+
+def ski2cv2(ski:np.array) -> np.ndarray:
+ return img_as_ubyte(ski)
+
+def cv22pil(cv2_img:np.ndarray) -> Image:
+ cv2_img = cv2.cvtColor(cv2_img, cv2.COLOR_BGR2RGB)
+ return Image.fromarray(cv2_img)
+
+def pil2cv2(pil_img:Image) -> np.array:
+ np_img_array = np.asarray(pil_img)
+ return cv2.cvtColor(np_img_array, cv2.COLOR_RGB2BGR)
+
+def pil2tensor(image:Image) -> torch.Tensor:
+ return torch.from_numpy(np.array(image).astype(np.float32) / 255.0).unsqueeze(0)
+
+def np2pil(np_image:np.ndarray) -> Image:
+ return Image.fromarray(np_image)
+
+def pil2np(pil_image:Image) -> np.array:
+ return np.ndarray(pil_image)
+
+def np2tensor(img_np: Union[np.ndarray, List[np.ndarray]]) -> torch.Tensor:
+ if isinstance(img_np, list):
+ return torch.cat([np2tensor(img) for img in img_np], dim=0)
+ return torch.from_numpy(img_np.astype(np.float32) / 255.0).unsqueeze(0)
+
+def tensor2np(tensor: torch.Tensor) -> List[np.ndarray]:
+ if len(tensor.shape) == 3: # Single image
+ return np.clip(255.0 * tensor.cpu().numpy(), 0, 255).astype(np.uint8)
+ else: # Batch of images
+ return [np.clip(255.0 * t.cpu().numpy(), 0, 255).astype(np.uint8) for t in tensor]
+
+def tensor2pil(t_image: torch.Tensor) -> Image:
+ return Image.fromarray(np.clip(255.0 * t_image.cpu().numpy().squeeze(), 0, 255).astype(np.uint8))
+
+def tensor2cv2(image:torch.Tensor) -> np.array:
+ if image.dim() == 4:
+ image = image.squeeze()
+ npimage = image.numpy()
+ cv2image = np.uint8(npimage * 255 / npimage.max())
+ return cv2.cvtColor(cv2image, cv2.COLOR_RGB2BGR)
+
+def image2mask(image:Image) -> torch.Tensor:
+ if image.mode == 'L':
+ return torch.tensor([pil2tensor(image)[0, :, :].tolist()])
+ else:
+ image = image.convert('RGB').split()[0]
+ return torch.tensor([pil2tensor(image)[0, :, :].tolist()])
+
+def mask2image(mask:torch.Tensor) -> Image:
+ masks = tensor2np(mask)
+ for m in masks:
+ _mask = Image.fromarray(m).convert("L")
+ _image = Image.new("RGBA", _mask.size, color='white')
+ _image = Image.composite(
+ _image, Image.new("RGBA", _mask.size, color='black'), _mask)
+ return _image
+
+'''Image Functions'''
+
+# 颜色加深
+def blend_color_burn(background_image:Image, layer_image:Image) -> Image:
+ img_1 = cv22ski(pil2cv2(background_image))
+ img_2 = cv22ski(pil2cv2(layer_image))
+ img = 1 - (1 - img_2) / (img_1 + 0.001)
+ mask_1 = img < 0
+ mask_2 = img > 1
+ img = img * (1 - mask_1)
+ img = img * (1 - mask_2) + mask_2
+ return cv22pil(ski2cv2(img))
+
+# 颜色减淡
+def blend_color_dodge(background_image:Image, layer_image:Image) -> Image:
+ img_1 = cv22ski(pil2cv2(background_image))
+ img_2 = cv22ski(pil2cv2(layer_image))
+ img = img_2 / (1.0 - img_1 + 0.001)
+ mask_2 = img > 1
+ img = img * (1 - mask_2) + mask_2
+ return cv22pil(ski2cv2(img))
+
+# 线性加深
+def blend_linear_burn(background_image:Image, layer_image:Image) -> Image:
+ img_1 = cv22ski(pil2cv2(background_image))
+ img_2 = cv22ski(pil2cv2(layer_image))
+ img = img_1 + img_2 - 1
+ mask_1 = img < 0
+ img = img * (1 - mask_1)
+ return cv22pil(ski2cv2(img))
+
+# 线性减淡
+def blend_linear_dodge(background_image:Image, layer_image:Image) -> Image:
+ img_1 = cv22ski(pil2cv2(background_image))
+ img_2 = cv22ski(pil2cv2(layer_image))
+ img = img_1 + img_2
+ mask_2 = img > 1
+ img = img * (1 - mask_2) + mask_2
+ return cv22pil(ski2cv2(img))
+
+# 变亮
+def blend_lighten(background_image:Image, layer_image:Image) -> Image:
+ img_1 = cv22ski(pil2cv2(background_image))
+ img_2 = cv22ski(pil2cv2(layer_image))
+ img = img_1 - img_2
+ mask = img > 0
+ img = img_1 * mask + img_2 * (1 - mask)
+ return cv22pil(ski2cv2(img))
+
+# 变暗
+def blend_dark(background_image:Image, layer_image:Image) -> Image:
+ img_1 = cv22ski(pil2cv2(background_image))
+ img_2 = cv22ski(pil2cv2(layer_image))
+ img = img_1 - img_2
+ mask = img < 0
+ img = img_1 * mask + img_2 * (1 - mask)
+ return cv22pil(ski2cv2(img))
+
+# 滤色
+def blend_screen(background_image:Image, layer_image:Image) -> Image:
+ img_1 = cv22ski(pil2cv2(background_image))
+ img_2 = cv22ski(pil2cv2(layer_image))
+ img = 1 - (1 - img_1) * (1 - img_2)
+ return cv22pil(ski2cv2(img))
+
+# 叠加
+def blend_overlay(background_image:Image, layer_image:Image) -> Image:
+ img_1 = cv22ski(pil2cv2(background_image))
+ img_2 = cv22ski(pil2cv2(layer_image))
+ mask = img_2 < 0.5
+ img = 2 * img_1 * img_2 * mask + (1 - mask) * (1 - 2 * (1 - img_1) * (1 - img_2))
+ return cv22pil(ski2cv2(img))
+
+# 柔光
+def blend_soft_light(background_image:Image, layer_image:Image) -> Image:
+ img_1 = cv22ski(pil2cv2(background_image))
+ img_2 = cv22ski(pil2cv2(layer_image))
+ mask = img_1 < 0.5
+ T1 = (2 * img_1 - 1) * (img_2 - img_2 * img_2) + img_2
+ T2 = (2 * img_1 - 1) * (np.sqrt(img_2) - img_2) + img_2
+ img = T1 * mask + T2 * (1 - mask)
+ return cv22pil(ski2cv2(img))
+
+# 强光
+def blend_hard_light(background_image:Image, layer_image:Image) -> Image:
+ img_1 = cv22ski(pil2cv2(background_image))
+ img_2 = cv22ski(pil2cv2(layer_image))
+ mask = img_1 < 0.5
+ T1 = 2 * img_1 * img_2
+ T2 = 1 - 2 * (1 - img_1) * (1 - img_2)
+ img = T1 * mask + T2 * (1 - mask)
+ return cv22pil(ski2cv2(img))
+
+# 亮光
+def blend_vivid_light(background_image:Image, layer_image:Image) -> Image:
+ img_1 = cv22ski(pil2cv2(background_image))
+ img_2 = cv22ski(pil2cv2(layer_image))
+ mask = img_1 < 0.5
+ T1 = 1 - (1 - img_2) / (2 * img_1 + 0.001)
+ T2 = img_2 / (2 * (1 - img_1) + 0.001)
+ mask_1 = T1 < 0
+ mask_2 = T2 > 1
+ T1 = T1 * (1 - mask_1)
+ T2 = T2 * (1 - mask_2) + mask_2
+ img = T1 * mask + T2 * (1 - mask)
+ return cv22pil(ski2cv2(img))
+
+# 点光
+def blend_pin_light(background_image:Image, layer_image:Image) -> Image:
+ img_1 = cv22ski(pil2cv2(background_image))
+ img_2 = cv22ski(pil2cv2(layer_image))
+ mask_1 = img_2 < (img_1 * 2 - 1)
+ mask_2 = img_2 > 2 * img_1
+ T1 = 2 * img_1 - 1
+ T2 = img_2
+ T3 = 2 * img_1
+ img = T1 * mask_1 + T2 * (1 - mask_1) * (1 - mask_2) + T3 * mask_2
+ return cv22pil(ski2cv2(img))
+
+# 线性光
+def blend_linear_light(background_image:Image, layer_image:Image) -> Image:
+ img_1 = cv22ski(pil2cv2(background_image))
+ img_2 = cv22ski(pil2cv2(layer_image))
+ img = img_2 + img_1 * 2 - 1
+ mask_1 = img < 0
+ mask_2 = img > 1
+ img = img * (1 - mask_1)
+ img = img * (1 - mask_2) + mask_2
+ return cv22pil(ski2cv2(img))
+
+def blend_hard_mix(background_image:Image, layer_image:Image) -> Image:
+ img_1 = cv22ski(pil2cv2(background_image))
+ img_2 = cv22ski(pil2cv2(layer_image))
+ img = img_1 + img_2
+ mask = img_1 + img_2 > 1
+ img = img * (1 - mask) + mask
+ img = img * mask
+ return cv22pil(ski2cv2(img))
+
+def shift_image(image:Image, distance_x:int, distance_y:int, background_color:str='#000000', cyclic:bool=False) -> Image:
+ width = image.width
+ height = image.height
+ ret_image = Image.new('RGB', size=(width, height), color=background_color)
+ for x in range(width):
+ for y in range(height):
+ if cyclic:
+ orig_x = x + distance_x
+ if orig_x > width-1 or orig_x < 0:
+ orig_x = abs(orig_x % width)
+ orig_y = y + distance_y
+ if orig_y > height-1 or orig_y < 0:
+ orig_y = abs(orig_y % height)
+
+ pixel = image.getpixel((orig_x, orig_y))
+ ret_image.putpixel((x, y), pixel)
+ else:
+ if x > -distance_x and y > -distance_y: # 防止回转
+ if x + distance_x < width and y + distance_y < height: # 防止越界
+ pixel = image.getpixel((x + distance_x, y + distance_y))
+ ret_image.putpixel((x, y), pixel)
+ return ret_image
+
+def chop_image(background_image:Image, layer_image:Image, blend_mode:str, opacity:int) -> Image:
+ ret_image = background_image
+ if blend_mode == 'normal':
+ ret_image = copy.deepcopy(layer_image)
+ if blend_mode == 'multply':
+ ret_image = ImageChops.multiply(background_image,layer_image)
+ if blend_mode == 'screen':
+ ret_image = ImageChops.screen(background_image, layer_image)
+ if blend_mode == 'add':
+ ret_image = ImageChops.add(background_image, layer_image, 1, 0)
+ if blend_mode == 'subtract':
+ ret_image = ImageChops.subtract(background_image, layer_image, 1, 0)
+ if blend_mode == 'difference':
+ ret_image = ImageChops.difference(background_image, layer_image)
+ if blend_mode == 'darker':
+ ret_image = ImageChops.darker(background_image, layer_image)
+ if blend_mode == 'lighter':
+ ret_image = ImageChops.lighter(background_image, layer_image)
+ if blend_mode == 'color_burn':
+ ret_image = blend_color_burn(background_image, layer_image)
+ if blend_mode == 'color_dodge':
+ ret_image = blend_color_dodge(background_image, layer_image)
+ if blend_mode == 'linear_burn':
+ ret_image = blend_linear_burn(background_image, layer_image)
+ if blend_mode == 'linear_dodge':
+ ret_image = blend_linear_dodge(background_image, layer_image)
+ if blend_mode == 'overlay':
+ ret_image = blend_overlay(background_image, layer_image)
+ if blend_mode == 'soft_light':
+ ret_image = blend_soft_light(background_image, layer_image)
+ if blend_mode == 'hard_light':
+ ret_image = blend_hard_light(background_image, layer_image)
+ if blend_mode == 'vivid_light':
+ ret_image = blend_vivid_light(background_image, layer_image)
+ if blend_mode == 'pin_light':
+ ret_image = blend_pin_light(background_image, layer_image)
+ if blend_mode == 'linear_light':
+ ret_image = blend_linear_light(background_image, layer_image)
+ if blend_mode == 'hard_mix':
+ ret_image = blend_hard_mix(background_image, layer_image)
+ # opacity
+ if opacity == 0:
+ ret_image = background_image
+ elif opacity < 100:
+ alpha = 1.0 - float(opacity) / 100
+ ret_image = Image.blend(ret_image, background_image, alpha)
+ return ret_image
+
+def chop_image_v2(background_image:Image, layer_image:Image, blend_mode:str, opacity:int) -> Image:
+
+ backdrop_prepped = np.asfarray(background_image.convert('RGBA'))
+ source_prepped = np.asfarray(layer_image.convert('RGBA'))
+ blended_np = BLEND_MODES[blend_mode](backdrop_prepped, source_prepped, opacity / 100)
+
+ # final_tensor = (torch.from_numpy(blended_np / 255)).unsqueeze(0)
+ # return tensor2pil(_tensor)
+
+ return Image.fromarray(np.uint8(blended_np)).convert('RGB')
+
+def remove_background(image:Image, mask:Image, color:str) -> Image:
+ width = image.width
+ height = image.height
+ ret_image = Image.new('RGB', size=(width, height), color=color)
+ ret_image.paste(image, mask=mask)
+ return ret_image
+
+def sharpen(image:Image) -> Image:
+ img = pil2cv2(image)
+ Laplace_kernel = np.array([[-1, -1, -1],
+ [-1, 9, -1],
+ [-1, -1, -1]], dtype=np.float32)
+ ret_image = cv2.filter2D(img, -1, Laplace_kernel)
+ return cv22pil(ret_image)
+
+def gaussian_blur(image:Image, radius:int) -> Image:
+ # image = image.convert("RGBA")
+ ret_image = image.filter(ImageFilter.GaussianBlur(radius=radius))
+ return ret_image
+
+def motion_blur(image:Image, angle:int, blur:int) -> Image:
+ angle += 45
+ blur *= 5
+ image = np.array(pil2cv2(image))
+ M = cv2.getRotationMatrix2D((blur / 2, blur / 2), angle, 1)
+ motion_blur_kernel = np.diag(np.ones(blur))
+ motion_blur_kernel = cv2.warpAffine(motion_blur_kernel, M, (blur, blur))
+ motion_blur_kernel = motion_blur_kernel / blur
+ blurred = cv2.filter2D(image, -1, motion_blur_kernel)
+ # convert to uint8
+ cv2.normalize(blurred, blurred, 0, 255, cv2.NORM_MINMAX)
+ blurred = np.array(blurred, dtype=np.uint8)
+ ret_image = cv22pil(blurred)
+ return ret_image
+
+def __apply_vignette(image, vignette):
+ # If image needs to be normalized (0-1 range)
+ needs_normalization = image.max() > 1
+ if needs_normalization:
+ image = image.astype(np.float32) / 255
+ final_image = np.clip(image * vignette[..., np.newaxis], 0, 1)
+ if needs_normalization:
+ final_image = (final_image * 255).astype(np.uint8)
+ return final_image
+def vignette_image(image:Image, intensity: float, center_x: float, center_y: float) -> Image:
+ image = pil2tensor(image)
+ _, height, width, _ = image.shape
+ # Generate the vignette for each image in the batch
+ # Create linear space but centered around the provided center point ratios
+ x = np.linspace(-1, 1, width)
+ y = np.linspace(-1, 1, height)
+ X, Y = np.meshgrid(x - (2 * center_x - 1), y - (2 * center_y - 1))
+ # Calculate distances to the furthest corner
+ distances_to_corners = [
+ np.sqrt((0 - center_x) ** 2 + (0 - center_y) ** 2),
+ np.sqrt((1 - center_x) ** 2 + (0 - center_y) ** 2),
+ np.sqrt((0 - center_x) ** 2 + (1 - center_y) ** 2),
+ np.sqrt((1 - center_x) ** 2 + (1 - center_y) ** 2)
+ ]
+ max_distance_to_corner = np.max(distances_to_corners)
+ radius = np.sqrt(X ** 2 + Y ** 2)
+ radius = radius / (max_distance_to_corner * np.sqrt(2)) # Normalize radius
+ opacity = np.clip(intensity, 0, 1)
+ vignette = 1 - radius * opacity
+ tensor_image = image.numpy()
+ # Apply vignette
+ vignette_image = __apply_vignette(tensor_image, vignette)
+ return tensor2pil(torch.from_numpy(vignette_image).unsqueeze(0))
+
+def RGB2YCbCr(t):
+ YCbCr = t.detach().clone()
+ YCbCr[:,:,:,0] = 0.2123 * t[:,:,:,0] + 0.7152 * t[:,:,:,1] + 0.0722 * t[:,:,:,2]
+ YCbCr[:,:,:,1] = 0 - 0.1146 * t[:,:,:,0] - 0.3854 * t[:,:,:,1] + 0.5 * t[:,:,:,2]
+ YCbCr[:,:,:,2] = 0.5 * t[:,:,:,0] - 0.4542 * t[:,:,:,1] - 0.0458 * t[:,:,:,2]
+ return YCbCr
+
+def YCbCr2RGB(t):
+ RGB = t.detach().clone()
+ RGB[:,:,:,0] = t[:,:,:,0] + 1.5748 * t[:,:,:,2]
+ RGB[:,:,:,1] = t[:,:,:,0] - 0.1873 * t[:,:,:,1] - 0.4681 * t[:,:,:,2]
+ RGB[:,:,:,2] = t[:,:,:,0] + 1.8556 * t[:,:,:,1]
+ return RGB
+
+# gaussian blur a tensor image batch in format [B x H x W x C] on H/W (spatial, per-image, per-channel)
+def cv_blur_tensor(images, dx, dy):
+ if min(dx, dy) > 100:
+ np_img = torch.nn.functional.interpolate(images.detach().clone().movedim(-1,1), scale_factor=0.1, mode='bilinear').movedim(1,-1).cpu().numpy()
+ for index, image in enumerate(np_img):
+ np_img[index] = cv2.GaussianBlur(image, (dx // 20 * 2 + 1, dy // 20 * 2 + 1), 0)
+ return torch.nn.functional.interpolate(torch.from_numpy(np_img).movedim(-1,1), size=(images.shape[1], images.shape[2]), mode='bilinear').movedim(1,-1)
+ else:
+ np_img = images.detach().clone().cpu().numpy()
+ for index, image in enumerate(np_img):
+ np_img[index] = cv2.GaussianBlur(image, (dx, dy), 0)
+ return torch.from_numpy(np_img)
+
+def image_add_grain(image:Image, scale:float=0.5, strength:float=0.5, saturation:float=0.7, toe:float=0.0, seed:int=0) -> Image:
+
+ image = pil2tensor(image.convert("RGB"))
+ t = image.detach().clone()
+ torch.manual_seed(seed)
+ grain = torch.rand(t.shape[0], int(t.shape[1] // scale), int(t.shape[2] // scale), 3)
+
+ YCbCr = RGB2YCbCr(grain)
+ YCbCr[:, :, :, 0] = cv_blur_tensor(YCbCr[:, :, :, 0], 3, 3)
+ YCbCr[:, :, :, 1] = cv_blur_tensor(YCbCr[:, :, :, 1], 15, 15)
+ YCbCr[:, :, :, 2] = cv_blur_tensor(YCbCr[:, :, :, 2], 11, 11)
+
+ grain = (YCbCr2RGB(YCbCr) - 0.5) * strength
+ grain[:, :, :, 0] *= 2
+ grain[:, :, :, 2] *= 3
+ grain += 1
+ grain = grain * saturation + grain[:, :, :, 1].unsqueeze(3).repeat(1, 1, 1, 3) * (1 - saturation)
+
+ grain = torch.nn.functional.interpolate(grain.movedim(-1, 1), size=(t.shape[1], t.shape[2]),
+ mode='bilinear').movedim(1, -1)
+ t[:, :, :, :3] = torch.clip((1 - (1 - t[:, :, :, :3]) * grain) * (1 - toe) + toe, 0, 1)
+ return tensor2pil(t)
+
+def filmgrain_image(image:Image, scale:float, grain_power:float,
+ shadows:float, highs:float, grain_sat:float,
+ sharpen:int=1, grain_type:int=4, src_gamma:float=1.0,
+ gray_scale:bool=False, seed:int=0) -> Image:
+ # image = pil2tensor(image)
+ # grain_type, 1=fine, 2=fine simple, 3=coarse, 4=coarser
+ grain_type_index = 3
+
+ # Apply grain
+ from .filmgrainer import filmgrainer as fg
+ grain_image = fg.process(image, scale=scale, src_gamma=src_gamma, grain_power=grain_power,
+ shadows=shadows, highs=highs, grain_type=grain_type_index,
+ grain_sat=grain_sat, gray_scale=gray_scale, sharpen=sharpen, seed=seed)
+ return tensor2pil(torch.from_numpy(grain_image).unsqueeze(0))
+
+def __apply_radialblur(image, blur_strength, radial_mask, focus_spread, steps):
+ from .filmgrainer import processing as processing_utils
+ needs_normalization = image.max() > 1
+ if needs_normalization:
+ image = image.astype(np.float32) / 255
+ blurred_images = processing_utils.generate_blurred_images(image, blur_strength, steps, focus_spread)
+ final_image = processing_utils.apply_blurred_images(image, blurred_images, radial_mask)
+ if needs_normalization:
+ final_image = np.clip(final_image * 255, 0, 255).astype(np.uint8)
+ return final_image
+
+def radialblur_image(image:Image, blur_strength:float, center_x:float, center_y:float, focus_spread:float, steps:int=5) -> Image:
+ width, height = image.size
+ image = pil2tensor(image)
+ if image.dim() == 4:
+ image = image[0]
+
+ # _, height, width, = image.shape
+ # Generate the vignette for each image in the batch
+ c_x, c_y = int(width * center_x), int(height * center_y)
+ # Calculate distances to all corners from the center
+ distances_to_corners = [
+ np.sqrt((c_x - 0)**2 + (c_y - 0)**2),
+ np.sqrt((c_x - width)**2 + (c_y - 0)**2),
+ np.sqrt((c_x - 0)**2 + (c_y - height)**2),
+ np.sqrt((c_x - width)**2 + (c_y - height)**2)
+ ]
+ max_distance_to_corner = max(distances_to_corners)
+ # Create and adjust radial mask
+ X, Y = np.meshgrid(np.arange(width) - c_x, np.arange(height) - c_y)
+ radial_mask = np.sqrt(X**2 + Y**2) / max_distance_to_corner
+ tensor_image = image.numpy()
+ # Apply blur
+ blur_image = __apply_radialblur(tensor_image, blur_strength, radial_mask, focus_spread, steps)
+ return tensor2pil(torch.from_numpy(blur_image).unsqueeze(0))
+
+def __apply_depthblur(image, depth_map, blur_strength, focal_depth, focus_spread, steps):
+ from .filmgrainer import processing as processing_utils
+ # Normalize the input image if needed
+ needs_normalization = image.max() > 1
+ if needs_normalization:
+ image = image.astype(np.float32) / 255
+ # Normalize the depth map if needed
+ depth_map = depth_map.astype(np.float32) / 255 if depth_map.max() > 1 else depth_map
+ # Resize depth map to match the image dimensions
+ depth_map_resized = cv2.resize(depth_map, (image.shape[1], image.shape[0]), interpolation=cv2.INTER_LINEAR)
+ if len(depth_map_resized.shape) > 2:
+ depth_map_resized = cv2.cvtColor(depth_map_resized, cv2.COLOR_BGR2GRAY)
+ # Adjust the depth map based on the focal plane
+ depth_mask = np.abs(depth_map_resized - focal_depth)
+ depth_mask = np.clip(depth_mask / np.max(depth_mask), 0, 1)
+ # Generate blurred versions of the image
+ blurred_images = processing_utils.generate_blurred_images(image, blur_strength, steps, focus_spread)
+ # Use the adjusted depth map as a mask for applying blurred images
+ final_image = processing_utils.apply_blurred_images(image, blurred_images, depth_mask)
+ # Convert back to original range if the image was normalized
+ if needs_normalization:
+ final_image = np.clip(final_image * 255, 0, 255).astype(np.uint8)
+ return final_image
+
+def depthblur_image(image:Image, depth_map:Image, blur_strength:float, focal_depth:float, focus_spread:float, steps:int=5) -> Image:
+ width, height = image.size
+ image = pil2tensor(image)
+ depth_map = pil2tensor(depth_map)
+ if image.dim() == 4:
+ image = image[0]
+ if depth_map.dim() == 4:
+ depth_map = depth_map[0]
+ tensor_image = image.numpy()
+ tensor_image_depth = depth_map.numpy()
+ # Apply blur
+ blur_image = __apply_depthblur(tensor_image, tensor_image_depth, blur_strength, focal_depth, focus_spread, steps)
+ return tensor2pil(torch.from_numpy(blur_image).unsqueeze(0))
+
+def fit_resize_image(image:Image, target_width:int, target_height:int, fit:str, resize_sampler:str, background_color:str = '#000000') -> Image:
+ image = image.convert('RGB')
+ orig_width, orig_height = image.size
+ if image is not None:
+ if fit == 'letterbox':
+ if orig_width / orig_height > target_width / target_height: # 更宽,上下留黑
+ fit_width = target_width
+ fit_height = int(target_width / orig_width * orig_height)
+ else: # 更瘦,左右留黑
+ fit_height = target_height
+ fit_width = int(target_height / orig_height * orig_width)
+ fit_image = image.resize((fit_width, fit_height), resize_sampler)
+ ret_image = Image.new('RGB', size=(target_width, target_height), color=background_color)
+ ret_image.paste(fit_image, box=((target_width - fit_width)//2, (target_height - fit_height)//2))
+ elif fit == 'crop':
+ if orig_width / orig_height > target_width / target_height: # 更宽,裁左右
+ fit_width = int(orig_height * target_width / target_height)
+ fit_image = image.crop(
+ ((orig_width - fit_width)//2, 0, (orig_width - fit_width)//2 + fit_width, orig_height))
+ else: # 更瘦,裁上下
+ fit_height = int(orig_width * target_height / target_width)
+ fit_image = image.crop(
+ (0, (orig_height-fit_height)//2, orig_width, (orig_height-fit_height)//2 + fit_height))
+ ret_image = fit_image.resize((target_width, target_height), resize_sampler)
+ else:
+ ret_image = image.resize((target_width, target_height), resize_sampler)
+ return ret_image
+
+def __rotate_expand(image:Image, angle:float, SSAA:int=0, method:str="lanczos") -> Image:
+ images = pil2tensor(image)
+ expand = "true"
+ height, width = images[0, :, :, 0].shape
+
+ def rotate_tensor(tensor):
+ resize_sampler = Image.LANCZOS
+ rotate_sampler = Image.BICUBIC
+ if method == "bicubic":
+ resize_sampler = Image.BICUBIC
+ rotate_sampler = Image.BICUBIC
+ elif method == "hamming":
+ resize_sampler = Image.HAMMING
+ rotate_sampler = Image.BILINEAR
+ elif method == "bilinear":
+ resize_sampler = Image.BILINEAR
+ rotate_sampler = Image.BILINEAR
+ elif method == "box":
+ resize_sampler = Image.BOX
+ rotate_sampler = Image.NEAREST
+ elif method == "nearest":
+ resize_sampler = Image.NEAREST
+ rotate_sampler = Image.NEAREST
+ img = tensor2pil(tensor)
+ if SSAA > 1:
+ img_us_scaled = img.resize((width * SSAA, height * SSAA), resize_sampler)
+ img_rotated = img_us_scaled.rotate(angle, rotate_sampler, expand == "true", fillcolor=(0, 0, 0, 0))
+ img_down_scaled = img_rotated.resize((img_rotated.width // SSAA, img_rotated.height // SSAA), resize_sampler)
+ result = pil2tensor(img_down_scaled)
+ else:
+ img_rotated = img.rotate(angle, rotate_sampler, expand == "true", fillcolor=(0, 0, 0, 0))
+ result = pil2tensor(img_rotated)
+ return result
+
+ if angle == 0.0 or angle == 360.0:
+ return tensor2pil(images)
+ else:
+ rotated_tensor = torch.stack([rotate_tensor(images[i]) for i in range(len(images))])
+ return tensor2pil(rotated_tensor).convert('RGB')
+
+def image_rotate_extend_with_alpha(image:Image, angle:float, alpha:Image=None, method:str="lanczos", SSAA:int=0) -> tuple:
+ _image = __rotate_expand(image.convert('RGB'), angle, SSAA, method)
+ if angle is not None:
+ _alpha = __rotate_expand(alpha.convert('RGB'), angle, SSAA, method)
+ ret_image = RGB2RGBA(_image, _alpha)
+ else:
+ ret_image = _image
+ return (_image, _alpha.convert('L'), ret_image)
+
+def create_box_gradient(start_color_inhex:str, end_color_inhex:str, width:int, height:int, scale:int=50) -> Image:
+ # scale is percent of border to center for the rectangle
+ if scale > 100:
+ scale = 100
+ elif scale < 1:
+ scale = 1
+ start_color = Hex_to_RGB(start_color_inhex)
+ end_color = Hex_to_RGB(end_color_inhex)
+ ret_image = Image.new("RGB", (width, height), start_color)
+ draw = ImageDraw.Draw(ret_image)
+ step = int(min(width, height) * scale / 100 / 2)
+ if step > 0:
+ for i in range(step):
+ R = int(start_color[0] * (step - i) / step + end_color[0] * i / step)
+ G = int(start_color[1] * (step - i) / step + end_color[1] * i / step)
+ B = int(start_color[2] * (step - i) / step + end_color[2] * i / step)
+ color = (R, G, B)
+ draw.rectangle((i, i, width - i, height - i), fill=color)
+ draw.rectangle((step, step, width - step, height - step), fill=end_color)
+ return ret_image
+
+def create_gradient(start_color_inhex:str, end_color_inhex:str, width:int, height:int, direction:str='bottom') -> Image:
+ # direction = one of top, bottom, left, right
+ start_color = Hex_to_RGB(start_color_inhex)
+ end_color = Hex_to_RGB(end_color_inhex)
+ ret_image = Image.new("RGB", (width, height), start_color)
+ draw = ImageDraw.Draw(ret_image)
+ if direction == 'bottom':
+ for i in range(height):
+ R = int(start_color[0] * (height - i) / height + end_color[0] * i / height)
+ G = int(start_color[1] * (height - i) / height + end_color[1] * i / height)
+ B = int(start_color[2] * (height - i) / height + end_color[2] * i / height)
+ color = (R, G, B)
+ draw.line((0, i, width, i), fill=color)
+ elif direction == 'top':
+ for i in range(height):
+ R = int(end_color[0] * (height - i) / height + start_color[0] * i / height)
+ G = int(end_color[1] * (height - i) / height + start_color[1] * i / height)
+ B = int(end_color[2] * (height - i) / height + start_color[2] * i / height)
+ color = (R, G, B)
+ draw.line((0, i, width, i), fill=color)
+ elif direction == 'right':
+ for i in range(width):
+ R = int(start_color[0] * (width - i) / width + end_color[0] * i / width)
+ G = int(start_color[1] * (width - i) / width + end_color[1] * i / width)
+ B = int(start_color[2] * (width - i) / width + end_color[2] * i / width)
+ color = (R, G, B)
+ draw.line((i, 0, i, height), fill=color)
+ elif direction == 'left':
+ for i in range(width):
+ R = int(end_color[0] * (width - i) / width + start_color[0] * i / width)
+ G = int(end_color[1] * (width - i) / width + start_color[1] * i / width)
+ B = int(end_color[2] * (width - i) / width + start_color[2] * i / width)
+ color = (R, G, B)
+ draw.line((i, 0, i, height), fill=color)
+ else:
+ log(f'A argument error of imagefunc.create_gradient(), '
+ f'"direction=" must one of "top, bottom, left, right".',
+ message_type='error')
+
+ return ret_image
+
+def gradient(start_color_inhex:str, end_color_inhex:str, width:int, height:int, angle:float, ) -> Image:
+ radius = int((width + height) / 4)
+ g = create_gradient(start_color_inhex, end_color_inhex, radius, radius)
+ _canvas = Image.new('RGB', size=(radius, radius*3), color=start_color_inhex)
+ top = Image.new('RGB', size=(radius, radius), color=start_color_inhex)
+ bottom = Image.new('RGB', size=(radius, radius),color=end_color_inhex)
+ _canvas.paste(top, box=(0, 0, radius, radius))
+ _canvas.paste(g, box=(0, radius, radius, radius * 2))
+ _canvas.paste(bottom,box=(0, radius * 2, radius, radius * 3))
+ _canvas = _canvas.resize((radius * 3, radius * 3))
+ _canvas = __rotate_expand(_canvas,angle)
+ center = int(_canvas.width / 2)
+ _x = int(width / 3)
+ _y = int(height / 3)
+ ret_image = _canvas.crop((center - _x, center - _y, center + _x, center + _y))
+ ret_image = ret_image.resize((width, height))
+ return ret_image
+
+def draw_rect(image:Image, x:int, y:int, width:int, height:int, line_color:str, line_width:int,
+ box_color:str=None) -> Image:
+ draw = ImageDraw.Draw(image)
+ draw.rectangle((x, y, x + width, y + height), fill=box_color, outline=line_color, width=line_width, )
+ return image
+
+def draw_border(image:Image, border_width:int, color:str='#FFFFFF') -> Image:
+ return ImageOps.expand(image, border=border_width, fill=color)
+
+# 对灰度图像进行直方图均衡化
+def normalize_gray(image:Image) -> Image:
+ if image.mode != 'L':
+ image = image.convert('L')
+ img = np.asarray(image)
+ balanced_img = img.copy()
+ hist, bins = np.histogram(img.reshape(-1), 256, (0, 256))
+ bmin = np.min(np.where(hist > (hist.sum() * 0.0005)))
+ bmax = np.max(np.where(hist > (hist.sum() * 0.0005)))
+ balanced_img = np.clip(img, bmin, bmax)
+ balanced_img = ((balanced_img - bmin) / (bmax - bmin) * 255)
+ return Image.fromarray(balanced_img).convert('L')
+
+def remap_pixel(pixel:int, min_brightness:int, max_brightness:int) -> int:
+ return int((pixel - min_brightness) / (max_brightness - min_brightness) * 255)
+def histogram_range(image:Image, black_point:int, black_range:int, white_point:int, white_range:int) -> Image:
+
+ if image.mode != 'L':
+ image = image.convert('L')
+
+ if black_point == 255:
+ black_point = 254
+ if white_point == 0:
+ white_point = 1
+ if black_point + black_range > 255:
+ black_range = 255 - black_point
+ if white_range > white_point:
+ white_range = white_point
+
+ white_image = Image.new("L", size=image.size, color="white")
+ black_image = Image.new("L", size=image.size, color="black")
+
+ if black_point == white_point:
+ return white_image
+
+
+ # draw white part
+ white_part = black_image
+ if white_point < 255 or white_range > 0:
+ for y in (range(image.height)):
+ for x in range(image.width):
+ pixel = image.getpixel((x, y))
+ if pixel > white_point: # put white
+ white_part.putpixel((x, y), 255)
+ elif pixel > white_point - white_range:
+ pixel = remap_pixel(pixel, white_point - white_range, white_point)
+ white_part.putpixel((x, y), pixel)
+ white_part = ImageChops.invert(white_part)
+
+
+ # draw black part
+ black_part = black_image
+ if black_point > 0 or black_range > 0:
+ for y in (range(image.height)):
+ for x in range(image.width):
+ pixel = image.getpixel((x, y))
+ if pixel < black_point: # put black
+ black_part.putpixel((x, y), 255)
+ elif pixel < black_point + black_range:
+ pixel = remap_pixel(pixel, black_point, black_point + black_range)
+ black_part.putpixel((x, y), 255 - pixel)
+ black_part = ImageChops.invert(black_part)
+
+ ret_image = chop_image_v2(white_part, black_part, blend_mode='darken', opacity=100)
+
+ return ret_image
+
+def histogram_equalization(image:Image, mask:Image=None, gamma_strength=0.5) -> Image:
+
+ if image.mode != 'L':
+ image = image.convert('L')
+
+ if mask is not None:
+ if mask.mode != 'L':
+ mask = mask.convert('L')
+ else:
+ mask = Image.new('L', size=image.size, color = 'white')
+
+ # calculate Min/Max brightness pixel
+ min_brightness = 255
+ max_brightness = 0
+ average_brightness = 0
+ total_pixel = 0
+ for y in range(image.height):
+ for x in range(image.width):
+ if mask.getpixel((x, y)) == 0:
+ continue
+ else:
+ pixel = image.getpixel((x, y))
+ if pixel < min_brightness:
+ min_brightness = pixel
+ if pixel > max_brightness:
+ max_brightness = pixel
+ average_brightness += pixel
+ total_pixel += 1
+ if total_pixel == 0:
+ log(f"histogram_equalization: mask is not available, return orinianl image.")
+ return image
+ average_brightness = int(average_brightness / total_pixel)
+
+ for y in range(image.height):
+ for x in range(image.width):
+ pixel = image.getpixel((x, y))
+ image.putpixel((x, y), remap_pixel(pixel, min_brightness, max_brightness))
+
+ image = gamma_trans(image, (average_brightness - 127) / 127 * gamma_strength * 0.66 + 1)
+
+ return image.convert('L')
+
+def adjust_levels(image:Image, input_black:int=0, input_white:int=255, midtones:float=1.0,
+ output_black:int=0, output_white:int=255) -> Image:
+
+ if input_black == input_white or output_black == output_white:
+ return Image.new('RGB', size=image.size, color='gray')
+
+ img = pil2cv2(image).astype(np.float64)
+
+ if input_black > input_white:
+ input_black, input_white = input_white, input_black
+ if output_black > output_white:
+ output_black, output_white = output_white, output_black
+
+
+ # input_levels remap
+ if input_black > 0 or input_white < 255:
+ img = 255 * ((img - input_black) / (input_white - input_black))
+ img[img < 0] = 0
+ img[img > 255] = 255
+
+ # # mid_tone
+ if midtones != 1.0:
+ img = 255 * np.power(img / 255, 1.0 / midtones)
+
+ img[img < 0] = 0
+ img[img > 255] = 255
+
+ # output_levels remap
+ if output_black > 0 or output_white < 255:
+ img = (img / 255) * (output_white - output_black) + output_black
+ img[img < 0] = 0
+ img[img > 255] = 255
+
+ img = img.astype(np.uint8)
+ return cv22pil(img)
+
+def get_image_color_tone(image:Image, mask:Image=None) -> str:
+ image = image.convert('RGB')
+ max_score = 0.0001
+ dominant_color = (255, 255, 255)
+ if mask is not None:
+ if mask.mode != 'L':
+ mask = mask.convert('L')
+ canvas = Image.new('RGB', size=image.size, color='black')
+ canvas.paste(image, mask=mask)
+ image = canvas
+
+ all_colors = image.getcolors(image.width * image.height)
+ for count, (r, g, b) in all_colors:
+ if mask is not None:
+ if r + g + b < 2: # 忽略黑色
+ continue
+ saturation = rgb_to_hsv(r / 255.0, g / 255.0, b / 255.0)[1]
+ y = min(abs(r * 2104 + g * 4130 + b * 802 + 4096 + 131072) >> 13,235)
+ y = (y - 16.0) / (235 - 16)
+ score = (saturation+0.1)*count
+ if score > max_score:
+ max_score = score
+ dominant_color = (r, g, b)
+ ret_color = RGB_to_Hex(dominant_color)
+ return ret_color
+
+def get_image_color_average(image:Image, mask:Image=None) -> str:
+ image = image.convert('RGB')
+ width, height = image.size
+ total_red = 0
+ total_green = 0
+ total_blue = 0
+ total_pixel =0
+ for y in range(height):
+ for x in range(width):
+ if mask is not None:
+ if mask.mode != 'L':
+ mask = mask.convert('L')
+ if mask.getpixel((x, y)) <= 127:
+ continue
+ rgb = image.getpixel((x, y))
+ total_red += rgb[0]
+ total_green += rgb[1]
+ total_blue += rgb[2]
+ total_pixel += 1
+
+ average_red = total_red // total_pixel
+ average_green = total_green // total_pixel
+ average_blue = total_blue // total_pixel
+ color = (average_red, average_green, average_blue)
+ ret_color = RGB_to_Hex(color)
+ return ret_color
+
+def get_gray_average(image:Image, mask:Image=None) -> int:
+ # image.mode = 'HSV', mask.mode = 'L'
+ image = image.convert('HSV')
+
+ if mask is not None:
+ if mask.mode != 'L':
+ mask = mask.convert('L')
+ else:
+ mask = Image.new('L', size=image.size, color='white')
+ _, _, _v = image.convert('HSV').split()
+ _v = np.array(_v)
+ average_gray = _v[np.array(mask) > 16].mean()
+ # width, height = image.size
+ # total_gray = 0
+ # valid_pixels = 0
+ # for y in range(height):
+ # for x in range(width):
+ # if mask is not None:
+ # if mask.getpixel((x, y)) > 16: #mask亮度低于16的忽略不计
+ # gray = _v.getpixel((x, y))
+ # total_gray += gray
+ # valid_pixels += 1
+ # else:
+ # gray = _v.getpixel((x, y))
+ # total_gray += gray
+ # valid_pixels += 1
+ # average_gray = total_gray // valid_pixels
+ return average_gray
+
+def calculate_shadow_highlight_level(gray:int) -> float:
+ range = 255
+ shadow_exponent = 3
+ highlight_exponent = 2
+ shadow_ratio = gray ** shadow_exponent / range ** shadow_exponent
+ highlight_ratio = gray ** highlight_exponent / range ** highlight_exponent
+ shadow_level = shadow_ratio * 100 + (1 - shadow_ratio) * 32
+ highlight_level = highlight_ratio * 100 + (1 - highlight_ratio) * 32
+ return shadow_level, highlight_level
+
+def luminance_keyer(image:Image, low:float=0, high:float=1, gamma:float=1) -> Image:
+ image = pil2tensor(image)
+ t = image[:, :, :, :3].detach().clone()
+ alpha = 0.2126 * t[:, :, :, 0] + 0.7152 * t[:, :, :, 1] + 0.0722 * t[:, :, :, 2]
+ if low == high:
+ alpha = (alpha > high).to(t.dtype)
+ else:
+ alpha = (alpha - low) / (high - low)
+ if gamma != 1.0:
+ alpha = torch.pow(alpha, 1 / gamma)
+ alpha = torch.clamp(alpha, min=0, max=1).unsqueeze(3).repeat(1, 1, 1, 3)
+ return tensor2pil(alpha).convert('L')
+
+def get_image_bright_average(image:Image) -> int:
+ image = image.convert('L')
+ width, height = image.size
+ total_bright = 0
+ pixels = 0
+ for y in range(height):
+ for x in range(width):
+ b = image.getpixel((x, y))
+ if b > 1: # 排除死黑
+ pixels += 1
+ total_bright += b
+ return int(total_bright / pixels)
+
+def image_channel_split(image:Image, mode = 'RGBA') -> tuple:
+ _image = image.convert('RGBA')
+ channel1 = Image.new('L', size=_image.size, color='black')
+ channel2 = Image.new('L', size=_image.size, color='black')
+ channel3 = Image.new('L', size=_image.size, color='black')
+ channel4 = Image.new('L', size=_image.size, color='black')
+ if mode == 'RGBA':
+ channel1, channel2, channel3, channel4 = _image.split()
+ if mode == 'RGB':
+ channel1, channel2, channel3 = _image.convert('RGB').split()
+ if mode == 'YCbCr':
+ channel1, channel2, channel3 = _image.convert('YCbCr').split()
+ if mode == 'LAB':
+ channel1, channel2, channel3 = _image.convert('LAB').split()
+ if mode == 'HSV':
+ channel1, channel2, channel3 = _image.convert('HSV').split()
+ return channel1, channel2, channel3, channel4
+
+def image_channel_merge(channels:tuple, mode = 'RGB' ) -> Image:
+ channel1 = channels[0].convert('L')
+ channel2 = channels[1].convert('L')
+ channel3 = channels[2].convert('L')
+ channel4 = Image.new('L', size=channel1.size, color='white')
+ if mode == 'RGBA':
+ if len(channels) > 3:
+ channel4 = channels[3].convert('L')
+ ret_image = Image.merge('RGBA',[channel1, channel2, channel3, channel4])
+ elif mode == 'RGB':
+ ret_image = Image.merge('RGB', [channel1, channel2, channel3])
+ elif mode == 'YCbCr':
+ ret_image = Image.merge('YCbCr', [channel1, channel2, channel3]).convert('RGB')
+ elif mode == 'LAB':
+ ret_image = Image.merge('LAB', [channel1, channel2, channel3]).convert('RGB')
+ elif mode == 'HSV':
+ ret_image = Image.merge('HSV', [channel1, channel2, channel3]).convert('RGB')
+ return ret_image
+
+def image_gray_offset(image:Image, offset:int) -> Image:
+ image = image.convert('L')
+ image_array = np.array(image, dtype=np.int16)
+ image_array = np.clip(image_array + offset, 0, 255).astype(np.uint8)
+ ret_image = Image.fromarray(image_array, mode='L')
+ return ret_image
+
+def image_gray_ratio(image:Image, ratio:float) -> Image:
+ image = image.convert('L')
+ image_array = np.array(image, dtype=np.float32)
+ image_array = np.clip(image_array * ratio, 0, 255).astype(np.uint8)
+ ret_image = Image.fromarray(image_array, mode='L')
+ return ret_image
+
+def image_hue_offset(image:Image, offset:int) -> Image:
+ image = image.convert('L')
+ image_array = np.array(image, dtype=np.int16)
+ image_array = (image_array + offset) % 256
+ image_array = image_array.astype(np.uint8)
+ ret_image = Image.fromarray(image_array, mode='L')
+
+ return ret_image
+
+def gamma_trans(image:Image, gamma:float) -> Image:
+ cv2_image = pil2cv2(image)
+ gamma_table = [np.power(x/255.0,gamma)*255.0 for x in range(256)]
+ gamma_table = np.round(np.array(gamma_table)).astype(np.uint8)
+ _corrected = cv2.LUT(cv2_image,gamma_table)
+ return cv22pil(_corrected)
+
+
+def read_LUT_IridasCube_encode_utf8(path: str):
+ from colour.utilities import as_float_array, as_int_scalar
+ from colour.io.luts.lut import LUT3x1D, LUT3D
+ title = re.sub("_|-|\\.", " ", os.path.splitext(os.path.basename(path))[0])
+ domain_min, domain_max = np.array([0, 0, 0]), np.array([1, 1, 1])
+ dimensions: int = 3
+ size: int = 2
+ data = []
+ comments = []
+
+ with open(path, encoding='utf-8') as cube_file:
+ lines = cube_file.readlines()
+ for line in lines:
+
+ line = line.strip() # noqa: PLW2901
+
+ if len(line) == 0:
+ continue
+
+ if line.startswith("#"):
+ comments.append(line[1:].strip())
+ continue
+
+ tokens = line.split()
+ if tokens[0] == "TITLE":
+ title = " ".join(tokens[1:])[1:-1]
+ elif tokens[0] == "DOMAIN_MIN":
+ domain_min = as_float_array(tokens[1:])
+ elif tokens[0] == "DOMAIN_MAX":
+ domain_max = as_float_array(tokens[1:])
+ elif tokens[0] == "LUT_1D_SIZE":
+ dimensions = 2
+ size = as_int_scalar(tokens[1])
+ elif tokens[0] == "LUT_3D_SIZE":
+ dimensions = 3
+ size = as_int_scalar(tokens[1])
+ else:
+ data.append(tokens)
+
+ table = as_float_array(data)
+
+ LUT: LUT3x1D | LUT3D
+ if dimensions == 2:
+ LUT = LUT3x1D(
+ table,
+ title,
+ np.vstack([domain_min, domain_max]),
+ comments=comments,
+ )
+ elif dimensions == 3:
+ # The lines of table data shall be in ascending index order,
+ # with the first component index (Red) changing most rapidly,
+ # and the last component index (Blue) changing least rapidly.
+ table = table.reshape([size, size, size, 3], order="F")
+
+ LUT = LUT3D(
+ table,
+ title,
+ np.vstack([domain_min, domain_max]),
+ comments=comments,
+ )
+
+ return LUT
+
+
+def apply_lut(image:Image, lut_file:str, colorspace:str, strength:int, clip_values:bool=True) -> Image:
+ """
+ Apply a LUT to an image.
+ :param image: Image to apply the LUT to.
+ :param lut_file: LUT file to apply.
+ :param colorspace: Colorspace to convert the image to before applying the LUT.
+ :param clip_values: Clip the values of the LUT to the domain of the LUT.
+ :param strength: Strength of the LUT.
+ :return: Image with the LUT applied.
+ """
+ log_colorspace = False
+ if colorspace == "log":
+ log_colorspace = True
+
+ # from colour.io.luts.iridas_cube import read_LUT_IridasCube
+
+ lut = read_LUT_IridasCube_encode_utf8(lut_file)
+ lut.name = lut_file
+
+ if clip_values:
+ if lut.domain[0].max() == lut.domain[0].min() and lut.domain[1].max() == lut.domain[1].min():
+ lut.table = np.clip(lut.table, lut.domain[0, 0], lut.domain[1, 0])
+ else:
+ if len(lut.table.shape) == 2: # 3x1D
+ for dim in range(3):
+ lut.table[:, dim] = np.clip(lut.table[:, dim], lut.domain[0, dim], lut.domain[1, dim])
+ else: # 3D
+ for dim in range(3):
+ lut.table[:, :, :, dim] = np.clip(lut.table[:, :, :, dim], lut.domain[0, dim], lut.domain[1, dim])
+
+ img = pil2tensor(image)
+ lut_img = img.numpy().copy()
+ is_non_default_domain = not np.array_equal(lut.domain, np.array([[0., 0., 0.], [1., 1., 1.]]))
+ dom_scale = None
+ if is_non_default_domain:
+ dom_scale = lut.domain[1] - lut.domain[0]
+ lut_img = lut_img * dom_scale + lut.domain[0]
+ if log_colorspace:
+ lut_img = lut_img ** (1/2.2)
+ lut_img = lut.apply(lut_img)
+ if log_colorspace:
+ lut_img = lut_img ** (2.2)
+ if is_non_default_domain:
+ lut_img = (lut_img - lut.domain[0]) / dom_scale
+ lut_img = torch.from_numpy(lut_img)
+ if strength < 100:
+ strength /= 100
+ lut_img = strength * lut_img + (1 - strength) * img
+
+ return tensor2pil(lut_img)
+
+def color_adapter(image:Image, ref_image:Image) -> Image:
+ image = pil2cv2(image)
+ ref_image = pil2cv2(ref_image)
+ image = cv2.cvtColor(image, cv2.COLOR_BGR2LAB)
+ image_mean, image_std = calculate_mean_std(image)
+ ref_image = cv2.cvtColor(ref_image, cv2.COLOR_BGR2LAB)
+ ref_image_mean, ref_image_std = calculate_mean_std(ref_image)
+ _image = ((image - image_mean) * (ref_image_std / image_std)) + ref_image_mean
+ np.putmask(_image, _image > 255, values=255)
+ np.putmask(_image, _image < 0, values=0)
+ ret_image = cv2.cvtColor(cv2.convertScaleAbs(_image), cv2.COLOR_LAB2BGR)
+ return cv22pil(ret_image)
+
+def calculate_mean_std(image:Image):
+ mean, std = cv2.meanStdDev(image)
+ mean = np.hstack(np.around(mean, decimals=2))
+ std = np.hstack(np.around(std, decimals=2))
+ return mean, std
+
+def image_watercolor(image:Image, level:int=50) -> Image:
+ img = pil2cv2(image)
+ img = cv2.cvtColor(img, cv2.COLOR_BGR2RGB)
+ factor = (level / 128.0) ** 2
+ sigmaS= int((image.width + image.height) / 5.0 * factor) + 1
+ sigmaR = sigmaS / 32.0 * factor + 0.002
+ img_color = cv2.stylization(img, sigma_s=sigmaS, sigma_r=sigmaR)
+ ret_image = cv2.cvtColor(img_color, cv2.COLOR_BGR2RGB)
+ return cv22pil(ret_image)
+
+
+def image_beauty(image:Image, level:int=50) -> Image:
+ img = pil2cv2(image)
+ img = cv2.cvtColor(img, cv2.COLOR_BGR2RGB)
+ factor = (level / 50.0)**2
+ d = int((image.width + image.height) / 256 * factor)
+ sigmaColor = int((image.width + image.height) / 256 * factor)
+ sigmaSpace = int((image.width + image.height) / 160 * factor)
+ img_bit = cv2.bilateralFilter(src=img, d=d, sigmaColor=sigmaColor, sigmaSpace=sigmaSpace)
+ ret_image = cv2.cvtColor(img_bit, cv2.COLOR_BGR2RGB)
+ return cv22pil(ret_image)
+
+
+def pixel_spread(image:Image, mask:Image) -> Image:
+ from pymatting import estimate_foreground_ml
+ i1 = pil2tensor(image)
+ if mask.mode != 'RGB':
+ mask = mask.convert('RGB')
+ i_dup = copy.deepcopy(i1.cpu().numpy().astype(np.float64))
+ a_dup = copy.deepcopy(pil2tensor(mask).cpu().numpy().astype(np.float64))
+ fg = copy.deepcopy(i1.cpu().numpy().astype(np.float64))
+
+ for index, img in enumerate(i_dup):
+ alpha = a_dup[index][:, :, 0]
+ fg[index], _ = estimate_foreground_ml(img, np.array(alpha), return_background=True)
+
+ return tensor2pil(torch.from_numpy(fg.astype(np.float32)))
+
+
+def generate_text_image(text:str, font_path:str, font_size:int, text_color:str="#FFFFFF",
+ vertical:bool=True, stroke_width:int=1, stroke_color:str="#000000",
+ spacing:int=0, leading:int=0) -> tuple:
+
+ lines = text.split("\n")
+ if vertical:
+ layout = "vertical"
+ else:
+ layout = "horizontal"
+ char_coordinates = []
+ if layout == "vertical":
+ x = 0
+ y = 0
+ for i in range(len(lines)):
+ line = lines[i]
+ for char in line:
+ char_coordinates.append((x, y))
+ y += font_size + spacing
+ x += font_size + leading
+ y = 0
+ else:
+ x = 0
+ y = 0
+ for line in lines:
+ for char in line:
+ char_coordinates.append((x, y))
+ x += font_size + spacing
+ y += font_size + leading
+ x = 0
+ if layout == "vertical":
+ width = (len(lines) * (font_size + spacing)) - spacing
+ height = ((len(max(lines, key=len)) + 1) * (font_size + spacing)) + spacing
+ else:
+ width = (len(max(lines, key=len)) * (font_size + spacing)) - spacing
+ height = ((len(lines) - 1) * (font_size + spacing)) + font_size
+
+ image = Image.new('RGBA', size=(width, height), color=stroke_color)
+ draw = ImageDraw.Draw(image)
+ font = ImageFont.truetype(font_path, font_size)
+ index = 0
+ for i, line in enumerate(lines):
+ for j, char in enumerate(line):
+ x, y = char_coordinates[index]
+ if stroke_width > 0:
+ draw.text((x - stroke_width, y), char, font=font, fill=stroke_color)
+ draw.text((x + stroke_width, y), char, font=font, fill=stroke_color)
+ draw.text((x, y - stroke_width), char, font=font, fill=stroke_color)
+ draw.text((x, y + stroke_width), char, font=font, fill=stroke_color)
+ draw.text((x, y), char, font=font, fill=text_color)
+ index += 1
+ return (image.convert('RGB'), image.split()[3])
+
+def watermark_image_size(image:Image) -> int:
+ size = int(math.sqrt(image.width * image.height * 0.015625) * 0.9)
+ return size
+
+def add_invisibal_watermark(image:Image, watermark_image:Image) -> Image:
+ """
+ Adds an invisible watermark to an image.
+ """
+ orig_image_mode = image.mode
+ temp_dir = os.path.join(folder_paths.get_temp_directory(), generate_random_name('_watermark_', '_temp', 16))
+ if os.path.isdir(temp_dir):
+ shutil.rmtree(temp_dir)
+ image_dir = os.path.join(temp_dir, 'image')
+ wm_dir = os.path.join(temp_dir, 'wm')
+ result_dir = os.path.join(temp_dir, 'result')
+
+ try:
+ os.makedirs(image_dir)
+ os.makedirs(wm_dir)
+ os.makedirs(result_dir)
+ except Exception as e:
+ # print(e)
+ log(f"Error: {NODE_NAME} skipped, because unable to create temporary folder.", message_type='error')
+ return (image,)
+
+ image_file_name = os.path.join(generate_random_name('watermark_orig_', '_temp', 16) + '.png')
+ wm_file_name = os.path.join(generate_random_name('watermark_image_', '_temp', 16) + '.png')
+ output_file_name = os.path.join(generate_random_name('watermark_output_', '_temp', 16) + '.png')
+
+ try:
+ if image.mode != "RGB":
+ image = image.convert("RGB")
+ image.save(os.path.join(image_dir, image_file_name))
+ watermark_image.save(os.path.join(wm_dir, wm_file_name))
+ except IOError as e:
+ # print(e)
+ log(f"Error: {NODE_NAME} skipped, because unable to create temporary file.", message_type='error')
+ return (image,)
+
+ from blind_watermark import WaterMark
+ bwm1 = WaterMark(password_img=1, password_wm=1)
+ bwm1.read_img(os.path.join(image_dir, image_file_name))
+ bwm1.read_wm(os.path.join(wm_dir, wm_file_name))
+ output_image = os.path.join(result_dir, output_file_name)
+ bwm1.embed(output_image, compression_ratio=100)
+
+ return Image.open(output_image).convert(orig_image_mode)
+
+def decode_watermark(image:Image, watermark_image_size:int=94) -> Image:
+ temp_dir = os.path.join(folder_paths.get_temp_directory(), generate_random_name('_watermark_', '_temp', 16))
+ if os.path.isdir(temp_dir):
+ shutil.rmtree(temp_dir)
+ image_dir = os.path.join(temp_dir, 'decode_image')
+ result_dir = os.path.join(temp_dir, 'decode_result')
+
+ try:
+ os.makedirs(image_dir)
+ os.makedirs(result_dir)
+ except Exception as e:
+ # print(e)
+ log(f"Error: {NODE_NAME} skipped, because unable to create temporary folder.", message_type='error')
+ return (image,)
+
+ image_file_name = os.path.join(generate_random_name('watermark_decode_', '_temp', 16) + '.png')
+ output_file_name = os.path.join(generate_random_name('watermark_decode_output_', '_temp', 16) + '.png')
+
+ try:
+ image.save(os.path.join(image_dir, image_file_name))
+ except IOError as e:
+ # print(e)
+ log(f"Error: {NODE_NAME} skipped, because unable to create temporary file.", message_type='error')
+ return (image,)
+
+ from blind_watermark import WaterMark
+ bwm1 = WaterMark(password_img=1, password_wm=1)
+ decode_image = os.path.join(image_dir, image_file_name)
+ output_image = os.path.join(result_dir, output_file_name)
+
+ try:
+ bwm1.extract(filename=decode_image, wm_shape=(watermark_image_size, watermark_image_size),
+ out_wm_name=os.path.join(output_image),)
+ ret_image = Image.open(output_image)
+ except Exception as e:
+ log(f"blind watermark extract fail, {e}")
+ ret_image = Image.new("RGB", (64, 64), color="black")
+ ret_image = normalize_gray(ret_image)
+ return ret_image
+
+def generate_text_image(width:int, height:int, text:str, font_file:str, text_scale:float=1, font_color:str="#FFFFFF",) -> Image:
+ image = Image.new("RGBA", (width, height), (0, 0, 0, 0))
+ draw = ImageDraw.Draw(image)
+ font_size = int(width / len(text) * text_scale)
+ font = ImageFont.truetype(font_file, font_size)
+ bbox = draw.textbbox((0, 0), text, font=font)
+ text_width, text_height = bbox[2] - bbox[0], bbox[3] - bbox[1]
+ x = int((width - text_width) / 2)
+ y = int((height - text_height) / 2) - int(font_size / 2)
+ draw.text((x, y), text, font=font, fill=font_color)
+ return image
+
+'''Mask Functions'''
+
+def create_mask_from_color_cv2(image:Image, color:str, tolerance:int=0) -> Image:
+ (r, g, b) = Hex_to_RGB(color)
+ target_color = (b, g, r)
+ tolerance = 127 + int(tolerance * 1.28)
+ # tolerance = 255 - tolerance
+ # 将RGB颜色转换为HSV颜色空间
+ image = pil2cv2(image)
+ hsv_image = cv2.cvtColor(image, cv2.COLOR_BGR2HSV)
+
+ # 定义目标颜色的HSV范围
+ lower_color = np.array([max(target_color[0] - tolerance, 0), max(target_color[1] - tolerance, 0), max(target_color[2] - tolerance, 0)])
+ upper_color = np.array([min(target_color[0] + tolerance, 255), min(target_color[1] + tolerance, 255), min(target_color[2] + tolerance, 255)])
+
+ # 创建掩码
+ mask = cv2.inRange(hsv_image, lower_color, upper_color)
+
+ return cv22pil(mask).convert("L")
+
+def create_mask_from_color_tensor(image:Image, color:str, tolerance:int=0) -> Image:
+ threshold = int(tolerance * 1.28)
+ (red, green, blue) = Hex_to_RGB(color)
+ image = pil2tensor(image).squeeze()
+ temp = (torch.clamp(image, 0, 1.0) * 255.0).round().to(torch.int)
+ color_value = torch.tensor([red, green, blue])
+ lower_bound = (color_value - threshold).clamp(min=0)
+ upper_bound = (color_value + threshold).clamp(max=255)
+ lower_bound = lower_bound.view(1, 1, 1, 3)
+ upper_bound = upper_bound.view(1, 1, 1, 3)
+ mask = (temp >= lower_bound) & (temp <= upper_bound)
+ mask = mask.all(dim=-1)
+ mask = mask.float()
+ return tensor2pil(mask).convert("L")
+
+@lru_cache(maxsize=1, typed=False)
+def load_RMBG_model():
+ from .briarmbg import BriaRMBG
+ current_directory = os.path.dirname(os.path.abspath(__file__))
+ device = "cuda" if torch.cuda.is_available() else "cpu"
+ net = BriaRMBG()
+ model_path = ""
+ try:
+ model_path = os.path.join(os.path.normpath(folder_paths.folder_names_and_paths['rmbg'][0][0]), "model.pth")
+ except:
+ pass
+ if not os.path.exists(model_path):
+ model_path = os.path.join(folder_paths.models_dir, "rmbg", "RMBG-1.4", "model.pth")
+ if not os.path.exists(model_path):
+ model_path = os.path.join(os.path.dirname(current_directory), "RMBG-1.4", "model.pth")
+ net.load_state_dict(torch.load(model_path, map_location=device, weights_only=True))
+ net.to(device)
+ net.eval()
+ return net
+
+
+def RMBG(image:Image) -> Image:
+ rmbgmodel = load_RMBG_model()
+ w, h = image.size
+ im_np = np.array(image.resize((1024, 1024), Image.BILINEAR))
+ im_tensor = torch.tensor(im_np, dtype=torch.float32).permute(2, 0, 1)
+ im_tensor = torch.divide(torch.unsqueeze(im_tensor, 0), 255.0)
+ im_tensor = TF.normalize(im_tensor, [0.5, 0.5, 0.5], [1.0, 1.0, 1.0])
+ if torch.cuda.is_available():
+ im_tensor = im_tensor.cuda()
+ result = rmbgmodel(im_tensor)
+ result = torch.squeeze(F.interpolate(result[0][0], size=(h, w), mode='bilinear'), 0)
+ ma = torch.max(result)
+ mi = torch.min(result)
+ result = (result - mi) / (ma - mi)
+ im_array = (result * 255).cpu().data.numpy().astype(np.uint8)
+ _mask = torch.from_numpy(np.squeeze(im_array).astype(np.float32))
+ return tensor2pil(_mask)
+
+def guided_filter_alpha(image:torch.Tensor, mask:torch.Tensor, filter_radius:int) -> torch.Tensor:
+ sigma = 0.15
+ d = filter_radius + 1
+ mask = pil2tensor(tensor2pil(mask).convert('RGB'))
+ if not bool(d % 2):
+ d += 1
+ s = sigma / 10
+ i_dup = copy.deepcopy(image.cpu().numpy())
+ a_dup = copy.deepcopy(mask.cpu().numpy())
+ for index, image in enumerate(i_dup):
+ alpha_work = a_dup[index]
+ i_dup[index] = guidedFilter(image, alpha_work, d, s)
+ return torch.from_numpy(i_dup)
+
+#pymatting edge detail
+def mask_edge_detail(image:torch.Tensor, mask:torch.Tensor, detail_range:int=8, black_point:float=0.01, white_point:float=0.99) -> torch.Tensor:
+ from pymatting import fix_trimap, estimate_alpha_cf
+ d = detail_range * 5 + 1
+ mask = pil2tensor(tensor2pil(mask).convert('RGB'))
+ if not bool(d % 2):
+ d += 1
+ i_dup = copy.deepcopy(image.cpu().numpy().astype(np.float64))
+ a_dup = copy.deepcopy(mask.cpu().numpy().astype(np.float64))
+ for index, img in enumerate(i_dup):
+ trimap = a_dup[index][:, :, 0] # convert to single channel
+ if detail_range > 0:
+ trimap = cv2.GaussianBlur(trimap, (d, d), 0)
+ trimap = fix_trimap(trimap, black_point, white_point)
+ alpha = estimate_alpha_cf(img, trimap, laplacian_kwargs={"epsilon": 1e-6},
+ cg_kwargs={"maxiter": 500})
+ a_dup[index] = np.stack([alpha, alpha, alpha], axis=-1) # convert back to rgb
+ return torch.from_numpy(a_dup.astype(np.float32))
+
+class VITMatteModel:
+ def __init__(self,model,processor):
+ self.model = model
+ self.processor = processor
+
+def load_VITMatte_model(model_name:str, local_files_only:bool=False) -> object:
+ if local_files_only:
+ model_name = Path(os.path.join(folder_paths.models_dir, "vitmatte"))
+ # model_name = Path(os.path.join(folder_paths.models_dir, "vitmatte"))
+ from transformers import VitMatteImageProcessor, VitMatteForImageMatting
+ model = VitMatteForImageMatting.from_pretrained(model_name, local_files_only=local_files_only)
+ processor = VitMatteImageProcessor.from_pretrained(model_name, local_files_only=local_files_only)
+ vitmatte = VITMatteModel(model, processor)
+ return vitmatte
+
+def generate_VITMatte(image:Image, trimap:Image, local_files_only:bool=False, device:str="cpu", max_megapixels:float=2.0) -> Image:
+ if image.mode != 'RGB':
+ image = image.convert('RGB')
+ if trimap.mode != 'L':
+ trimap = trimap.convert('L')
+ max_megapixels *= 1048576
+ width, height = image.size
+ ratio = width / height
+ target_width = math.sqrt(ratio * max_megapixels)
+ target_height = target_width / ratio
+ target_width = int(target_width)
+ target_height = int(target_height)
+ if width * height > max_megapixels:
+ image = image.resize((target_width, target_height), Image.BILINEAR)
+ trimap = trimap.resize((target_width, target_height), Image.BILINEAR)
+ # log(f"vitmatte image size {width}x{height} too large, resize to {target_width}x{target_height} for processing.")
+ model_name = "hustvl/vitmatte-small-composition-1k"
+ if device=="cpu":
+ device = torch.device('cpu')
+ else:
+ if torch.cuda.is_available():
+ device = torch.device('cuda')
+ else:
+ log("vitmatte device is set to cuda, but not available, using cpu instead.")
+ device = torch.device('cpu')
+ vit_matte_model = load_VITMatte_model(model_name=model_name, local_files_only=local_files_only)
+ vit_matte_model.model.to(device)
+ # log(f"vitmatte processing, image size = {image.width}x{image.height}, device = {device}.")
+ inputs = vit_matte_model.processor(images=image, trimaps=trimap, return_tensors="pt")
+ with torch.no_grad():
+ inputs = {k: v.to(device) for k, v in inputs.items()}
+ predictions = vit_matte_model.model(**inputs).alphas
+ if torch.cuda.is_available():
+ torch.cuda.empty_cache()
+ torch.cuda.ipc_collect()
+ mask = tensor2pil(predictions).convert('L')
+ mask = mask.crop(
+ (0, 0, image.width, image.height)) # remove padding that the prediction appends (works in 32px tiles)
+ if width * height > max_megapixels:
+ mask = mask.resize((width, height), Image.BILINEAR)
+ return mask
+
+def generate_VITMatte_trimap(mask:torch.Tensor, erode_kernel_size:int, dilate_kernel_size:int) -> Image:
+ def g_trimap(mask, erode_kernel_size=10, dilate_kernel_size=10):
+ erode_kernel = np.ones((erode_kernel_size, erode_kernel_size), np.uint8)
+ dilate_kernel = np.ones((dilate_kernel_size, dilate_kernel_size), np.uint8)
+ eroded = cv2.erode(mask, erode_kernel, iterations=5)
+ dilated = cv2.dilate(mask, dilate_kernel, iterations=5)
+ trimap = np.zeros_like(mask)
+ trimap[dilated == 255] = 128
+ trimap[eroded == 255] = 255
+ return trimap
+
+ mask = mask.squeeze(0).cpu().detach().numpy().astype(np.uint8) * 255
+ trimap = g_trimap(mask, erode_kernel_size, dilate_kernel_size).astype(np.float32)
+ trimap[trimap == 128] = 0.5
+ trimap[trimap == 255] = 1
+ trimap = torch.from_numpy(trimap).unsqueeze(0)
+
+ return tensor2pil(trimap).convert('L')
+
+
+def get_a_person_mask_generator_model_path() -> str:
+ model_folder_name = 'mediapipe'
+ model_name = 'selfie_multiclass_256x256.tflite'
+
+ model_file_path = ""
+ try:
+ model_file_path = os.path.join(os.path.normpath(folder_paths.folder_names_and_paths[model_folder_name][0][0]), model_name)
+ except:
+ pass
+ if not os.path.exists(model_file_path):
+ model_file_path = os.path.join(folder_paths.models_dir, model_folder_name, model_name)
+
+ if not os.path.exists(model_file_path):
+ import wget
+ model_url = f'https://storage.googleapis.com/mediapipe-models/image_segmenter/selfie_multiclass_256x256/float32/latest/{model_name}'
+ log(f"Downloading '{model_name}' model")
+ os.makedirs(os.path.dirname(model_file_path), exist_ok=True)
+ wget.download(model_url, model_file_path)
+ return model_file_path
+
+def mask_fix(images:torch.Tensor, radius:int, fill_holes:int, white_threshold:float, extra_clip:float) -> torch.Tensor:
+ d = radius * 2 + 1
+ i_dup = copy.deepcopy(images.cpu().numpy())
+ for index, image in enumerate(i_dup):
+ cleaned = cv2.bilateralFilter(image, 9, 0.05, 8)
+ alpha = np.clip((image - white_threshold) / (1 - white_threshold), 0, 1)
+ rgb = image * alpha
+ alpha = cv2.GaussianBlur(alpha, (d, d), 0) * 0.99 + np.average(alpha) * 0.01
+ rgb = cv2.GaussianBlur(rgb, (d, d), 0) * 0.99 + np.average(rgb) * 0.01
+ rgb = rgb / np.clip(alpha, 0.00001, 1)
+ rgb = rgb * extra_clip
+ cleaned = np.clip(cleaned / rgb, 0, 1)
+ if fill_holes > 0:
+ fD = fill_holes * 2 + 1
+ gamma = cleaned * cleaned
+ kD = np.ones((fD, fD), np.uint8)
+ kE = np.ones((fD + 2, fD + 2), np.uint8)
+ gamma = cv2.dilate(gamma, kD, iterations=1)
+ gamma = cv2.erode(gamma, kE, iterations=1)
+ gamma = cv2.GaussianBlur(gamma, (fD, fD), 0)
+ cleaned = np.maximum(cleaned, gamma)
+ i_dup[index] = cleaned
+ return torch.from_numpy(i_dup)
+
+def histogram_remap(image:torch.Tensor, blackpoint:float, whitepoint:float) -> torch.Tensor:
+ bp = min(blackpoint, whitepoint - 0.001)
+ scale = 1 / (whitepoint - bp)
+ i_dup = copy.deepcopy(image.cpu().numpy())
+ i_dup = np.clip((i_dup - bp) * scale, 0.0, 1.0)
+ return torch.from_numpy(i_dup)
+
+def expand_mask(mask:torch.Tensor, grow:int, blur:int) -> torch.Tensor:
+ # grow
+ c = 0
+ kernel = np.array([[c, 1, c],
+ [1, 1, 1],
+ [c, 1, c]])
+ growmask = mask.reshape((-1, mask.shape[-2], mask.shape[-1]))
+ out = []
+ for m in growmask:
+ output = m.numpy()
+ for _ in range(abs(grow)):
+ if grow < 0:
+ output = scipy.ndimage.grey_erosion(output, footprint=kernel)
+ else:
+ output = scipy.ndimage.grey_dilation(output, footprint=kernel)
+ output = torch.from_numpy(output)
+ out.append(output)
+ # blur
+ for idx, tensor in enumerate(out):
+ pil_image = tensor2pil(tensor.cpu().detach())
+ pil_image = pil_image.filter(ImageFilter.GaussianBlur(blur))
+ out[idx] = pil2tensor(pil_image)
+ ret_mask = torch.cat(out, dim=0)
+ return ret_mask
+
+def mask_invert(mask:torch.Tensor) -> torch.Tensor:
+ return 1 - mask
+
+def subtract_mask(masks_a:torch.Tensor, masks_b:torch.Tensor) -> torch.Tensor:
+ return torch.clamp(masks_a - masks_b, 0, 255)
+
+def add_mask(masks_a:torch.Tensor, masks_b:torch.Tensor) -> torch.Tensor:
+ mask = chop_image(tensor2pil(masks_a), tensor2pil(masks_b), blend_mode='add', opacity=100)
+ return image2mask(mask)
+
+def RGB2RGBA(image:Image, mask:Image) -> Image:
+ (R, G, B) = image.convert('RGB').split()
+ return Image.merge('RGBA', (R, G, B, mask.convert('L')))
+
+def mask_area(image:Image) -> tuple:
+ cv2_image = pil2cv2(image.convert('RGBA'))
+ gray = cv2.cvtColor(cv2_image, cv2.COLOR_BGR2GRAY)
+ _, thresh = cv2.threshold(gray, 127, 255, 0)
+ locs = np.where(thresh == 255)
+ x1 = np.min(locs[1]) if len(locs[1]) > 0 else 0
+ x2 = np.max(locs[1]) if len(locs[1]) > 0 else image.width
+ y1 = np.min(locs[0]) if len(locs[0]) > 0 else 0
+ y2 = np.max(locs[0]) if len(locs[0]) > 0 else image.height
+ x1, y1, x2, y2 = min(x1, x2), min(y1, y2), max(x1, x2), max(y1, y2)
+ return (x1, y1, x2 - x1, y2 - y1)
+
+def min_bounding_rect(image:Image) -> tuple:
+ cv2_image = pil2cv2(image)
+ gray = cv2.cvtColor(cv2_image, cv2.COLOR_BGR2GRAY)
+ ret, thresh = cv2.threshold(gray, 127, 255, 0)
+ contours, _ = cv2.findContours(thresh, 1, 2)
+ x, y, width, height = 0, 0, 0, 0
+ area = 0
+ for contour in contours:
+ _x, _y, _w, _h = cv2.boundingRect(contour)
+ _area = _w * _h
+ if _area > area:
+ area = _area
+ x, y, width, height = _x, _y, _w, _h
+ return (x, y, width, height)
+
+def max_inscribed_rect(image:Image) -> tuple:
+ img = pil2cv2(image)
+ img_gray = cv2.cvtColor(img, cv2.COLOR_BGR2GRAY)
+ ret, img_bin = cv2.threshold(img_gray, 127, 255, cv2.THRESH_BINARY)
+ contours, _ = cv2.findContours(img_bin, cv2.RETR_CCOMP, cv2.CHAIN_APPROX_SIMPLE)
+ contour = contours[0].reshape(len(contours[0]), 2)
+ rect = []
+ for i in range(len(contour)):
+ x1, y1 = contour[i]
+ for j in range(len(contour)):
+ x2, y2 = contour[j]
+ area = abs(y2 - y1) * abs(x2 - x1)
+ rect.append(((x1, y1), (x2, y2), area))
+ all_rect = sorted(rect, key=lambda x: x[2], reverse=True)
+ if all_rect:
+ best_rect_found = False
+ index_rect = 0
+ nb_rect = len(all_rect)
+ while not best_rect_found and index_rect < nb_rect:
+ rect = all_rect[index_rect]
+ (x1, y1) = rect[0]
+ (x2, y2) = rect[1]
+ valid_rect = True
+ x = min(x1, x2)
+ while x < max(x1, x2) + 1 and valid_rect:
+ if any(img[y1, x]) == 0 or any(img[y2, x]) == 0:
+ valid_rect = False
+ x += 1
+ y = min(y1, y2)
+ while y < max(y1, y2) + 1 and valid_rect:
+ if any(img[y, x1]) == 0 or any(img[y, x2]) == 0:
+ valid_rect = False
+ y += 1
+ if valid_rect:
+ best_rect_found = True
+ index_rect += 1
+ #较小的数值排前面
+ x1, y1, x2, y2 = min(x1, x2), min(y1, y2), max(x1, x2), max(y1, y2)
+ return (x1, y1, x2 - x1, y2 - y1)
+
+def gray_threshold(image:Image, thresh:int=127, otsu:bool=False) -> Image:
+ cv2_image = pil2cv2(image)
+ gray = cv2.cvtColor(cv2_image, cv2.COLOR_BGR2GRAY)
+ if otsu:
+ _, thresh = cv2.threshold(gray,0,255,cv2.THRESH_BINARY+cv2.THRESH_OTSU)
+ else:
+ _, thresh = cv2.threshold(gray, thresh, 255, cv2.THRESH_TOZERO)
+ return cv22pil(thresh).convert('L')
+
+def image_to_colormap(image:Image, index:int) -> Image:
+ return cv22pil(cv2.applyColorMap(pil2cv2(image), index))
+
+# 检查mask有效区域面积比例
+def mask_white_area(mask:Image, white_point:int) -> float:
+ if mask.mode != 'L':
+ mask.convert('L')
+ white_pixels = 0
+ for y in range(mask.height):
+ for x in range(mask.width):
+ mask.getpixel((x, y)) > 16
+ if mask.getpixel((x, y)) > white_point:
+ white_pixels += 1
+ return white_pixels / (mask.width * mask.height)
+
+'''Color Functions'''
+
+def color_balance(image:Image, shadows:list, midtones:list, highlights:list,
+ shadow_center:float=0.15, midtone_center:float=0.5, highlight_center:float=0.8,
+ shadow_max:float=0.1, midtone_max:float=0.3, highlight_max:float=0.2,
+ preserve_luminosity:bool=False) -> Image:
+
+ img = pil2tensor(image)
+ # Create a copy of the img tensor
+ img_copy = img.clone()
+
+ # Calculate the original luminance if preserve_luminosity is True
+ if preserve_luminosity:
+ original_luminance = 0.2126 * img_copy[..., 0] + 0.7152 * img_copy[..., 1] + 0.0722 * img_copy[..., 2]
+
+ # Define the adjustment curves
+ def adjust(x, center, value, max_adjustment):
+ # Scale the adjustment value
+ value = value * max_adjustment
+
+ # Define control points
+ points = torch.tensor([[0, 0], [center, center + value], [1, 1]])
+
+ # Create cubic spline
+ from scipy.interpolate import CubicSpline
+ cs = CubicSpline(points[:, 0], points[:, 1])
+
+ # Apply the cubic spline to the color channel
+ return torch.clamp(torch.from_numpy(cs(x)), 0, 1)
+
+ # Apply the adjustments to each color channel
+ # shadows, midtones, highlights are lists of length 3 (for R, G, B channels) with values between -1 and 1
+ for i, (s, m, h) in enumerate(zip(shadows, midtones, highlights)):
+ img_copy[..., i] = adjust(img_copy[..., i], shadow_center, s, shadow_max)
+ img_copy[..., i] = adjust(img_copy[..., i], midtone_center, m, midtone_max)
+ img_copy[..., i] = adjust(img_copy[..., i], highlight_center, h, highlight_max)
+
+ # If preserve_luminosity is True, adjust the RGB values to match the original luminance
+ if preserve_luminosity:
+ current_luminance = 0.2126 * img_copy[..., 0] + 0.7152 * img_copy[..., 1] + 0.0722 * img_copy[..., 2]
+ img_copy *= (original_luminance / current_luminance).unsqueeze(-1)
+
+ return tensor2pil(img_copy)
+
+def RGB_to_Hex(RGB:tuple) -> str:
+ color = '#'
+ for i in RGB:
+ num = int(i)
+ color += str(hex(num))[-2:].replace('x', '0').upper()
+ return color
+
+def Hex_to_RGB(inhex:str) -> tuple:
+ if not inhex.startswith('#'):
+ raise ValueError(f'Invalid Hex Code in {inhex}')
+ else:
+ rval = inhex[1:3]
+ gval = inhex[3:5]
+ bval = inhex[5:]
+ rgb = (int(rval, 16), int(gval, 16), int(bval, 16))
+ return tuple(rgb)
+
+def RGB_to_HSV(RGB:tuple) -> list:
+ HSV = rgb_to_hsv(RGB[0] / 255.0, RGB[1] / 255.0, RGB[2] / 255.0)
+ return [int(x * 360) for x in HSV]
+
+def Hex_to_HSV_255level(inhex:str) -> list:
+ if not inhex.startswith('#'):
+ raise ValueError(f'Invalid Hex Code in {inhex}')
+ else:
+ rval = inhex[1:3]
+ gval = inhex[3:5]
+ bval = inhex[5:]
+ RGB = (int(rval, 16), int(gval, 16), int(bval, 16))
+ HSV = rgb_to_hsv(RGB[0] / 255.0, RGB[1] / 255.0, RGB[2] / 255.0)
+ return [int(x * 255) for x in HSV]
+
+def HSV_255level_to_Hex(HSV: list) -> str:
+ if len(HSV) != 3 or any((not isinstance(v, int) or v < 0 or v > 255) for v in HSV):
+ raise ValueError('Invalid HSV values, each value should be an integer between 0 and 255')
+
+ H, S, V = HSV
+ RGB = tuple(int(x * 255) for x in hsv_to_rgb(H / 255.0, S / 255.0, V / 255.0))
+
+ # Convert RGB values to hexadecimal format
+ hex_r = format(RGB[0], '02x')
+ hex_g = format(RGB[1], '02x')
+ hex_b = format(RGB[2], '02x')
+
+ return '#' + hex_r + hex_g + hex_b
+
+# 返回补色色值
+def complementary_color(color: str) -> str:
+ color = Hex_to_RGB(color)
+ return RGB_to_Hex((255 - color[0], 255 - color[1], 255 - color[2]))
+
+# 返回颜色对应灰度值
+def rgb2gray(color:str)->int:
+ (r, g, b) = Hex_to_RGB(color)
+ return int((r * 299 + g * 587 + b * 114) / 1000)
+
+'''Value Functions'''
+def is_valid_mask(tensor:torch.Tensor) -> bool:
+ return not bool(torch.all(tensor == 0).item())
+
+def step_value(start_value, end_value, total_step, step) -> float: # 按当前步数在总步数中的位置返回比例值
+ factor = step / total_step
+ return (end_value - start_value) * factor + start_value
+
+def step_color(start_color_inhex:str, end_color_inhex:str, total_step:int, step:int) -> str: # 按当前步数在总步数中的位置返回比例颜色
+ start_color = tuple(Hex_to_RGB(start_color_inhex))
+ end_color = tuple(Hex_to_RGB(end_color_inhex))
+ start_R, start_G, start_B = start_color[0], start_color[1], start_color[2]
+ end_R, end_G, end_B = end_color[0], end_color[1], end_color[2]
+ ret_color = (int(step_value(start_R, end_R, total_step, step)),
+ int(step_value(start_G, end_G, total_step, step)),
+ int(step_value(start_B, end_B, total_step, step)),
+ )
+ return RGB_to_Hex(ret_color)
+
+def has_letters(string:str) -> bool:
+ pattern = r'[a-zA-Z]'
+ match = re.search(pattern, string)
+ if match:
+ return True
+ else:
+ return False
+
+
+def replace_case(old:str, new:str, text:str) -> str:
+ index = text.lower().find(old.lower())
+ if index == -1:
+ return text
+ return replace_case(old, new, text[:index] + new + text[index + len(old):])
+
+def random_numbers(total:int, random_range:int, seed:int=0, sum_of_numbers:int=0) -> list:
+ random.seed(seed)
+ numbers = [random.randint(-random_range//2, random_range//2) for _ in range(total - 1)]
+ avg = sum(numbers) // total
+ ret_list = []
+ for i in numbers:
+ ret_list.append(i - avg)
+ ret_list.append((sum_of_numbers - sum(ret_list)) // 2)
+ return ret_list
+
+# 四舍五入取整数倍
+def num_round_to_multiple(number:int, multiple:int) -> int:
+ remainder = number % multiple
+ if remainder == 0 :
+ return number
+ else:
+ factor = int(number / multiple)
+ if number - factor * multiple > multiple / 2:
+ factor += 1
+ return factor * multiple
+
+# 向上取整数倍
+def num_round_up_to_multiple(number: int, multiple: int) -> int:
+ remainder = number % multiple
+ if remainder == 0:
+ return number
+ else:
+ factor = (number + multiple - 1) // multiple # 向上取整的计算方式
+ return factor * multiple
+
+def calculate_side_by_ratio(orig_width:int, orig_height:int, ratio:float, longest_side:int=0) -> int:
+
+ if orig_width > orig_height:
+ if longest_side:
+ target_width = longest_side
+ else:
+ target_width = orig_width
+ target_height = int(target_width / ratio)
+ else:
+ if longest_side:
+ target_height = longest_side
+ else:
+ target_height = orig_height
+ target_width = int(target_height * ratio)
+
+ if ratio < 1:
+ if longest_side:
+ _r = longest_side / target_height
+ target_height = longest_side
+ else:
+ _r = orig_height / target_height
+ target_height = orig_height
+ target_width = int(target_width * _r)
+
+ return target_width, target_height
+
+def generate_random_name(prefix:str, suffix:str, length:int) -> str:
+ name = ''.join(random.choice("abcdefghijklmnopqrstupvxyz1234567890") for x in range(length))
+ return prefix + name + suffix
+
+def check_image_file(file_name:str, interval:int) -> object:
+ while True:
+ if os.path.isfile(file_name):
+ try:
+ image = Image.open(file_name)
+ ret_image = copy.deepcopy(image)
+ image.close()
+ return ret_image
+ except Exception as e:
+ log(e)
+ return None
+ break
+ time.sleep(interval / 1000)
+
+# 判断字符串是否包含中文
+def is_contain_chinese(check_str:str) -> bool:
+ for ch in check_str:
+ if u'\u4e00' <= ch <= u'\u9fff':
+ return True
+ return False
+
+# 生成随机颜色
+def generate_random_color():
+ """
+ Generate a random color in hexadecimal format.
+ """
+ # random.seed(int(time.time()))
+ return "#{:06x}".format(random.randint(0x101010, 0xFFFFFF))
+
+# 提取字符串中的int数为列表
+def extract_numbers(string):
+ return [int(s) for s in re.findall(r'\d+', string)]
+
+# 提取字符串中的数值, 返回为列表
+def extract_all_numbers_from_str(string, checkint:bool=False):
+ # 定义浮点数的正则表达式模式
+ number_pattern = r'[-+]?\d*\.?\d+(?:[eE][-+]?\d+)?'
+ # 使用re.findall找到所有匹配的字符串
+ matches = re.findall(number_pattern, string)
+ # 转换为浮点数
+ numbers = [float(match) for match in matches]
+ number_list = []
+ # 如果需要检查是否为整数,则将浮点数转换为整数
+ if checkint:
+ for num in numbers:
+ int_num = int(num)
+ if math.isclose(num, int_num, rel_tol=1e-19):
+ number_list.append(int_num)
+ else:
+ number_list.append(num)
+ else:
+ number_list = numbers
+
+ return number_list
+
+
+
+# 提取字符串中用"," ";" " "分开的字符串, 返回为列表
+def extract_substr_from_str(string) -> list:
+ return re.split(r'[,\s;,;]+', string)
+
+def clear_memory():
+ import gc
+ # Cleanup
+ gc.collect()
+ if torch.cuda.is_available():
+ torch.cuda.empty_cache()
+ torch.cuda.ipc_collect()
+
+def tensor_info(tensor:object) -> str:
+ value = ''
+ if isinstance(tensor, torch.Tensor):
+ value += f"\n Input dim = {tensor.dim()}, shape[0] = {tensor.shape[0]} \n"
+ for i in range(tensor.shape[0]):
+ t = tensor[i]
+ image = tensor2pil(t)
+ value += f'\n index {i}: Image.size = {image.size}, Image.mode = {image.mode}, dim = {t.dim()}, '
+ for j in range(t.dim()):
+ value += f'shape[{j}] = {t.shape[j]}, '
+ else:
+ value = f"tensor_info: Not tensor, type is {type(tensor)}"
+ return value
+
+# 去除空行
+def remove_empty_lines(text):
+ lines = text.split('\n')
+ non_empty_lines = [line for line in lines if line.strip() != '']
+ return '\n'.join(non_empty_lines)
+
+# 去除重复的句子
+def remove_duplicate_string(text:str) -> str:
+ sentences = re.split(r'(?<=[:;,.!?])\s+', text)
+ unique_sentences = []
+ seen = set()
+ for sentence in sentences:
+ if sentence not in seen:
+ seen.add(sentence)
+ unique_sentences.append(sentence)
+ return ' '.join(unique_sentences)
+
+files_for_uform_gen2_qwen = Path(os.path.join(folder_paths.models_dir, "LLavacheckpoints", "files_for_uform_gen2_qwen"))
+class StopOnTokens(StoppingCriteria):
+ def __call__(self, input_ids: torch.LongTensor, scores: torch.FloatTensor, **kwargs) -> bool:
+ stop_ids = [151645] # Define stop tokens as per your model's specifics
+ for stop_id in stop_ids:
+ if input_ids[0][-1] == stop_id:
+ return True
+ return False
+
+class UformGen2QwenChat:
+
+ def __init__(self):
+ from huggingface_hub import snapshot_download
+ # self.model_path = snapshot_download("unum-cloud/uform-gen2-qwen-500m",
+ # local_dir=files_for_uform_gen2_qwen,
+ # force_download=False, # Set to True if you always want to download, regardless of local copy
+ # local_files_only=False, # Set to False to allow downloading if not available locally
+ # local_dir_use_symlinks="auto") # or set to True/False based on your symlink preference
+ self.model_path = files_for_uform_gen2_qwen
+ self.device = "cuda" if torch.cuda.is_available() else "cpu"
+ self.model = AutoModel.from_pretrained(self.model_path, trust_remote_code=True).to(self.device)
+ self.processor = AutoProcessor.from_pretrained(self.model_path, trust_remote_code=True)
+
+ def chat_response(self, message, history, image_path):
+ stop = StopOnTokens()
+ messages = [{"role": "system", "content": "You are a helpful Assistant."}]
+
+ for user_msg, assistant_msg in history:
+ messages.append({"role": "user", "content": user_msg})
+ messages.append({"role": "assistant", "content": assistant_msg})
+
+ if len(messages) == 1:
+ message = f" {message}"
+
+ messages.append({"role": "user", "content": message})
+
+ model_inputs = self.processor.tokenizer.apply_chat_template(
+ messages,
+ add_generation_prompt=True,
+ return_tensors="pt"
+ )
+
+ image = Image.open(image_path) # Load image using PIL
+ image_tensor = (
+ self.processor.feature_extractor(image)
+ .unsqueeze(0)
+ )
+
+ attention_mask = torch.ones(
+ 1, model_inputs.shape[1] + self.processor.num_image_latents - 1
+ )
+
+ model_inputs = {
+ "input_ids": model_inputs,
+ "images": image_tensor,
+ "attention_mask": attention_mask
+ }
+
+ model_inputs = {k: v.to(self.device) for k, v in model_inputs.items()}
+
+ with torch.inference_mode():
+ output = self.model.generate(
+ **model_inputs,
+ max_new_tokens=512,
+ do_sample=True,
+ temperature=0.3,
+ repetition_penalty=1.2,
+ stopping_criteria=StoppingCriteriaList([stop])
+ )
+
+ response_text = self.processor.tokenizer.decode(output[0], skip_special_tokens=True)
+ response_text = remove_duplicate_string(response_text)
+ return response_text
+
+'''CLASS'''
+
+class AnyType(str):
+ """A special class that is always equal in not equal comparisons. Credit to pythongosssss"""
+ def __eq__(self, __value: object) -> bool:
+ return True
+ def __ne__(self, __value: object) -> bool:
+ return False
+
+
+
+'''Load File'''
+
+def download_hg_model(model_id:str,exDir:str='') -> str:
+ # 下载本地
+ model_checkpoint = os.path.join(folder_paths.models_dir, exDir, os.path.basename(model_id))
+ if not os.path.exists(model_checkpoint):
+ from huggingface_hub import snapshot_download
+ snapshot_download(repo_id=model_id, local_dir=model_checkpoint, local_dir_use_symlinks=False)
+ return model_checkpoint
+
+
+def get_files(model_path: str, file_ext_list:list) -> dict:
+ file_list = []
+ for ext in file_ext_list:
+ file_list.extend(glob.glob(os.path.join(model_path, '*' + ext)))
+ files_dict = {}
+ for i in range(len(file_list)):
+ _, filename = os.path.split(file_list[i])
+ files_dict[filename] = file_list[i]
+ return files_dict
+
+# def load_inference_prompt() -> str:
+# inference_prompt_file = os.path.join(os.path.dirname(os.path.dirname(os.path.normpath(__file__))), "resource",
+# "inference.prompt")
+# ret_value = ''
+# try:
+# with open(inference_prompt_file, 'r') as f:
+# ret_value = f.readlines()
+# except Exception as e:
+# log(f'Warning: {inference_prompt_file} ' + repr(e) + f", check it to be correct. ", message_type='warning')
+# return ''.join(ret_value)
+
+def load_custom_size() -> list:
+ custom_size_file = os.path.join(os.path.dirname(os.path.dirname(os.path.normpath(__file__))), "custom_size.ini")
+ ret_value = ['1024 x 1024',
+ '768 x 512',
+ '512 x 768',
+ '1280 x 720',
+ '720 x 1280',
+ '1344 x 768',
+ '768 x 1344',
+ '1536 x 640',
+ '640 x 1536'
+ ]
+ try:
+ with open(custom_size_file, 'r') as f:
+ ini = f.readlines()
+ for line in ini:
+ if not line.startswith(f'#'):
+ ret_value.append(line.strip())
+ except Exception as e:
+ pass
+ # log(f'Warning: {custom_size_file} not found' + f", use default size. ")
+ return ret_value
+
+def get_api_key(api_name:str) -> str:
+ api_key_ini_file = os.path.join(os.path.dirname(os.path.dirname(os.path.normpath(__file__))), "api_key.ini")
+ ret_value = ''
+ try:
+ with open(api_key_ini_file, 'r') as f:
+ ini = f.readlines()
+ for line in ini:
+ if line.startswith(f'{api_name}='):
+ ret_value = line[line.find('=') + 1:].rstrip().lstrip()
+ break
+ except Exception as e:
+ log(f'Warning: {api_key_ini_file} ' + repr(e) + f", check it to be correct. ", message_type='warning')
+ remove_char = ['"', "'", '“', '”', '‘', '’']
+ for i in remove_char:
+ if i in ret_value:
+ ret_value = ret_value.replace(i, '')
+ if len(ret_value) < 4:
+ log(f'Warning: Invalid API-key, Check the key in {api_key_ini_file}.', message_type='warning')
+ return ret_value
+
+# 判断文件名后缀是否包括在列表中(忽略大小写)
+def file_is_extension(filename:str, ext_list:tuple) -> bool:
+ # 获取文件的真实后缀(包括点)
+ true_ext = os.path.splitext(filename)[1]
+ if true_ext.lower() in ext_list:
+ return True
+ return False
+
+# 遍历目录下包括子目录指定后缀文件,返回字典
+def collect_files(root_dir:str, suffixes:tuple, default_dir:str=""):
+ result = {}
+ for dirpath, _, filenames in os.walk(root_dir):
+ for file in filenames:
+ if file_is_extension(file, suffixes):
+ # 获取文件的完整路径作为 value
+ full_path = os.path.join(dirpath, file)
+ # 如果是default_dir 则去掉路径,使用文件名作为 key
+ if dirpath == default_dir:
+ relative_path = os.path.relpath(full_path, root_dir)
+ result.update({relative_path: full_path})
+ else:
+ result.update({full_path: full_path})
+ return result
+
+
+def get_resource_dir() -> list:
+ default_lut_dir = []
+ default_lut_dir.append(os.path.join(os.path.dirname(os.path.dirname(os.path.normpath(__file__))), 'lut'))
+ default_font_dir = []
+ default_font_dir.append(os.path.join(os.path.dirname(os.path.dirname(os.path.normpath(__file__))), 'font'))
+ resource_dir_ini_file = os.path.join(os.path.dirname(os.path.dirname(os.path.normpath(__file__))),
+ "resource_dir.ini")
+ try:
+ with open(resource_dir_ini_file, 'r') as f:
+ ini = f.readlines()
+ for line in ini:
+ if line.startswith('LUT_dir='):
+ _ldir = line[line.find('=') + 1:].rstrip().lstrip()
+ for dir in extract_substr_from_str(_ldir) :
+ if os.path.exists(dir):
+ default_lut_dir.append(dir)
+ elif line.startswith('FONT_dir='):
+ _fdir = line[line.find('=') + 1:].rstrip().lstrip()
+ for dir in extract_substr_from_str(_fdir):
+ if os.path.exists(dir):
+ default_font_dir.append(dir)
+ except Exception as e:
+ pass
+ # log(f'Warning: {resource_dir_ini_file} not found' + f", default directory to be used. ")
+
+
+ LUT_DICT = {}
+ for dir in default_lut_dir:
+ LUT_DICT.update(collect_files(root_dir=dir, suffixes= ('.cube'), default_dir=default_lut_dir[0] )) # 后缀要小写
+ LUT_LIST = list(LUT_DICT.keys())
+
+ FONT_DICT = {}
+ for dir in default_font_dir:
+ FONT_DICT.update(collect_files(root_dir=dir, suffixes=('.ttf', '.otf'), default_dir=default_font_dir[0])) # 后缀要小写
+ FONT_LIST = list(FONT_DICT.keys())
+
+ return (LUT_DICT, FONT_DICT)
+
+# (LUT_DICT, FONT_DICT) = get_resource_dir()
+# FONT_LIST = list(FONT_DICT.keys())
+# LUT_LIST = list(LUT_DICT.keys())
+
+# def get_models_dir() -> dict:
+# models_dir_ini_file = os.path.join(os.path.dirname(os.path.dirname(os.path.normpath(__file__))), "models_dir.ini")
+# MODELS_DIR = {}
+# model_dir_list = [
+# "birefnet_dir",
+# "evf-sam_dir",
+# "florence2_dir",
+# "lama_dir",
+# "rmbg_dir",
+# "segformerB2_dir",
+# "segformerB3_clothes_dir",
+# "segformerB3_fashion_dir",
+# "sam2_dir",
+# "transparent-background_dir",
+# "yolo8_dir",
+# "yolo_world_dir"
+# ]
+# try:
+# with open(models_dir_ini_file, 'r') as f:
+# ini = f.readlines()
+# for line in ini:
+# for model_dir in model_dir_list:
+# if line.startswith(model_dir):
+# path = line[line.find('=') + 1:].rstrip().lstrip()
+# if os.path.exists(path):
+# MODELS_DIR[model_dir] = path
+# log(f'Find {len(MODELS_DIR)} path(s) in {models_dir_ini_file}.')
+# except Exception as e:
+# log(f'Warning: {models_dir_ini_file} not found' + f', default directory to be used.')
+#
+# return MODELS_DIR
+#
+# MODELS_DIR = get_models_dir()
+
+def draw_bounding_boxes(image: Image, bboxes: list, color: str = "#FF0000", line_width: int = 5) -> Image:
+ """
+ Draw bounding boxes on the image using the coordinates provided in the bboxes dictionary.
+ """
+
+ (_, FONT_DICT) = get_resource_dir()
+
+ font_size = 25
+ font = ImageFont.truetype(list(FONT_DICT.items())[0][1], font_size)
+
+ if len(bboxes) > 0:
+ draw = ImageDraw.Draw(image)
+ width, height = image.size
+ if line_width < 0: # auto line width
+ line_width = (image.width + image.height) // 1000
+
+ for index, box in enumerate(bboxes):
+ random_color = generate_random_color()
+ if color != "random":
+ random_color = color
+ xmin = min(box[0], box[2])
+ xmax = max(box[0], box[2])
+ ymin = min(box[1], box[3])
+ ymax = max(box[1], box[3])
+ draw.rectangle([xmin, ymin, xmax, ymax], outline=random_color, width=line_width)
+ draw.text((xmin, ymin - font_size*1.2), str(index), font=font, fill=random_color)
+
+ return image
+
+def draw_bbox(image: Image, bbox: tuple, color: str = "#FF0000", line_width: int = 5, title: str = "", font_size: int = 10) -> Image:
+ """
+ Draw bounding boxes on the image using the coordinates provided in the bboxes dictionary.
+ """
+
+ (_, FONT_DICT) = get_resource_dir()
+
+ font = ImageFont.truetype(list(FONT_DICT.items())[0][1], font_size)
+
+ draw = ImageDraw.Draw(image)
+ width, height = image.size
+ if line_width < 0: # auto line width
+ line_width = (image.width + image.height) // 1000
+
+ random_color = generate_random_color()
+ if color != "random":
+ random_color = color
+ xmin = min(bbox[0], bbox[2])
+ xmax = max(bbox[0], bbox[2])
+ ymin = min(bbox[1], bbox[3])
+ ymax = max(bbox[1], bbox[3])
+ draw.rectangle([xmin, ymin, xmax, ymax], outline=random_color, width=line_width)
+ if title != "":
+ draw.text((xmin, ymin - font_size*1.2), title, font=font, fill=random_color)
+
+ return image
+
+
+
+'''Constant'''
+
+chop_mode = [
+ 'normal',
+ 'multply',
+ 'screen',
+ 'add',
+ 'subtract',
+ 'difference',
+ 'darker',
+ 'lighter',
+ 'color_burn',
+ 'color_dodge',
+ 'linear_burn',
+ 'linear_dodge',
+ 'overlay',
+ 'soft_light',
+ 'hard_light',
+ 'vivid_light',
+ 'pin_light',
+ 'linear_light',
+ 'hard_mix'
+ ]
+
+# Blend Mode from Virtuoso Pack https://github.com/chrisfreilich/virtuoso-nodes
+chop_mode_v2 = list(BLEND_MODES.keys())
+
+gemini_generate_config = {
+ "temperature": 0,
+ "top_p": 1,
+ "top_k": 1,
+ "max_output_tokens": 400
+}
+
+gemini_safety_settings = [
+ {
+ "category": "HARM_CATEGORY_HARASSMENT",
+ "threshold": "BLOCK_NONE"
+ },
+ {
+ "category": "HARM_CATEGORY_HATE_SPEECH",
+ "threshold": "BLOCK_NONE"
+ },
+ {
+ "category": "HARM_CATEGORY_SEXUALLY_EXPLICIT",
+ "threshold": "BLOCK_NONE"
+ },
+ {
+ "category": "HARM_CATEGORY_DANGEROUS_CONTENT",
+ "threshold": "BLOCK_NONE"
+ }
+]
+
+minicpm_llama3_v25_prompts = """
+ # MISSION
+ You are an imagine generator for a slide deck tool. You will be given the text or description of a slide and you'll generate a few image descriptions that will be fed to an AI image generator. It will need to have a particular format (seen below). You will also be given some examples below. Think metaphorically and symbolically.
+
+ # FORMAT
+ The format should follow this general pattern:
+
+ , , ,