8 ms·
Keras Core: Keras for TensorFlow, Jax, and PyTorch
- dewitt 3y agoFrom the announcement: "We're excited to share with you a new library called Keras Core, a preview version of the future of Keras. In Fall 2023, this library will become Keras 3.0. Keras Core is a full rewrite of the Keras codebase that rebases it on top of a modular backend architecture. It makes it possible to run Keras workflows on top of arbitrary frameworks — starting with TensorFlow, JAX, and PyTorch." Excited about this one. Please let us know if you have any questions.
- albertzeyer 3y agoThat looks very interesting. I actually have developed (and am developing) sth very similar, what we call the RETURNN frontend, a new frontend + new backends for our RETURNN framework. The new frontend is supporting very similar Python code to define models as you see in PyTorch or Keras, i.e. a core Tensor class, a base Module class you can derive, a Parameter class, and then a core functional API to perform all the computations. That supports multiple backends, currently mostly TensorFlow (graph-based) and PyTorch, but JAX was something I also planned. Some details here: https://github.com/rwth-i6/returnn/issues/1120 https://github.com/rwth-i6/returnn/issues/1120 (Note that we went a bit further ahead and made named dimensions a core principle of the framework.) (Example beam search implementation: https://github.com/rwth-i6/i6_experiments/blob/14b66c4dc74c0830ab92343f39d0cb181771098e/users/zeyer/experiments/exp2023_04_25_rf/conformer_import_moh_att_2023_04_24_BxqgICRSGkgb.py#L359 https://github.com/rwth-i6/i6_experiments/blob/14b66c4dc74c0...) One difficulty I found was how design the API in a way that works well both for eager-mode frameworks (PyTorch, TF eager-mode) and graph-based frameworks (TF graph-mode, JAX). That mostly involves everything where there is some state, or sth code which should not just execute in the inner training loop but e.g. for initialization only, or after each epoch, or whatever. So for example: - Parameter initialization. - Anything involving buffers, e.g. batch normalization. - Other custom training loops? Or e.g. an outer loop and an inner loop (e.g. like GAN training)? - How to implement sth like weight normalization? In PyTorch, the module.param is renamed, and then there is a pre-forward hook, which on-the-fly calculates module.param for each call for forward. So, just following the same logic for both eager-mode and graph-mode? - How to deal with control flow context, accessing values outside the loop which came from inside, etc. Those things are naturally possible eager-mode, where you would get the most recent value, and where there is no real control flow context. - Device logic: Have device defined explicitly for each tensor (like PyTorch), or automatically eagerly move tensors to the GPU (like TensorFlow)? Moving from one device to another (or CPU) is automatic or must be explicit? - How to you allow easy interop, e.g. mixing torch.nn.Module and Keras layers? I see that you have keras_core.callbacks.LambdaCallback which is maybe similar, but can you effectively update the logic of the module in there?
- kerasteam 3y agoI worked on the project, happy to answer any questions!
- ayhanfuat 3y agoThis is an amazing contribution to the NN world. Thank you all the team members.
- dbish 3y agoGreat to see this, but I’m curious, does this mean we’ll get fewer fchollet tweets that talk up TF and down PyTorch? Is the rivalry done?
- sbrother 3y agoThis looks awesome; I was a big fan of Keras back when it had pluggable backends and a much cleaner API than Tensorflow. Fast forward to now, and my biggest pain point is that all the new models are released on PyTorch, but the PyTorch serving story is still far behind TF Serving. Can this help convert a PyTorch model into a servable SavedModel?
- kerasteam 3y agoFor a Keras Core model to be usable with the TF Serving ecosystem, it must be implemented either via Keras APIs (Keras layers and Keras ops) or via TF APIs. To use pretrained models, you can take a look at KerasCV and KerasNLP, they have all the classics, like BERT, T5, OPT, Whisper, StableDiffusion, EfficientNet, YOLOv8, etc. They're adding new models regularly.
- binarymax 3y agoCongrats on the launch! I learned Keras back when I first got in to ML, so really happy to see it making a comeback. Are there some example architectures available/planned that are somewhat complex, and not just a couple layers (BERT, ResNet, etc.)?
- kerasteam 3y ago
- Narew 3y agoKeras was already that some years ago. It supported tensorflow, theano, mxnet if my memory is right. And then they ditched everything for tensorflow. At the time it was really hard to use keras without calling backend directly for lots of optimisation, unsupported feature on they API etc... This make the use of Keras not agnostic at all. What's different now ?
- minimaxir 3y ago> What's different now ? PyTorch adoption: back when Keras went hard into TensorFlow in 2018, both TF and PyTorch adoption were about the same with TF having a bit more popularity. Now, most of the papers and models released are PyTorch-first.
- Narew 3y agoYes I understand why they do the move (they want to attract pytorch user). What's the benefit for the user instead of directly using pytorch for example ? I see we can maybe use tpu by switching to jax etc... PS: sorry I'm a bit salty by my user experience of Keras.
- minimaxir 3y agoKeras has a cleaner API compared to base PyTorch, especially if you want to use the Sequential construction as demoed in the post.
- sva_ 3y agoHow so? You can use torch.nn.Sequential pretty much equivalently? https://pytorch.org/docs/stable/generated/torch.nn.Sequential.html https://pytorch.org/docs/stable/generated/torch.nn.Sequentia...
- minimaxir 3y agoHuh, didn't realize base PyTorch had an equivalent Sequential API. The point about the better API overall still stands (notably including the actual training part, as base PyTorch requires you to implement your own loop)
- syntaxing 3y agoDoes this mean the weights output can be backend agnostic? Also, are there any examples using this for the coral TPU?
- kerasteam 3y agoYes, model weights saved with Keras Core are backend-agnostic. You can train a model in one backend and reload it in another. Coral TPU could be used with Keras Core, but via the TensorFlow backend only.
- syntaxing 3y agoSuper cool, does that mean if someone trains something using the PyTorch backend, I can still use it with the coral if I load the weights using the tensorflow backend?
- kerasteam 3y agoThat's right, if the model is backend-agnostic you can train it with a PyTorch training loop and then reload it and use it with TF ecosystem tools, like serve it with TF-Serving or export it to Coral TPU.
- deleted 3y ago[deleted]
- deleted 3y ago[deleted]
- voz_ 3y agoIf I want to use this brave new keras with torch.compile, what does that look like?
- haifeng-jin 3y agoWe are still working on this feature. We try to have it in model.compile(jit_compile=True). https://github.com/keras-team/keras-core/blob/v0.1.0/keras_core/backend/torch/trainer.py#L109-L110 https://github.com/keras-team/keras-core/blob/v0.1.0/keras_c...
- bootsmann 3y agoWait so what happens if I use a model with torch backend now and call .compile()? Does it just return and then do normal torch jit when .fit() (or whatever the keras notation is, i have forgotten most of it) is called?
- voz_ 3y agoA lot of this seems like abstraction for abstractions sake. When would someone actually use this?
- Narew 3y agoSame question for the decorator tf.function ?
- jszymborski 3y agoKeras and PyTorch! I thought I'd never see the day! Glad to see the two communities bury the hatchet.
- p1esk 3y agoI don’t get it - why would you want Keras if you already use Pytorch?
- pilotneko 3y agoBecause sometimes you don’t want to write your own training loops, you just want a working method to train a model.
- thangngoc89 3y agoThere are a lot of libraries for that. For example Pytorch Lightning, Accelerate are very mature
- jszymborski 3y agoSure, and Keras is another, very mature library which allows you to do this...
- p1esk 3y agoKeras + Pytorch is not mature.
- jszymborski 3y agoThe same reason why you might want to use Keras if you use any of the other backends. They operate at different levels. Keras is a higher-level API. It means that you can prototype architectures quickly and you don't have to write a training loop. It's also really easy to extend. I currently use PyTorch Lightning to avoid having to write tonnes of boilerplate code, but I've been looking for a way to leave it for ages as I'm not a huge fan of the direction of the product. Keras seems like it might be the answer for me.
- riku_iki 3y agowill keras be backward compatible, or as always and now google/tf ecosystem will have 3 gens of frameworks: tf1, tf2 + keras, keras core 3.
- math_dandy 3y agoIIRC, Keras was added officially added to Tensorflow as part of the version 2.0 release. With Keras reverting to its backend-agnostic state, will it be removed from Tensorflow? Is this a divorce or are TF & Keras just opening up their relationship?
- dharmeshkakadia 3y agoSupporting multiple backends (especially Jax) is nice! Makes experimenting/migrating between them so much more approachable. Any timeline on when can we expect support for distributed Jax training? The doc currently seems to indicate only TF is supported for distributed training.
- martin-gorner 3y agoSupport for distributed JAX training demoed here: bit.ly/keras-on-jax-demo You have to write a custom training loop for now, but it works.
- dharmeshkakadia 3y agoThanks!
- boredumb 3y agoI must admit i've never actually used keras but this is interesting to see how they are implementing it with Jax, definitely worth a reminder to dig into one of these days.
- adolph 3y agoWouldn't a multi-framework wrapper be a subset of any supported framework's features common among all frameworks? Additionally would it always be at least a step behind any framework depending on the wrapper's release cycle?
- dkga 3y agoI think that is pretty cool - literally made me screen "Yes!" when I saw it and I don't do this for your everyday framework. I think the beauty of keras was the perfect balance between simplicity/abstraction and flexibility. I moved to PyTorch eventually but one thing I always missed was this. And now, to have it leapfrog the current fragmentation and just achieve what seems to be a true multi-backend is pretty awesome. Looking forward to the next steps!
- m_ke 3y agoAs someone who has dealt with countless breaking changes in keras and wasted days of my life attempting to upgrade, no thank you. My pytorch code from years ago still works with no issues, my old keras code would break all the time even in minor releases.
- ipunchghosts 3y agoAgreed. This will only break things, especially research code.
- kaycebasques 3y agoCan someone ELI5 the relationship between Keras and TensorFlow/Jax/PyTorch/etc? I kinda get the idea the Keras is the "frontend" and TF/Jax/PyTorch are the "backend" but I'm looking to solidify my understanding of the relationship. It might help to also comment on the key differences between TF/Jax/PyTorch/etc. Thank you.
- _Wintermute 3y agoKeras was a high-level wrapper around Theano or Tensorflow. The creator of Keras was then employed by Google to work on Keras, who promised everyone Keras would remain backend agnostic. Keras become part of Tensorflow as a high-level API and did not remain backend agnostic. There was lots of questionable twitter beef about Pytorch by the Keras creator. Keras is now once again backend agnostic, as a high-level API for Tensorflow/PyTorch/Jax. Likely as them seeing Tensorflow losing traction.
- martin-gorner 3y agodemo on JAX: https://bit.ly/keras-on-jax-demo https://bit.ly/keras-on-jax-demo
- ipunchghosts 3y agoLike everything keras, this will promise a lot and only deliver on conops deemed worthy by the keras team.