🤦
This commit is contained in:
+1
-1
@@ -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"
|
||||
@@ -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
@@ -3,4 +3,5 @@ torch
|
||||
torchaudio
|
||||
einops
|
||||
numpy
|
||||
aeiou
|
||||
aeiou
|
||||
flash_attn
|
||||
+133
-29
@@ -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)
|
||||
@@ -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();
|
||||
};
|
||||
}
|
||||
},
|
||||
});
|
||||
Reference in New Issue
Block a user