diff --git a/__init__.py b/__init__.py index 5f3c625..dd120d9 100644 --- a/__init__.py +++ b/__init__.py @@ -8,4 +8,4 @@ from .nodes import StableAudioSampler, NODE_CLASS_MAPPINGS, NODE_DISPLAY_NAME_MAPPINGS __all__ = ['NODE_CLASS_MAPPINGS', 'NODE_DISPLAY_NAME_MAPPINGS'] - +WEB_DIRECTORY = "./web" \ No newline at end of file diff --git a/nodes.py b/nodes.py index 3eb2ace..b0417a5 100644 --- a/nodes.py +++ b/nodes.py @@ -27,7 +27,7 @@ try: except ImportError as e: checker = PackageDependencyChecker() discrepancies = checker.check_version_discrepancies('requirements.txt') - instructions = checker.generate_user_instructions(discrepancies) + #instructions = checker.generate_user_instructions(discrepancies) # Find dependent discrepancies for all packages with issues dependent_discrepancies = [] @@ -43,7 +43,7 @@ except ImportError as e: for suggestion in solution_suggestions: out += f"{suggestion}\n" - raise ValueError(f"<>: You Have some Environment Problems...\n\n{instructions}\n{out}") + raise ValueError(f"<>: You Have some Environment Problems...\n\n{out}") # Test current setup # Add in Audio2Audio @@ -241,8 +241,9 @@ def generate_audio(cond_batch, steps, cfg_scale, sigma_min, sigma_max, sampler_t output = output.to(torch.float32).div(torch.max(torch.abs(output))).clamp(-1, 1).mul(32767).to(torch.int16).cpu() print("Transformed Output:", output) + filepaths = None if save: - save_audio_files(output, sample_rate, save_prefix, counter, data=gendata) + filepaths = save_audio_files(output, sample_rate, save_prefix, counter, data=gendata) spectrogram = audio_spectrogram_image(output, sample_rate=sample_rate) @@ -253,7 +254,7 @@ def generate_audio(cond_batch, steps, cfg_scale, sigma_min, sigma_max, sampler_t torch.cuda.empty_cache() gc.collect() - return audio_bytes, sample_rate, spectrogram + return audio_bytes, sample_rate, spectrogram, filepaths def get_model(model_filename=None, config=None, repo=None, half_precision=False, device_override=None): @@ -333,24 +334,28 @@ def load_model(model_config=None, model_ckpt_path=None, pretrained_name=None, pr return model, model_config import shutil +from urllib.parse import quote def save_audio_files(output, sample_rate, filename_prefix, counter, data=None, save_temp=True): filename_prefix += "" output_dir = "output" os.makedirs(output_dir, exist_ok=True) wavname = filename_prefix if not data else replace_variables(filename_prefix, data) + filepaths = [] for i, audio in enumerate(output): if i > 0: # TODO fix batches break - fpath = f"{wavname}_{counter:04}.wav" + fpath = f"{quote(wavname)}_{counter:04}.wav" file_path = os.path.join(output_dir, fpath) print(f"Saving audio to {file_path}") torchaudio.save(file_path, audio.unsqueeze(0), sample_rate) + filepaths.append(fpath) # Saves to temporary path so it can be used for streaming loops if save_temp: tpath = os.path.join(TEMP_FOLDER, "stableaudiosampler.wav") print(f"Saving temp audio to: {tpath}") shutil.copyfile(file_path, tpath) counter += 1 + return filepaths from aeiou.viz import spectrogram_image @@ -372,8 +377,8 @@ class StableAudioSampler: return { "required": { "audio_model": ("SAOMODEL", {"forceInput": True}), - "positive": ("SAOCOND", {"forceInput": True}), - "negative": ("SAOCOND", {"forceInput": True}), + "positive": ("CONDITIONING", {"forceInput": True}), + "negative": ("CONDITIONING", {"forceInput": True}), "seed": ("INT", {"default": -1, "min": -1, "max": MAX_FP32}), "steps": ("INT", {"default": 100, "min": 1, "max": 10000}), "cfg_scale": ("FLOAT", {"default": 7.0, "min": 0.0, "max": 100.0, "step": 0.1}), @@ -398,9 +403,10 @@ class StableAudioSampler: CATEGORY = "audio/samplers" def sample(self, audio_model, positive, negative, seed, steps, cfg_scale, sigma_min, sigma_max, sampler_type, denoise, save, save_prefix, audio=None): - audio_bytes, sample_rate, spectrogram = generate_audio((positive, negative), steps, cfg_scale, sigma_min, sigma_max, sampler_type, device, save, save_prefix, audio_model, seed=seed, counter=self.counter, init_noise_level=denoise, init_audio=audio) + audio_bytes, sample_rate, spectrogram, filepaths = generate_audio((positive, negative), steps, cfg_scale, sigma_min, sigma_max, sampler_type, device, save, save_prefix, audio_model, seed=seed, counter=self.counter, init_noise_level=denoise, init_audio=audio) spectrograms = create_image_batch([spectrogram], 1) - return (audio_bytes, sample_rate, spectrograms) + return {"ui": {"paths": filepaths}, "result": (audio_bytes, sample_rate, spectrograms)} + #return (audio_bytes, sample_rate, spectrograms) class StableLoadAudioModel: @classmethod @@ -433,12 +439,12 @@ class StableAudioPrompt: def INPUT_TYPES(s): return { "required": { - "conditioning": ("SAOCOND", {"forceInput": True}), + "conditioning": ("CONDITIONING", {"forceInput": True}), "prompt": ("STRING", {"multiline": True}), } } - RETURN_TYPES = ("SAOCOND", ) + RETURN_TYPES = ("CONDITIONING", ) RETURN_NAMES = ("conditioning", ) FUNCTION = "go" @@ -475,7 +481,7 @@ class StableAudioConditioning: } } - RETURN_TYPES = ("SAOCOND", ) + RETURN_TYPES = ("CONDITIONING", ) RETURN_NAMES = ("conditioning", ) FUNCTION = "go" diff --git a/requirements.txt b/requirements.txt index 684a9f6..21d93a6 100644 --- a/requirements.txt +++ b/requirements.txt @@ -3,4 +3,5 @@ torch torchaudio einops numpy -aeiou \ No newline at end of file +aeiou +flash_attn \ No newline at end of file diff --git a/util_dependencies.py b/util_dependencies.py index f482a3b..e7b95e7 100644 --- a/util_dependencies.py +++ b/util_dependencies.py @@ -21,20 +21,18 @@ class PackageDependencyChecker: if not line or line.startswith('#'): return None - if '==' in line: - package_name, required_version = line.split('==') - return self._check_version(package_name, required_version) - else: - parts = line.split() - if len(parts) == 2: - package_name, version_spec = parts - return self._check_version(package_name, version_spec) - else: - return None + try: + requirement = pkg_resources.Requirement.parse(line) + package_name = requirement.name + version_spec = str(requirement.specifier) + return self._check_version(package_name, version_spec) + except pkg_resources.extern.packaging.requirements.InvalidRequirement: + return None def _check_version(self, package_name, version_spec): if package_name in self.installed_packages: installed_version = self.installed_packages[package_name].version + #print(f"Checking version: {package_name} {version_spec}") requirement = pkg_resources.Requirement.parse(f"{package_name} {version_spec}") if not requirement.specifier.contains(installed_version): dependents = self.find_dependents(package_name) @@ -86,15 +84,20 @@ class PackageDependencyChecker: for dependent in dependents: for requirement in dependent[2]: if package_name in requirement: - requirement_version = requirement.split(package_name)[-1].strip() - if requirement_version: - discrepancy = self._check_version(package_name, requirement_version) - if discrepancy: - discrepancies.append(discrepancy) + try: + req = pkg_resources.Requirement.parse(requirement) + if req.name == package_name: + requirement_version = str(req.specifier) + discrepancy = self._check_version(package_name, requirement_version) + if discrepancy: + discrepancies.append(discrepancy) + except pkg_resources.extern.packaging.requirements.InvalidRequirement: + continue return discrepancies def analyze_discrepancies(self, discrepancies): """Analyzes discrepancies and finds possible version ranges.""" + print("Analyzing package discrepancies...") requirements = defaultdict(list) for discrepancy in discrepancies: requirements[discrepancy['package_name']].append(discrepancy['required_version']) @@ -152,18 +155,55 @@ class PackageDependencyChecker: break return versions + def analyze_possible_upgrades(self, discrepancies): + possible_upgrades = {} + for discrepancy in discrepancies: + package_name = discrepancy['package_name'] + current_version = discrepancy['installed_version'] + # Fetch possible upgrades from pip + available_versions = self.fetch_available_versions(package_name) + for version in available_versions: + if version > current_version: + impact = self.calculate_upgrade_impact(package_name, version, discrepancies) + possible_upgrades[(package_name, version)] = impact + return possible_upgrades + + def calculate_upgrade_impact(self, package_name, version, discrepancies): + # Simulate upgrading the package and recheck discrepancies + new_discrepancies = self.simulate_upgrade(package_name, version, discrepancies) + return len(discrepancies) - len(new_discrepancies) + + def simulate_upgrade(self, package_name, new_version, discrepancies): + original_version = self.installed_packages[package_name].version + self.installed_packages[package_name] = pkg_resources.Distribution(project_name=package_name, version=new_version) + new_discrepancies = [] + for discrepancy in discrepancies: + if discrepancy['package_name'] != package_name: + new_discrepancy = self._check_version(discrepancy['package_name'], discrepancy['required_version']) + if new_discrepancy: + new_discrepancies.append(new_discrepancy) + self.installed_packages[package_name] = pkg_resources.Distribution(project_name=package_name, version=original_version) + return new_discrepancies + def suggest_solutions(self, solutions): """Generates solution suggestions based on the analyzed discrepancies.""" suggestions = [] for package, version_range in solutions.items(): + dependents = self.find_dependents(package) + if version_range: min_version, max_version, excluded_versions = version_range excluded_versions_str = ", ".join(str(v) for v in sorted(excluded_versions)) + if excluded_versions: suggestions.append(f"Package '{package}' can be installed in the version range {min_version} - {max_version}, excluding versions: {excluded_versions_str}.") else: suggestions.append(f"Package '{package}' can be installed in the version range {min_version} - {max_version}.") - + + suggestions.append("This version range is required by the following packages:") + for dependent in dependents: + suggestions.append(f" - {dependent[0]} (version {dependent[1]})") + # Fetch available versions from pip available_versions = self.fetch_available_versions(package) valid_versions = [ @@ -181,7 +221,29 @@ class PackageDependencyChecker: else: suggestions.append(f"No valid versions found for '{package}' within the specified range.") else: - suggestions.append(f"No common version range found for package '{package}'.") + suggestions.append(f"No common version range found for package '{package}'. However, you can try the following versions required by other packages:") + package_commands = defaultdict(list) + for dependent in dependents: + for requirement in dependent[2]: + if package in requirement: + suggestions.append(f" - {requirement} required by {dependent[0]} (version {dependent[1]})") + available_versions = self.fetch_available_versions(package) + try: + req = pkg_resources.Requirement.parse(requirement) + for ver in available_versions: + ver_parsed = version.parse(ver) + if str(ver_parsed) in str(req.specifier): + package_commands[(dependent[0], dependent[1])].append(f" - pip install {package}=={ver}") + except (ValueError, IndexError): + continue + + if package_commands: + for packages, commands in package_commands.items(): + package_list = f"{packages[0]} (version {packages[1]})" + suggestions.append(f"To get the package {package_list} working, you can use the following commands:") + suggestions.extend(commands) + else: + suggestions.append(f"No valid versions found for '{package}' based on individual requirements.") return suggestions def filter_latest_versions(self, versions): @@ -193,6 +255,44 @@ class PackageDependencyChecker: latest_versions[major_minor] = ver return latest_versions.values() + def rank_solutions(self, solutions): + """Ranks version ranges based on the number of discrepancies they resolve.""" + version_range_counts = {} + for package, version_range in solutions.items(): + if version_range: + min_version, max_version, _ = version_range + range_key = (min_version, max_version) + if range_key not in version_range_counts: + version_range_counts[range_key] = 0 + version_range_counts[range_key] += 1 + + ranked_solutions = sorted(version_range_counts.items(), key=lambda x: -x[1]) + return ranked_solutions + + def rank_upgrades(self, possible_upgrades): + """Ranks upgrades based on the number of discrepancies they resolve.""" + ranked_upgrades = sorted(possible_upgrades.items(), key=lambda x: -x[1]) + return ranked_upgrades + + def find_best_upgrade_path(self, discrepancies): + """Finds the best upgrade path to resolve the most discrepancies.""" + solutions = self.analyze_discrepancies(discrepancies) + possible_upgrades = self.analyze_possible_upgrades(discrepancies) + ranked_solutions = self.rank_solutions(solutions) + ranked_upgrades = self.rank_upgrades(possible_upgrades) + + best_path = [] + + # Combine ranked solutions and upgrades to find the optimal path + for upgrade, impact in ranked_upgrades: + package, version = upgrade + best_path.append(f"Upgrading {package} to version {version} resolves {impact} discrepancies.") + + for version_range, count in ranked_solutions: + best_path.append(f"Installing packages in the version range {version_range} resolves {count} discrepancies.") + + return best_path + if __name__ == "__main__": import sys @@ -202,10 +302,10 @@ if __name__ == "__main__": # Find dependents of a specific package package_name = sys.argv[1] dependents = checker.find_dependents(package_name) - print(f"Dependents of '{package_name}':") + print(f"Dependents of '{package_name}':\n") for dependent in dependents: print(f"Package: {dependent[0]}, Version: {dependent[1]}") - print(f"Requires: {dependent[2]}") + print(f" Requires: {dependent[2]}") # Check version discrepancies based on requirements.txt # discrepancies = checker.check_version_discrepancies('requirements.txt') @@ -223,16 +323,16 @@ if __name__ == "__main__": # Check version discrepancies among dependents of the specific package dependent_discrepancies = checker.check_dependents_discrepancies(package_name) - if dependent_discrepancies: - print(f"\nDiscrepancies in dependents of '{package_name}':") - for discrepancy in dependent_discrepancies: - print(f"Package: {discrepancy['package_name']}, Installed: {discrepancy['installed_version']}, Required: {discrepancy['required_version']}") - print("Dependents:") - for dependent in discrepancy['dependents']: - print(f" - Dependent Package: {dependent[0]}, Version: {dependent[1]}") - print(f" Requires: {dependent[2]}") - else: - print(f"No discrepancies found among the dependents of '{package_name}'.") + # if dependent_discrepancies: + # print(f"\nDiscrepancies in dependents of '{package_name}':") + # for discrepancy in dependent_discrepancies: + # print(f"Package: {discrepancy['package_name']}, Installed: {discrepancy['installed_version']}, Required: {discrepancy['required_version']}") + # print("Dependents:") + # for dependent in discrepancy['dependents']: + # print(f" - Dependent Package: {dependent[0]}, Version: {dependent[1]}") + # print(f" Requires: {dependent[2]}") + # else: + # print(f"No discrepancies found among the dependents of '{package_name}'.") # Analyze discrepancies and suggest solutions solutions = checker.analyze_discrepancies(dependent_discrepancies) @@ -240,3 +340,7 @@ if __name__ == "__main__": print("\nSuggested solutions:") for suggestion in solution_suggestions: print(suggestion) + + best_upgrade_path = checker.find_best_upgrade_path(dependent_discrepancies) + for up in best_upgrade_path: + print(up) \ No newline at end of file diff --git a/web/js/playSound.js b/web/js/playSound.js new file mode 100644 index 0000000..7319166 --- /dev/null +++ b/web/js/playSound.js @@ -0,0 +1,37 @@ +/* + shouts to: pygoss, Fill, and Joviex! +*/ +import { app } from "../../../scripts/app.js"; + +console.log("StableAudioSampler") + +app.registerExtension({ + name: "lks-ai.StableAudioSampler", + async beforeRegisterNodeDef(nodeType, nodeData, app) { + if (nodeData.name === "StableAudioSampler") { + console.log(app); + console.log(nodeData); + const onExecuted = nodeType.prototype.onExecuted; + //console.log(onExecuted); + nodeType.prototype.onExecuted = async function (message) { + onExecuted?.apply(this, arguments); + // console.log(this.widgets); + // console.log(app.ui.lastQueueSize); + // console.log(message) + + // TODO can check this.widgets[] for specific controls + let file = message.paths[0]; + if (!file) { + file = "temp/stableaudiosampler.wav"; + } + + const url = new URL(`http://localhost:8188/view?filename=${encodeURIComponent(file)}&subfolder=&type=output&format=audio%2Fwav`); + console.log(import.meta.url) + console.log(url) + const audio = new Audio(url); + audio.volume = 1.0; //this.widgets[1].value; + audio.play(); + }; + } + }, +}); \ No newline at end of file