7 ms·
Speedup from switch to +=
- teruakohatu 4y agoI guess this is the beauty of making a model open source.
- myrryr 4y agoit is a hell of a good case study that is for sure.
- mhzsh 4y agoBut why is it faster? A non-associative translation to byte code (or however python works)?
- onedognight 4y agoMy guess is that it operates in place with no memory allocations or copying.
- actually_a_dog 4y agoNot exactly: >>> def f(x): x += 1 ... >>> def g(x): x = x + 1 ... >>> dis.dis(f) 1 0 LOAD_FAST 0 (x) 3 LOAD_CONST 1 (1) 6 INPLACE_ADD 7 STORE_FAST 0 (x) 10 LOAD_CONST 0 (None) 13 RETURN_VALUE >>> dis.dis(g) 1 0 LOAD_FAST 0 (x) 3 LOAD_CONST 1 (1) 6 BINARY_ADD 7 STORE_FAST 0 (x) 10 LOAD_CONST 0 (None) 13 RETURN_VALUE
- lnyan 4y agoFor PyTorch, `+=` is interpreted as an in-place operation
- eru 4y agoI wonder what version of Python they were using? I'm wondering, because recent version have improved performance a lot. 3.11 is much faster than 3.10, and what's in 3.12 is already much faster than 3.11.
- eminence32 4y agoThe upstream Stable Diffusion uses python 3.8: https://github.com/CompVis/stable-diffusion/blob/69ae4b35e0a0f6ee1af8bb9a5d0016ccb27e36dc/environment.yaml#L6 https://github.com/CompVis/stable-diffusion/blob/69ae4b35e0a...
- eru 4y agoThanks.
- deleted 4y ago[deleted]
- dahfizz 4y agoIs python in the fast path? Why not rewrite in a performant language for a XXX% speedup?
- bee_rider 4y agoThe += operator is almost certainly calling some method on sends out the real work to some tuned hardware-specific framework written in a fast language.
- dahfizz 4y agoSo python is marshalling data to and from an ffi in the fast path? That sounds even worse
- bee_rider 4y agoI'm not sure this is the conventional use of the phrase "fast path." But anyway, the idea is usually that the Python code calls out to the framework with operations that are in some way "large," and so the overhead is not so significant. Python probably doesn't have to do any marshaling, hypothetically the framework could just return an object that represents a pointer. Then the python code sends that pointer to another framework method.
- kjeetgill 4y agoI think you mean critical path (which is ironically usually the slowest path). A fast path is usually a hardcoded shortcut you can take for select cases.
- pclmulqdq 4y agoNot exactly: most of these frameworks essentially JIT compile the entire operation graph so that it can be executed, and the Python code only touches the data at the endpoints of the full computation. I don't know why the JIT compiler doesn't optimize a = b + a to a += b, but I guess they assumed that the JIT-ed code path would only be used once, so the compiler has to be fast.
- 4y ago
- JonathonW 4y agoIf they're seeing these kinds of gains from relatively minor changes to their Python code, I can't help but wonder how much faster the model would run in a compiled language or a language with a good JIT (way more optimization work's gone into the mainstream Javascript runtimes than CPython). I'd assumed that overall performance in Stable Diffusion was limited by the code running on the GPU, with Python performance being a fairly minor factor-- but I guess that's not the case?
- deleted 4y ago[deleted]
- sshine 4y agoI don't know anything about stable diffusion, but I've been optimizing a lot of prime-field arithmetic in Rust lately, and we experienced a similar speedup going from `+ x` to `+= x` (for scalars and especially for composite structures like vectors and polynomials).
- thayne 4y agoFor composite structures that isn't too surprising, but for scalars, I would have expected llvm to optimize the addition and assignment into a single in place addition.
- codeflo 4y agoThe short answer is yes. The long answer is that it’s not so clear what an “in place addition” even means at the level of CPU instructions after you consider register allocation. For example, if you have v = x; v += y; f(v); and the never mention v again, then the whole operation is performed directly in the register that is specified to receive the first argument in a function call, not in whatever register might have been allocated for v. That’s because, with some complications I don’t want to go into, compilers look at the dependency graph of values rather than at the variable names.
- bee_rider 4y ago
- brrrrrm 4y agoIt’s not clear a JIT compiled language would help much here unless the operations were baked into the JIT itself (which would have to identify the memory savings of an in-place call).
- Waterluvian 4y agoOne comment asks about putting it all on one line, and this is where interpreted languages without a JIT kinda blow. Many times I have had to decide if my Python code would be more legible or get free performance. The thing I like about JavaScript is that I can _usually_ trust the JIT to make my code faster than I could, meaning I can focus entirely on writing clean code. P.S. you can always hand optimize. If you do, just comment the heck out of it.
- NavinF 4y agoThis has nothing to do with python. A JITed/AoT compiled version of the old code should do exactly the same thing because it would build the same pytorch graph.
- nodja 4y ago> Many times I have had to decide if my Python code would be more legible or get free performance. This is rarely an option that has presented itself to me. If there's a clear performance issue in my code then I probably picked the wrong algorithm or my code has a bug, unless you decided for some reason to do heavy calculations in raw python. If you're doing operations on big chunks of data you should always use something like numpy or jax. Even OPs issue the clear reason is that it's doing an operation in place instead of creating a copy, for ML models this can only be done at inference time and not training time since you need to keep track of the whole network, hence why the code was in it's unoptimized state.
- ironhaven 4y agoBecause of operator overloading "+=" can call a more optimized method than "+". If this code was written in a language without operator overloading I don't think this would be a very interesting pull request. THis could be a example of why some people don't like operator overloading and why some programing languages (java, zig, etc) do not implment the feature.
- noobermin 4y agoIf python did not have operator overloading it would not be used for numeric programming to the extent it is. Overloading is key to its success in that field. The problem is thinking `+' and `+=' are the same, they are not and `+' should not be used when `+=' can be used.
- staticassertion 4y agoI don't think this is an operator overloading thing? It's just that `x = y + x` is equivalent to z = y + x x = z Basically, creating an object `z` just to throw it away. `x += y` just adds y to x directly without any intermediary. You could write this in any language pretty easily. For example, in Rust: let x = "abc".to_string(); let y = "123".to_string(); let x = x + &y; as opposed to the more efficient: let mut x = "abc".to_string(); let y = "123".to_string(); x.push_str(&y); It's just using an operation to mutate in place vs an immutable operation.
- masklinn 4y ago> I don't think this is an operator overloading thing? It’s the confusion / idea that this is trivial change which is the overload thing.
- NavinF 4y agoOperator overloading is a major reason why libraries like pytorch exist so IMO that's a moot point. Btw there's ongoing work to automatically optimize expressions like this. See the XLA compiler for example. Right now deep learning has a ton of seemingly obvious compute/memory optimisations that are not done automatically.
- datalopers 4y agoThis StackOverflow answer [1] goes into performance details of INPLACE_ADD versus BINARY_ADD. [1] https://stackoverflow.com/a/15376520 https://stackoverflow.com/a/15376520
- noobermin 4y agoWhenever I see things like this in highly visible code that people exclaim about across the internet it makes me really take a moment to absorb how much time I spend agonizing over minutae in my daily work and how people who really are just lucky can get away with much worse. Just a reminder about how the idea that "tech" is a meritocracy was never really true.
- WatchDog 4y agoI assume that you don't have thousands of people looking over your code, how can you know that it doesn't have similar or greater room for optimization?
- staticassertion 4y agoThis isn't a Python issue, this is a "I'm copying when I don't need to" issue. As I mention elsewhere, you can write this sort of "bug" in almost any language pretty easily (as I demonstrate with Rust). This isn't a case of "The Python interpreter is bad" it's just that the code is doing what the user asked it to do - create a completely new copy of the data, then overwrite the old copy with it. Immutable operations like this are slow, mutating the value (what += does) is fast. Granted, a compiled language could recognize that you're doing this, but it also might not - is `+` and `+=` semantically identical such that the compiler can replace one with the other? Maybe? Probably not, if I had to guess. The correct answer is to just use the faster operation, as it is with all language. I don't know the type of `x`, but I'd suggest another optimization here would be to: a) Preallocate the buffer rather before mutating it 3x (which is still likely forcing some allocations) b) Reuse that buffer if it's so important, store it in `self` and clear it before use.
- FabHK 4y agoPlot twist: it breaks the code...? > Changing this back to the original implementation fixed an error I was getting when doing textual inversion on Windows https://github.com/lstein/stable-diffusion/commit/62863ac586194a43ff952eba17a83cecf9956500#commitcomment-83696307 https://github.com/lstein/stable-diffusion/commit/62863ac586...
- staticassertion 4y agoLove to see it. A perfect example of why this optimization can't be done automatically - in the case of `else` you're working with a mutable reference to `x` passed in, which means that now your function is mutating something it used to not mutate. A "safe" way to do this is still straightforward, I think. from copy import copy def _forward(self, x, context=None): x = x.contiguous() if x.device.type == 'mps' else x x = copy(x) x += self.attn1(self.norm1(x)) x += self.attn2(self.norm2(x), context=context) x += self.ff(self.norm3(x)) return x It could be faster but I don't know what `x` is and I'm not going to guess. Also, `copy` may not be sufficient, `deepcopy` may be necessary - again, I don't know what `x` is so I can't figure that out. Pls use type annotations :)
- mgraczyk 4y agoThat's not safe if the problem is the in place mutation. You will still mutate x while reading from it.
- desmond373 4y agoMight be platform dependent whether that first line counts as a mutate or not, seeing as it can be converted to not do anything in some cases.
- mgraczyk 4y agoAll of the lines that have += mutate x in place
- deleted 4y ago[deleted]
- olliej 4y agoIs this a lookup overhead thing or a memcpy based overhead regression? In the case of the latter it seems like this may result in an unexpected mutation of the source data?
- deleted 4y ago[deleted]
- thweorui23432 4y agoSpeedup likely won't work for training the model.
- NavinF 4y agoYep, intermediate results (activations) are kept in memory during training.
- teo_zero 4y agoBut wait... x+=y is equivalent to x=x+y not to x=y+x. Only if + is commutative, then the three are equivalent. Are we sure the + operation is commutatve for this type of data? And does the compiler know it? It would be interesting to check whether changing every expression to x=x+y has a performance more similar to += or to ...+x
- nodja 4y agoI see lots of people answering why it's faster, but not many saying why the engineers chose the slower version. As everyone said, this is more performant because x is being modified in place, the reason this was not done in place is because you can't train a neural network if an instruction is being done in place. During training a network goes literally through all operations that were done and see how well they performed so they can be adjusted using a secondary value called a gradient, this is done during the backwards pass. If you replace something in place you're essentially overwriting the input values that were passed to that function, and by extension, the output values of the function called before, essentially breaking the network chain, unless you also copy the inputs together with the gradients, which would cause an even worse performance hit and be a memory hog. The breakage bug later in the issue is proof of this, when sampling to generate an image only the forward pass is done on the network, but textual inversion requires you to train the network and therefore do the backwards pass, triggering the error since the dependency graph is broken. I should also note that technically the add operation should be safe to do in place as it's reversible, but I'm not a pytorch expert so I'm not sure exactly what's going on in there.
- umvi 4y agoSee, this is a great example of where a comment needed to be added, but wasn't. If the engineers that originally implemented the function intentionally chose the slower version, a quick comment as to why would have prevented this from happening in the first place.
- hbogert 4y agoWill you be my colleague please? This is idd the time to place a comment, yet so many people don't do that.
- nodja 4y agoThis is common knowledge, so common that someone that hasn't coded anything besides some basic linear regression model like me knows about it. It's like commenting on why you'd put parenthesis in some formula, it's just gonna say "parenthesis here because this operation takes priority", similarly in a pytorch model, if it was done by those standards the code would be filled with "operation not done in place because it would break the network graph". You're more likely to encounter the opposite comment, "doing this operation in-place because it'll be discarded later" or something along those lines. One of the first things you're taught when learning pytorch is that you're not coding in python, but actually creating a network graph that is loaded and executed on a GPU. Other common sense things is knowing that you shouldn't use stuff that is in the stdlib or in numpy and use torch.* variants instead, not doing so will incur either undefined behavior, cause massive memory copies between the CPU and GPU or most likely, error out at runtime. Note that this is a repo that is forked from the official repo, it's a community repo focused on inference and thus doesn't care about training so it has completely different considerations than the original code.
- MaXtreeM 4y agoThere is a case in C# where using compound assignment is actually slower [0]. Based on comments this should be fixed in .NET7 I haven't checked it myself. [0]: https://mobile.twitter.com/badamczewski01/status/1561817158442782720 https://mobile.twitter.com/badamczewski01/status/15618171584...
- chillee 4y agoOk, I work on PyTorch, so probably should clear up some misconceptions in this thread. 1. In PyTorch (and other array programming libraries like Numpy), the operations being passed around are tensors/arrays (i.e. large chunks of memory). Thus, += is overloaded to mean "in-place write" to the arrays. So, `+` vs `+=` is the equivalent of a: float[1000] b: float[1000] for i in [0, 1000]: b[i] = a[i] + 2 vs. a: float[1000] for i in [0, 1000]: a[i] = a[i] + 2 The main performance advantage comes in 1. no need to allocate an extra array, 2. you're using less memory overall, so various caching levels can work better. It has nothing to do with python bytecodes. 2. As for whether it generally makes sense to do this optimization manually... Usually, PyTorch users don't use in-place operations as its a bit uglier mathematically and have various foot-guns/restrictions that users find confusing. Generally, it's best to have this optimization be done automatically by an optimizing compiler. 3. PyTorch in general does support using in-place operations during training, albeit with some caveats. (PS) 4. Putting everything on one line (as some folks suggest) is almost certainly not going to help performance - the primary performance bottlenecks here have almost nothing to do with CPU perf.
- teruakohatu 4y agoThanks for the input. Before I start throwing += into my PyTorch code can you explain what you mean here: > Generally, it's best to have this optimization be done automatically by an optimizing compiler. What compiler should be optimizing this operation? There are comments on the commit reporting errors under certain conditions.
- chillee 4y agoTo clarify, by "compilers" I mean "deep learning compilers". There's many different paths to optimizing compilers folks use with PyTorch. One with close integration is NVFuser (see https://www.reddit.com/r/MachineLearning/comments/xa75km/p_pytorchs_newest_nvfuser_on_stable_diffusion_to/ https://www.reddit.com/r/MachineLearning/comments/xa75km/p_p...), although there are other compilers like ONNXRuntime. Yes, handling autograd (during training) is a whole different thing, and not all compilers support that.
- eesmith 4y agoLincoln Stein. Now that's a name I've not heard in a long time. A long time. He's the author of the essay "How Perl Saved the Genome Project", the books "Network Programming with Perl" and "Writing Apache Modules with Perl and C", and a number of Perl packages including CGI.pm - which helped power the dot-com era - and GD.pm.
- spullara 4y agoMutation faster than making a new object.