This commit is contained in:
holonic
2024-06-17 11:34:23 +01:00
parent f88fc7c02f
commit b55c712b8a
5 changed files with 191 additions and 43 deletions
+1 -1
View File
@@ -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"
+18 -12
View File
@@ -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"<<StableAudioSampler>>: You Have some Environment Problems...\n\n{instructions}\n{out}")
raise ValueError(f"<<StableAudioSampler>>: 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"
+2 -1
View File
@@ -3,4 +3,5 @@ torch
torchaudio
einops
numpy
aeiou
aeiou
flash_attn
+133 -29
View File
@@ -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)
+37
View File
@@ -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();
};
}
},
});