fix florence

This commit is contained in:
aiXander
2024-08-14 15:21:41 -07:00
parent 2434d4846d
commit 682edd9333
5 changed files with 18 additions and 13 deletions
+3 -3
View File
@@ -3,10 +3,10 @@ GPU_ID="device=3"
cog predict --gpus $GPU_ID \
-i name="xander_sdxl_cog" \
-i lora_training_urls="https://storage.googleapis.com/public-assets-xander/A_workbox/lora_training_sets/xander_big.zip" \
-i lora_training_urls="https://storage.googleapis.com/public-assets-xander/A_workbox/lora_training_sets/xander.zip" \
-i concept_mode="face" \
-i sd_model_version="sdxl" \
-i max_train_steps="360" \
-i caption_model="blip" \
-i max_train_steps="300" \
-i caption_model="florence" \
-i debug="False" \
-i seed="0"
+7 -3
View File
@@ -175,9 +175,13 @@ class Predictor(BasePredictor):
# Add instructions README:
tar.add("instructions_README.md", arcname="README.md")
tar.add("comfyUI_workflow_lora_txt2img.json", arcname="comfyUI_workflow_lora_txt2img.json")
if sd_model_version == "sd15":
tar.add("comfyUI_workflow_lora_adiff.json", arcname="comfyUI_workflow_lora_adiff.json")
comfy_workflows_path = "ComfyUI_workflows"
if os.path.exists(comfy_workflows_path) and os.path.isdir(comfy_workflows_path):
for root, dirs, files in os.walk(comfy_workflows_path):
for file in files:
file_path = os.path.join(root, file)
arcname = os.path.relpath(file_path, os.path.dirname(comfy_workflows_path))
tar.add(file_path, arcname=arcname)
attributes = {}
attributes['grid_prompts'] = config.training_attributes["validation_prompts"]
+2 -1
View File
@@ -21,4 +21,5 @@ ujson==5.10.0
bitsandbytes==0.43.1
setuptools==70.3.0
torchtyping==0.1.5
einops
einops
timm
+2 -2
View File
@@ -1,5 +1,5 @@
{
"name": "clipx",
"name": "twisting_realities",
"sd_model_version": "sdxl",
"lora_training_urls": "https://edenartlab-lfs.s3.amazonaws.com/datasets/twisting_realities.zip",
"concept_mode": "style",
@@ -7,7 +7,7 @@
"seed": 0,
"resolution": 512,
"train_batch_size": 4,
"n_sample_imgs": 8,
"n_sample_imgs": 6,
"max_train_steps": 300,
"checkpointing_steps": 200,
+4 -4
View File
@@ -562,7 +562,7 @@ def florence_caption_dataset(images, captions):
with patch("transformers.dynamic_module_utils.get_imports", fixed_get_imports): #workaround for unnecessary flash_attn requirement
model = AutoModelForCausalLM.from_pretrained("microsoft/Florence-2-large", attn_implementation="sdpa", device_map=device, torch_dtype=torch_dtype,trust_remote_code=True)
processor = AutoProcessor.from_pretrained("microsoft/Florence-2-large", trust_remote_code=True)
processor = AutoProcessor.from_pretrained("microsoft/Florence-2-large", trust_remote_code=True, cache_dir = model_paths.get_path("FLORENCE"))
for i, image in enumerate(tqdm(images)):
if captions[i] is None:
@@ -675,15 +675,15 @@ def random_crop(image, scale=(0.85, 0.95)):
return image.crop((left, top, left + new_width, top + new_height))
def gaussian_blur(image):
return image.filter(ImageFilter.GaussianBlur(radius=1))
def gaussian_blur(image, radius = 1.0):
return image.filter(ImageFilter.GaussianBlur(radius=radius))
def augment_image(image):
image = hue_augmentation(image)
image = color_jitter(image)
image = random_crop(image)
if random.random() < 0.5:
image = gaussian_blur(image, blur = random.uniform(0.0, 1.0))
image = gaussian_blur(image, radius = random.uniform(0.0, 1.0))
return image
def round_to_nearest_multiple(x, multiple):